diff --git a/Q Research/README.md b/Q Research/README.md new file mode 100644 index 0000000..6d1a621 --- /dev/null +++ b/Q Research/README.md @@ -0,0 +1,41 @@ +# QuantResearch + +## 数据与回测提交须知 + +- 任何修改 `data/raw`、`data/derived`、`results/data_quality` 时,务必重新运行: + ```bash + python scripts/build_dataset_manifest.py --dirs data/raw data/derived --output data/_manifest.json + python scripts/check_data_integrity.py + ``` +- 如果 `data/signature_baseline.json` 中的哈希需要更新,请在 PR 描述中写明原因、影响范围,并附上新的 `results/data_quality/*.json` 报告路径。 +- 回测脚本会在 `results//summary.json` 记录本次运行的 KPI + 数据签名,提交代码时请一并引用该 run_id 便于审核。 +- 提交前可执行 `python scripts/validate_results.py results/`,快速检查 KPI 字段、数据报告引用是否完整。 + +## 风控与 Diagnostics 流程 + +1. **风险仿真守门** + ```bash + cd QuantResearch + RUN= ./scripts/run_risk_sim.sh + ./bin/backfill_risk.sh # 确保 results/risk/metrics.csv 同步最新 run + ``` + `run_risk_sim.sh` 默认使用 `QuantTrader/config/risk_limits_sim.yaml`(宽松额度,仅用于数据 gating)。在 PR 描述中标注对应 run_id,并说明若需要引用实盘限额,应改用 `QuantTrader/config/risk_limits.yaml`。 + > CI 会在 diagnostics workflow 中调用 `python scripts/watch_risk_metrics.py`;若 metrics 中仍存在 fail/run 缺失,PR 会直接失败,因此务必在本地先修复再推送。 + +2. **诊断图表** + ```bash + ./scripts/run_ci_diagnostics.sh # 自动挑选最新 batch/walkforward/MC 输出 + ``` + 图表保存在 `charts/ci/diagnostics_/`,CI 可上传该目录作为 artifact 供审阅。 + +3. **Slack 告警(可选)** + - 在本地或 CI 中设置 `SLACK_RISK_WEBHOOK=https://hooks.slack.com/services/...`,即可运行 `./scripts/notify_risk_metrics.sh`;脚本会调用 `watch_risk_metrics.py` 并在发现 fail 时发送告警。 + - 生产服务器可将同一命令写入 cron(示例:`*/30 * * * * cd /path/to/QuantResearch && source .env && ./scripts/notify_risk_metrics.sh`);请参考 `.env.example` 填写 webhook。 + +4. **Prometheus 推送(Phase 4)** + ```bash + python scripts/export_metrics_prom.py | curl --data-binary @- http://pushgateway:9091/metrics/job/risk_sim + ``` + CI 已支持该脚本;如需手动推送,请先在 `.env` 中配置 `PUSHGATEWAY_URL`。 + +更多运维细节见 `docs/runbook_paper_risk.md` 与 `docs/runbook_ops.md`。 diff --git a/Q Research/requirements.txt b/Q Research/requirements.txt new file mode 100644 index 0000000..1646ab7 --- /dev/null +++ b/Q Research/requirements.txt @@ -0,0 +1,12 @@ +# requirements.txt + +oandapyV20>=0.7.2 +pandas>=2.0.0 +numpy>=1.23.0 +loguru>=0.7.0 +python-dotenv>=1.0.0 +plotly>=5.0.0 +ta>=0.10.0 +pyyaml>=6.0.0 +streamlit>=1.28.0 +xgboost==1.7.6 diff --git a/Q Research/scripts/__pycache__/backtest_example.cpython-312.pyc b/Q Research/scripts/__pycache__/backtest_example.cpython-312.pyc new file mode 100644 index 0000000..937dbf9 Binary files /dev/null and b/Q Research/scripts/__pycache__/backtest_example.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/backtest_strategy.cpython-312.pyc b/Q Research/scripts/__pycache__/backtest_strategy.cpython-312.pyc new file mode 100644 index 0000000..f23bdd8 Binary files /dev/null and b/Q Research/scripts/__pycache__/backtest_strategy.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/backtest_strategy.cpython-313.pyc b/Q Research/scripts/__pycache__/backtest_strategy.cpython-313.pyc new file mode 100644 index 0000000..6d3edde Binary files /dev/null and b/Q Research/scripts/__pycache__/backtest_strategy.cpython-313.pyc differ diff --git a/Q Research/scripts/__pycache__/compute_indicators.cpython-312.pyc b/Q Research/scripts/__pycache__/compute_indicators.cpython-312.pyc new file mode 100644 index 0000000..7eaec1b Binary files /dev/null and b/Q Research/scripts/__pycache__/compute_indicators.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/get_candles.cpython-312.pyc b/Q Research/scripts/__pycache__/get_candles.cpython-312.pyc new file mode 100644 index 0000000..16d3ccd Binary files /dev/null and b/Q Research/scripts/__pycache__/get_candles.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/paper_trade.cpython-312.pyc b/Q Research/scripts/__pycache__/paper_trade.cpython-312.pyc new file mode 100644 index 0000000..0b93540 Binary files /dev/null and b/Q Research/scripts/__pycache__/paper_trade.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/paper_trade.cpython-313.pyc b/Q Research/scripts/__pycache__/paper_trade.cpython-313.pyc new file mode 100644 index 0000000..6e4e347 Binary files /dev/null and b/Q Research/scripts/__pycache__/paper_trade.cpython-313.pyc differ diff --git a/Q Research/scripts/__pycache__/plot_candles.cpython-312.pyc b/Q Research/scripts/__pycache__/plot_candles.cpython-312.pyc new file mode 100644 index 0000000..08a147d Binary files /dev/null and b/Q Research/scripts/__pycache__/plot_candles.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/plot_indicators.cpython-312.pyc b/Q Research/scripts/__pycache__/plot_indicators.cpython-312.pyc new file mode 100644 index 0000000..44f3232 Binary files /dev/null and b/Q Research/scripts/__pycache__/plot_indicators.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/scenario_utils.cpython-312.pyc b/Q Research/scripts/__pycache__/scenario_utils.cpython-312.pyc new file mode 100644 index 0000000..2547ac9 Binary files /dev/null and b/Q Research/scripts/__pycache__/scenario_utils.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/scenario_utils.cpython-313.pyc b/Q Research/scripts/__pycache__/scenario_utils.cpython-313.pyc new file mode 100644 index 0000000..5a37e02 Binary files /dev/null and b/Q Research/scripts/__pycache__/scenario_utils.cpython-313.pyc differ diff --git a/Q Research/scripts/__pycache__/simulate_execution.cpython-312.pyc b/Q Research/scripts/__pycache__/simulate_execution.cpython-312.pyc new file mode 100644 index 0000000..eab4ce9 Binary files /dev/null and b/Q Research/scripts/__pycache__/simulate_execution.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/train_xgb_usdjpy.cpython-312.pyc b/Q Research/scripts/__pycache__/train_xgb_usdjpy.cpython-312.pyc new file mode 100644 index 0000000..cb8b887 Binary files /dev/null and b/Q Research/scripts/__pycache__/train_xgb_usdjpy.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/validate_dataset.cpython-312.pyc b/Q Research/scripts/__pycache__/validate_dataset.cpython-312.pyc new file mode 100644 index 0000000..53bbace Binary files /dev/null and b/Q Research/scripts/__pycache__/validate_dataset.cpython-312.pyc differ diff --git a/Q Research/scripts/__pycache__/validate_dataset.cpython-313.pyc b/Q Research/scripts/__pycache__/validate_dataset.cpython-313.pyc new file mode 100644 index 0000000..99d0c87 Binary files /dev/null and b/Q Research/scripts/__pycache__/validate_dataset.cpython-313.pyc differ diff --git a/Q Research/scripts/aggregate_data_quality.py b/Q Research/scripts/aggregate_data_quality.py new file mode 100644 index 0000000..c3b0dd5 --- /dev/null +++ b/Q Research/scripts/aggregate_data_quality.py @@ -0,0 +1,101 @@ +#!/usr/bin/env python3 +""" +Aggregate data-quality reports (results/data_quality/*.json) into a tabular summary. +""" + +from __future__ import annotations + +import argparse +import csv +import json +from datetime import datetime +from pathlib import Path +from typing import List, Dict + + +DEFAULT_DIR = Path("results/data_quality") +DEFAULT_OUTPUT = Path("metrics/data_quality_summary.csv") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Aggregate data-quality reports.") + parser.add_argument("--input", default=str(DEFAULT_DIR), help="Directory containing data_quality JSON reports.") + parser.add_argument("--output", default=str(DEFAULT_OUTPUT), help="CSV file to write summary (default metrics/data_quality_summary.csv).") + parser.add_argument("--print", action="store_true", help="Print summary to stdout.") + return parser.parse_args() + + +def load_reports(report_dir: Path) -> List[Dict]: + rows: List[Dict] = [] + if not report_dir.exists(): + return rows + for path in sorted(report_dir.glob("*.json")): + try: + data = json.load(path.open("r", encoding="utf-8")) + except Exception: + continue + manifest = data.get("manifest") or {} + dataset_path = data.get("dataset_path") or manifest.get("path") + rows.append({ + "generated_at": data.get("generated_at"), + "dataset_path": dataset_path, + "symbol": infer_symbol(path.name, dataset_path), + "severity": data.get("severity"), + "gap_ratio": data.get("gap_ratio"), + "duplicate_timestamps": data.get("duplicate_timestamps"), + "null_max": max((data.get("null_counts") or {}).values() or [0]), + "outlier_columns": ", ".join((data.get("numeric_outliers") or {}).keys()), + "hash": manifest.get("sha256"), + "report_file": str(path), + }) + return rows + + +def infer_symbol(filename: str, dataset_path: str | None) -> str: + if dataset_path: + stem = Path(dataset_path).stem + parts = stem.split("_") + if parts: + return parts[0] + if "_" in filename: + return filename.split("_")[1] + return "UNKNOWN" + + +def write_csv(rows: List[Dict], output: Path) -> None: + output.parent.mkdir(parents=True, exist_ok=True) + headers = ["generated_at", "dataset_path", "symbol", "severity", "gap_ratio", "duplicate_timestamps", "null_max", "outlier_columns", "hash", "report_file"] + with output.open("w", newline="", encoding="utf-8") as fh: + writer = csv.DictWriter(fh, fieldnames=headers) + writer.writeheader() + for row in rows: + writer.writerow(row) + + +def print_table(rows: List[Dict]) -> None: + if not rows: + print("No reports found.") + return + print(f"{'Generated':25} {'Symbol':8} {'Severity':7} {'Gap%':7} {'Dup':5} {'Hash':64}") + for row in rows: + gap = f"{row['gap_ratio']:.4f}" if isinstance(row["gap_ratio"], (int, float)) else "n/a" + print( + f"{row['generated_at'][:23] if row['generated_at'] else '':25} " + f"{row['symbol']:8} {row['severity']:7} {gap:7} " + f"{row['duplicate_timestamps']!s:5} {row['hash'] or ''}" + ) + + +def main() -> None: + args = parse_args() + report_dir = Path(args.input) + rows = load_reports(report_dir) + if args.print: + print_table(rows) + output_path = Path(args.output) + write_csv(rows, output_path) + print(f"Wrote summary to {output_path} ({len(rows)} rows).") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/analyze_cost_profiles.py b/Q Research/scripts/analyze_cost_profiles.py new file mode 100644 index 0000000..b64dfec --- /dev/null +++ b/Q Research/scripts/analyze_cost_profiles.py @@ -0,0 +1,123 @@ +from __future__ import annotations + +import argparse +import json +import os +from datetime import datetime +from pathlib import Path +from typing import Any, Dict, List + +import pandas as pd +import yaml +from loguru import logger + + +DEFAULT_SESSIONS: List[Dict[str, Any]] = [ + {"name": "asia_open", "weekdays": [0, 1, 2, 3, 4], "start_hour": 0, "end_hour": 7}, + {"name": "europe", "weekdays": [0, 1, 2, 3, 4], "start_hour": 7, "end_hour": 13}, + {"name": "us_session", "weekdays": [0, 1, 2, 3, 4], "start_hour": 13, "end_hour": 22}, +] + + +def _load_sessions(path: str | None) -> List[Dict[str, Any]]: + if not path: + return DEFAULT_SESSIONS + fp = Path(path) + if not fp.exists(): + raise FileNotFoundError(f"Session file not found: {path}") + if fp.suffix.lower() in {".yml", ".yaml"}: + data = yaml.safe_load(fp.read_text(encoding="utf-8")) or {} + else: + data = json.loads(fp.read_text(encoding="utf-8")) + if isinstance(data, dict): + sessions = data.get("sessions") or data.get("profiles") + else: + sessions = data + if not isinstance(sessions, list): + raise ValueError("session file must contain a list of session definitions") + return sessions + + +def _match_session(row: pd.Series, session: Dict[str, Any]) -> bool: + hour = row["hour"] + weekday = row["weekday"] + weekdays = session.get("weekdays") + if weekdays and weekday not in weekdays: + return False + start = session.get("start_hour") + end = session.get("end_hour") + if start is None and end is None: + return True + start = 0 if start is None else float(start) + end = 24 if end is None else float(end) + if start < end: + return start <= hour < end + return hour >= start or hour < end + + +def main(): + parser = argparse.ArgumentParser(description="Aggregate spread/slippage samples into cost profiles.") + parser.add_argument("--input", required=True, help="CSV with columns ts,spread_pips,slip_pips (plus optional fields).") + parser.add_argument("--symbol", required=True, help="Symbol name, e.g. USDJPY.") + parser.add_argument("--out", default="data/cost_profiles/profile.yaml", help="Output YAML path.") + parser.add_argument("--sessions", help="Optional YAML/JSON file describing session windows.") + parser.add_argument("--min-samples", type=int, default=20, help="Minimum rows required for a session to be emitted.") + args = parser.parse_args() + + df = pd.read_csv(args.input) + if "ts" not in df.columns: + raise ValueError("input CSV must contain 'ts' column (timestamp).") + if "spread_pips" not in df.columns or "slip_pips" not in df.columns: + raise ValueError("input CSV must contain 'spread_pips' and 'slip_pips'.") + + df["ts"] = pd.to_datetime(df["ts"], utc=True, errors="coerce") + df = df.dropna(subset=["ts"]) + df["hour"] = df["ts"].dt.hour + df["ts"].dt.minute / 60.0 + df["weekday"] = df["ts"].dt.weekday + + sessions = _load_sessions(args.sessions) + profiles: List[Dict[str, Any]] = [] + for session in sessions: + mask = df.apply(lambda row: _match_session(row, session), axis=1) + subset = df.loc[mask] + if len(subset) < args.min_samples: + logger.warning(f"Session {session.get('name')} skipped (samples={len(subset)} < {args.min_samples}).") + continue + profile = { + "name": session.get("name", f"session_{len(profiles)}"), + "weekdays": session.get("weekdays"), + "start_hour": session.get("start_hour"), + "end_hour": session.get("end_hour"), + "spread": round(subset["spread_pips"].mean(), 4), + "slip": round(subset["slip_pips"].mean(), 4), + "comm": session.get("comm"), + "samples": int(len(subset)), + "spread_p95": round(subset["spread_pips"].quantile(0.95), 4), + "slip_p95": round(subset["slip_pips"].quantile(0.95), 4), + } + if session.get("priority") is not None: + profile["priority"] = session["priority"] + profiles.append(profile) + + if not profiles: + raise RuntimeError("No sessions met the minimum sample requirement; nothing to write.") + + default_profile = min(profiles, key=lambda p: p.get("priority", float("inf"))) + default_profile["default"] = True + + payload = { + "symbol": args.symbol.upper(), + "generated_at": datetime.utcnow().isoformat() + "Z", + "source": os.path.abspath(args.input), + "profiles": profiles, + } + + out_path = Path(args.out) + out_path.parent.mkdir(parents=True, exist_ok=True) + with out_path.open("w", encoding="utf-8") as fh: + yaml.safe_dump(payload, fh, allow_unicode=True, sort_keys=False) + logger.info(f"Cost profile saved to {out_path}") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/backfill_risk_metrics.py b/Q Research/scripts/backfill_risk_metrics.py new file mode 100644 index 0000000..382e520 --- /dev/null +++ b/Q Research/scripts/backfill_risk_metrics.py @@ -0,0 +1,146 @@ +#!/usr/bin/env python3 +""" +Backfill results/risk/metrics.csv using historical execution runs. + +For each run under results/execution//, the script inspects +sim_results.json, derives reject/kill counts, infers status (pass/fail), +and appends any missing records to results/risk/metrics.csv. +""" + +from __future__ import annotations + +import argparse +import csv +import json +from datetime import datetime, timezone +from pathlib import Path +from typing import Dict, List, Optional + +ROOT = Path(__file__).resolve().parents[1] +EXECUTION_DIR = ROOT / "results" / "execution" +METRICS_PATH = ROOT / "results" / "risk" / "metrics.csv" + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Backfill risk metrics from historical execution runs.") + parser.add_argument( + "--runs", + type=str, + help="Comma-separated run IDs to backfill. Defaults to all directories in results/execution/.", + ) + parser.add_argument( + "--force", + action="store_true", + help="Re-write entries even if run_id already exists in metrics.csv (otherwise skipped).", + ) + return parser.parse_args() + + +def list_runs(filter_ids: Optional[List[str]]) -> List[str]: + if filter_ids: + return filter_ids + if not EXECUTION_DIR.exists(): + return [] + return sorted([p.name for p in EXECUTION_DIR.iterdir() if p.is_dir()]) + + +def load_sim_results(run_id: str) -> Optional[Dict]: + sim_path = EXECUTION_DIR / run_id / "sim_results.json" + if not sim_path.exists(): + return None + return json.loads(sim_path.read_text(encoding="utf-8")) + + +def infer_timestamp(run_id: str) -> str: + summary_path = ROOT / "results" / run_id / "summary.json" + if summary_path.exists(): + try: + data = json.loads(summary_path.read_text(encoding="utf-8")) + if "created_at" in data: + return data["created_at"] + except json.JSONDecodeError: + pass + sim_path = EXECUTION_DIR / run_id / "sim_results.json" + if sim_path.exists(): + return datetime.fromtimestamp(sim_path.stat().st_mtime, tz=timezone.utc).isoformat() + return datetime.now(timezone.utc).isoformat() + + +def summarize_run(run_id: str) -> Optional[Dict[str, str]]: + payload = load_sim_results(run_id) + if payload is None: + return None + rejects = len(payload.get("rejects", [])) + kills = len(payload.get("kill_switch_events", [])) + status = "pass" if rejects == 0 and kills == 0 else "fail" + return { + "timestamp": infer_timestamp(run_id), + "run_id": run_id, + "rejects": str(rejects), + "kills": str(kills), + "status": status, + } + + +def load_existing() -> Dict[str, Dict[str, str]]: + if not METRICS_PATH.exists(): + return {} + with METRICS_PATH.open("r", encoding="utf-8") as fh: + reader = csv.DictReader(fh) + return {row["run_id"]: row for row in reader if row.get("run_id")} + + +def append_rows(rows: List[Dict[str, str]], replace: bool = False) -> None: + METRICS_PATH.parent.mkdir(parents=True, exist_ok=True) + existing = load_existing() + if replace: + for row in rows: + existing[row["run_id"]] = row + with METRICS_PATH.open("w", encoding="utf-8", newline="") as fh: + writer = csv.DictWriter(fh, fieldnames=["timestamp", "run_id", "rejects", "kills", "status"]) + writer.writeheader() + for row in sorted(existing.values(), key=lambda r: r["timestamp"]): + writer.writerow(row) + else: + new_file = not METRICS_PATH.exists() + with METRICS_PATH.open("a", encoding="utf-8", newline="") as fh: + writer = csv.DictWriter(fh, fieldnames=["timestamp", "run_id", "rejects", "kills", "status"]) + if new_file: + writer.writeheader() + for row in rows: + writer.writerow(row) + + +def main() -> None: + args = parse_args() + runs = list_runs(args.runs.split(",") if args.runs else None) + if not runs: + print("No runs to process.") + return + existing = load_existing() + new_rows: List[Dict[str, str]] = [] + replacements: List[Dict[str, str]] = [] + for run_id in runs: + record = summarize_run(run_id) + if not record: + print(f"[skip] sim_results.json missing for {run_id}") + continue + if run_id in existing and not args.force: + print(f"[skip] {run_id} already in metrics.csv") + continue + if args.force and run_id in existing: + replacements.append(record) + else: + new_rows.append(record) + if replacements: + append_rows(replacements, replace=True) + print(f"Replaced {len(replacements)} entries (force mode).") + if new_rows: + append_rows(new_rows, replace=False) + print(f"Appended {len(new_rows)} new entries.") + if not replacements and not new_rows: + print("No changes made.") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/backtest_strategy.py b/Q Research/scripts/backtest_strategy.py new file mode 100644 index 0000000..0092a10 --- /dev/null +++ b/Q Research/scripts/backtest_strategy.py @@ -0,0 +1,828 @@ +from __future__ import annotations + +import os +import sys +from pathlib import Path +from datetime import datetime, timezone + +sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +import json +import pandas as pd +import yaml +import numpy as np +from queue import Empty, Queue +from typing import Any, List, Optional +from loguru import logger + +from core.backtest.strategy_engine import ( + FXRateProvider, + StrategyEngine, + StrategySpec, + parse_strategy_specs, + _coerce_fx_rates, + _merge_fx_rates, +) +from data.csv_feed import CSVFeed # 你的CSVFeed +from scripts.validate_dataset import ( + DEFAULT_MANIFEST, + compute_report, + load_manifest_entry, +) + +# ===== FX 元数据与换算工具 ===== +BASE_DIR = os.path.dirname(os.path.dirname(__file__)) +DATA_DIR = os.path.join(BASE_DIR, "data") +RAW_DATA_DIR = os.path.join(DATA_DIR, "raw") +DERIVED_DATA_DIR = os.path.join(DATA_DIR, "derived") +OUTPUT_DIR = os.path.join(DATA_DIR, "outputs") +EQUITY_DIR = os.path.join(OUTPUT_DIR, "equity") +TRADES_DIR = os.path.join(OUTPUT_DIR, "trades") +STATS_DIR = os.path.join(OUTPUT_DIR, "stats") +DATA_REPORT_DIR = os.path.join(STATS_DIR, "data_reports") +RESULTS_DIR = os.path.join(BASE_DIR, "results") +GRID_DIR = os.path.join(DATA_DIR, "grid") +PARAMS_DIR = os.path.join(DATA_DIR, "params") + +for _dir in [RAW_DATA_DIR, DERIVED_DATA_DIR, EQUITY_DIR, TRADES_DIR, STATS_DIR, DATA_REPORT_DIR, RESULTS_DIR, GRID_DIR, PARAMS_DIR]: + os.makedirs(_dir, exist_ok=True) + + +def _load_manifest_entry(csv_path: Path, manifest_path: Optional[str]) -> Optional[dict]: + if not manifest_path: + manifest_path = os.path.join(DATA_DIR, "_manifest.json") + manifest_file = Path(manifest_path) + if not manifest_file.exists(): + return None + try: + entry = load_manifest_entry(manifest_file, csv_path) + return entry + except Exception as exc: + logger.warning(f"Failed to read manifest entry for {csv_path}: {exc}") + return None + + +def _validate_input_dataset(csv_path: str, manifest_path: Optional[str] = DEFAULT_MANIFEST) -> dict: + dataset_path = Path(csv_path).expanduser().resolve() + manifest_entry = _load_manifest_entry(dataset_path, manifest_path) + report = compute_report(dataset_path, manifest_entry, z_threshold=5.0) + severity = report.get("severity") + gap_ratio = report.get("gap_ratio") + gap_ratio_str = f"{gap_ratio:.4f}" if isinstance(gap_ratio, (int, float)) else "n/a" + logger.info( + f"Data validation severity={severity} " + f"duplicates={report.get('duplicate_timestamps')} gap_ratio={gap_ratio_str}" + ) + if severity == "error": + raise RuntimeError( + f"Dataset validation failed for {dataset_path}. Messages: {report.get('messages')}" + ) + return report + + +def _relpath_or_abs(path: Optional[str]) -> Optional[str]: + if not path: + return None + try: + return str(Path(path).resolve().relative_to(BASE_DIR)) + except Exception: + return str(path) + + +def _load_structured_data(path: Optional[str]): + if not path: + return None + file_path = Path(path).expanduser() + if not file_path.exists(): + raise FileNotFoundError(f"Config file not found: {file_path}") + with file_path.open("r", encoding="utf-8") as fh: + if file_path.suffix.lower() in (".yaml", ".yml"): + return yaml.safe_load(fh) + return json.load(fh) + + +def _write_data_report(report: Optional[dict], symbol: str, fast_win: int, slow_win: int, suffix: str) -> Optional[str]: + if not report: + return None + report_path = Path(DATA_REPORT_DIR) / f"data_{symbol}_H1_{fast_win}x{slow_win}_{suffix}.json" + with report_path.open("w", encoding="utf-8") as fh: + json.dump(report, fh, indent=2, ensure_ascii=False) + try: + return str(report_path.relative_to(BASE_DIR)) + except ValueError: + return str(report_path) + + +def _prepare_run_dir(enabled: bool, results_dir: Optional[str]) -> tuple[Optional[str], Optional[Path]]: + if not enabled or not results_dir: + return None, None + run_id = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S") + run_path = Path(results_dir) / run_id + run_path.mkdir(parents=True, exist_ok=True) + return run_id, run_path + + +def _write_run_summary(run_path: Path, summary: dict) -> str: + summary_path = run_path / "summary.json" + with summary_path.open("w", encoding="utf-8") as fh: + json.dump(summary, fh, indent=2, ensure_ascii=False) + metrics_path = run_path / "metrics.json" + with metrics_path.open("w", encoding="utf-8") as fh: + json.dump(summary.get("metrics", {}), fh, indent=2, ensure_ascii=False) + return str(summary_path) + + +def run_once( + symbol: str = "EURUSD", + csv_path: str = os.path.join(RAW_DATA_DIR, "EURUSD_H1.csv"), + initial_cash: float = 100000.0, + qty: int = 10_000, + account_ccy: str = "USD", + fx_rates: FXRateProvider = None, + fast_win: int = 20, # 改为更敏感的短期均线 + slow_win: int = 100, # 改为中期均线 + spread_pips: float = 2.0, + commission_per_million: float = 0.25, + slippage_pips: float = 0.3, + stop_loss_pips: float = 50, + take_profit_pips: float | None = None, + atr_sl: float | None = 1.5, # 默认使用1.5倍ATR止损 + atr_tp: float | None = 3.0, # 默认使用3倍ATR止盈 + atr_window: int = 14, # ATR窗口保持14天 + # RSI & trailing defaults + rsi_period: int = 14, + rsi_long_thresh: Optional[float] = None, + rsi_short_thresh: Optional[float] = None, + enable_trailing: bool = False, + trailing_enable_atr_mult: float = 1.0, + trailing_atr_mult: float = 0.5, + htf_factor: int = 4, + htf_ema_window: Optional[int] = None, + htf_rsi_period: Optional[int] = None, + regime_ema_window: int = 200, + regime_slope_min: Optional[float] = None, + regime_atr_min: Optional[float] = None, + regime_atr_percentile_min: Optional[float] = None, + regime_atr_percentile_window: int = 500, + regime_trend_min_bars: int = 0, + long_only_above_slow: bool = False, + slope_lookback: int = 0, + cooldown: int = 0, + allow_short: bool = True, + short_only_below_slow: bool = False, + strategies: Optional[List[StrategySpec]] = None, + cost_profiles: Optional[Any] = None, + slippage_model: Optional[Any] = None, + strategy_mode: str = "first_hit", + strategy_vote_threshold: float = 0.0, + stress_cost_spread_mult: float = 1.0, + stress_cost_comm_mult: float = 1.0, + stress_slippage_mult: float = 1.0, + stress_price_vol_mult: float = 1.0, + stress_skip_trade_pct: float = 0.0, + risk_per_trade_pct: Optional[float] = None, + max_drawdown_pct: Optional[float] = None, + max_position_units: Optional[float] = None, + skip_outlier_entries: bool = False, + validate_data: bool = True, + manifest_path: Optional[str] = DEFAULT_MANIFEST, + results_dir: Optional[str] = RESULTS_DIR, + write_summary: bool = True, +): + os.makedirs(EQUITY_DIR, exist_ok=True) + os.makedirs(TRADES_DIR, exist_ok=True) + os.makedirs(STATS_DIR, exist_ok=True) + + run_id, run_dir = _prepare_run_dir(write_summary, results_dir) + + data_report = None + if validate_data and csv_path: + try: + data_report = _validate_input_dataset(csv_path, manifest_path) + except RuntimeError: + raise + except Exception as exc: + raise RuntimeError(f"Data validation error: {exc}") from exc + + engine = StrategyEngine( + symbol=symbol, + fast_win=fast_win, + slow_win=slow_win, + spread_pips=spread_pips, + commission_per_million=commission_per_million, + slippage_pips=slippage_pips, + stop_loss_pips=stop_loss_pips, + take_profit_pips=take_profit_pips, + atr_sl=atr_sl, + atr_tp=atr_tp, + atr_window=atr_window, + regime_ema_window=regime_ema_window, + regime_slope_min=regime_slope_min, + regime_atr_min=regime_atr_min, + regime_atr_percentile_min=regime_atr_percentile_min, + regime_atr_percentile_window=regime_atr_percentile_window, + regime_trend_min_bars=regime_trend_min_bars, + rsi_period=rsi_period, + rsi_long_thresh=rsi_long_thresh, + rsi_short_thresh=rsi_short_thresh, + enable_trailing=enable_trailing, + trailing_enable_atr_mult=trailing_enable_atr_mult, + trailing_atr_mult=trailing_atr_mult, + htf_factor=htf_factor, + htf_ema_window=htf_ema_window, + htf_rsi_period=htf_rsi_period, + long_only_above_slow=long_only_above_slow, + slope_lookback=slope_lookback, + cooldown=cooldown, + qty=qty, + account_ccy=account_ccy, + fx_rates=fx_rates, + strategy_specs=strategies, + cost_profiles=cost_profiles, + slippage_model=slippage_model, + strategy_combine_mode=strategy_mode, + strategy_vote_threshold=strategy_vote_threshold, + stress_cost_spread_mult=stress_cost_spread_mult, + stress_cost_comm_mult=stress_cost_comm_mult, + stress_slippage_mult=stress_slippage_mult, + stress_price_vol_mult=stress_price_vol_mult, + stress_skip_trade_pct=stress_skip_trade_pct, + skip_outlier_bars=skip_outlier_entries, + allow_short=allow_short, + short_only_below_slow=short_only_below_slow, + risk_per_trade_pct=risk_per_trade_pct, + max_drawdown_pct=max_drawdown_pct, + max_position_units=max_position_units, + output_dirs={ + "equity": EQUITY_DIR, + "trades": TRADES_DIR, + "stats": STATS_DIR, + }, + ) + engine.set_initial_cash(initial_cash) + + q = Queue() + data = CSVFeed(q, path=str(csv_path), symbol=symbol) + logger.info(f"Using CSV: {csv_path} for {symbol}") + data.start() + + while True: + try: + ev = q.get(timeout=0.05) + except Empty: + if hasattr(data, "pump"): + data.pump(n=50) + if getattr(data, "finished", False): + break + continue + if ev.get("type") != "bar": + continue + engine.handle_bar(ev) + + engine.finalize() + + suffix = engine.compute_suffix() + output_files = engine.export_outputs(fast_win, slow_win, suffix) + data_report_path = _write_data_report(data_report, symbol, fast_win, slow_win, suffix) + result = engine.summary(fast_win, slow_win, suffix) + + final_equity = result["final_equity"] if result["final_equity"] is not None else engine.cash + ret_pct = (final_equity / initial_cash - 1.0) * 100.0 + logger.info(f"Bars processed: {engine.bar_count}, Trades executed: {engine.trade_count}") + logger.info(f"策略最终权益: {final_equity:.2f},累计收益: {ret_pct:.4f}%") + if all(result.get(k) is not None for k in ("sharpe", "ann_return", "ann_vol", "max_drawdown")): + logger.info( + f"Sharpe={result['sharpe']:.3f} AnnRet={result['ann_return']*100:.2f}% " + f"AnnVol={result['ann_vol']*100:.2f}% MaxDD={result['max_drawdown']*100:.2f}%" + ) + + data_summary = None + if data_report: + manifest_info = (data_report.get("manifest") or {}) + data_summary = { + "severity": data_report.get("severity"), + "gap_ratio": data_report.get("gap_ratio"), + "duplicate_timestamps": data_report.get("duplicate_timestamps"), + "hash": manifest_info.get("sha256"), + "path": manifest_info.get("path"), + } + logger.info( + "Data signature: severity={} hash={}", + data_summary["severity"], + data_summary["hash"], + ) + + if data_report_path: + result["data_report"] = data_report_path + result["data_validation"] = { + "severity": data_summary["severity"] if data_summary else None, + "messages": data_report.get("messages") if data_report else None, + } + + if run_dir: + param_snapshot = { + "fast_win": fast_win, + "slow_win": slow_win, + "spread_pips": spread_pips, + "commission_per_million": commission_per_million, + "slippage_pips": slippage_pips, + "stop_loss_pips": stop_loss_pips, + "take_profit_pips": take_profit_pips, + "atr_sl": atr_sl, + "atr_tp": atr_tp, + "atr_window": atr_window, + "rsi_period": rsi_period, + "regime_ema_window": regime_ema_window, + "skip_outlier_entries": skip_outlier_entries, + "strategy_mode": strategy_mode, + "strategy_vote_threshold": strategy_vote_threshold, + "stress_cost_spread_mult": stress_cost_spread_mult, + "stress_cost_comm_mult": stress_cost_comm_mult, + "stress_slippage_mult": stress_slippage_mult, + "stress_price_vol_mult": stress_price_vol_mult, + "stress_skip_trade_pct": stress_skip_trade_pct, + } + summary = { + "run_id": run_id, + "timestamp": datetime.now(timezone.utc).isoformat(), + "symbol": symbol, + "csv_path": os.path.relpath(csv_path, BASE_DIR) if csv_path else None, + "parameters": param_snapshot, + "metrics": result, + "data_report": data_summary, + "artifacts": { + "equity": _relpath_or_abs(output_files.get("equity")), + "trades": _relpath_or_abs(output_files.get("trades")), + "trade_stats": _relpath_or_abs(output_files.get("trade_stats")), + }, + } + summary_path = _write_run_summary(run_dir, summary) + result["run_id"] = run_id + result["summary_path"] = summary_path + + return result + +def main(**kwargs): + """ + Backwards-compatible wrapper for legacy callers that imported + scripts.backtest_strategy.main. It simply proxies to run_once(). + """ + return run_once(**kwargs) + +def grid_search(symbol="EURUSD", + csv_path=None, + qty=10_000, + initial_cash=100000.0, + account_ccy="USD", + fx_rates: FXRateProvider = None, + # 单值默认;若传入 *_list 则以列表为准 + spread=1.0, + slip=0.2, + comm=2.0, + atr_window=14, + skip_outlier_entries: bool = False, + # 维度开关;传 None 使用默认网格 + fast_list=None, + slow_list=None, + atr_sl_list=None, + atr_tp_list=None, + long_only_list=None, + cooldown_list=None, + slope_list=None, + spread_list=None, + slip_list=None, + comm_list=None): + """ + 多维参数网格搜索。 + - 若 *_list 为 None,则采用合理的默认网格;否则使用传入列表。 + - 结果会输出: + data/grid/grid_{symbol}_H1_ATR.csv + data/grid/grid_top10_by_sharpe_{symbol}.csv + data/params/best_params_grid_{symbol}.json + """ + import pandas as pd + import json + + # --- 默认网格(可被参数列表覆盖) --- + fast_list = fast_list or [20, 30, 50] + slow_list = slow_list or [100, 150, 200] + atr_sl_list = atr_sl_list or [1.5, 2.0] # 止损倍数 + atr_tp_list = atr_tp_list or [None, 2.0, 3.0] # 含不设止盈 + long_only_list = long_only_list or [False, True] + cooldown_list = cooldown_list or [0, 6, 12, 24] + slope_list = slope_list or [0, 3] + spread_list = spread_list or [spread] + slip_list = slip_list or [slip] + comm_list = comm_list or [comm] + + rows = [] + total = 0 + for f in fast_list: + for s in slow_list: + if f >= s: + continue + for k in atr_sl_list: + for m in atr_tp_list: + for lo in long_only_list: + for cd in cooldown_list: + for slp in slope_list: + for sp in spread_list: + for sp_slip in slip_list: + for cm in comm_list: + total += 1 + logger.info( + f"[GRID] sym={symbol} fast={f} slow={s} " + f"SL=ATR×{k} TP={'None' if m is None else 'ATR×'+str(m)} " + f"ABOVE={lo} CD={cd} SLOPE={slp} " + f"spread={sp} slip={sp_slip} comm={cm}" + ) + res = run_once( + symbol=symbol, + csv_path=csv_path, + fast_win=int(f), slow_win=int(s), + spread_pips=float(sp), + commission_per_million=float(cm), + slippage_pips=float(sp_slip), + # 关闭固定 pips,启用 ATR + stop_loss_pips=None, + take_profit_pips=None, + atr_sl=float(k) if k is not None else None, + atr_tp=float(m) if m is not None else None, + atr_window=int(atr_window), + qty=int(qty), + initial_cash=float(initial_cash), + account_ccy=str(account_ccy), + fx_rates=fx_rates, + long_only_above_slow=bool(lo), + cooldown=int(cd), + slope_lookback=int(slp), + skip_outlier_entries=skip_outlier_entries, + write_summary=False, + ) + # 把当前维度也写入结果,便于回看 + res.update({ + "symbol": symbol, + "spread": float(sp), + "slip": float(sp_slip), + "comm": float(cm), + "long_only_above_slow": bool(lo), + "cooldown": int(cd), + "slope_lookback": int(slp), + }) + rows.append(res) + + df = pd.DataFrame(rows) + os.makedirs(GRID_DIR, exist_ok=True) + os.makedirs(PARAMS_DIR, exist_ok=True) + out = os.path.join(GRID_DIR, f"grid_{symbol}_H1_ATR.csv") + df.to_csv(out, index=False) + logger.info(f"[GRID] 扫描完成(组合数={total}),已保存: {out}") + + try: + if df.empty: + logger.warning("[GRID] 无结果,跳过排名/保存。") + return + # 确保 sharpe 可排序 + df["sharpe"] = pd.to_numeric(df["sharpe"], errors="coerce") + df_sorted = df.sort_values("sharpe", ascending=False, na_position="last") + + logger.info("\n[GRID] Top 10 by Sharpe:\n" + df_sorted.head(10).to_string(index=False)) + + # 保存最优参数 + best = df_sorted.iloc[0].to_dict() + best_path = os.path.join(PARAMS_DIR, f"best_params_grid_{symbol}.json") + with open(best_path, "w", encoding="utf-8") as f: + json.dump(best, f, ensure_ascii=False, indent=2) + logger.info(f"[GRID] 最优参数已保存: {best_path}") + + # 保存 Top-10 + top10_path = os.path.join(GRID_DIR, f"grid_top10_by_sharpe_{symbol}.csv") + df_sorted.head(10).to_csv(top10_path, index=False) + logger.info(f"[GRID] Top 10 已保存: {top10_path}") + except Exception as e: + logger.warning(f"[GRID] 排序/保存失败: {e}") + +if __name__ == "__main__": + import argparse + ap = argparse.ArgumentParser() + ap.add_argument("--grid", action="store_true", help="启用参数网格扫描(含 ATR)") + ap.add_argument("--symbol", type=str, default="EURUSD", help="交易品种") + ap.add_argument("--csv", type=str, default=os.path.join(RAW_DATA_DIR, "EURUSD_H1.csv"), help="CSV 路径") + ap.add_argument("--fast", type=int, default=50, help="快速均线窗口") + ap.add_argument("--slow", type=int, default=200, help="慢速均线窗口(应 > fast)") + ap.add_argument("--qty", type=int, default=10_000, help="下单数量(名义)") + ap.add_argument("--cash", type=float, default=100000.0, help="初始资金") + ap.add_argument("--account-ccy", type=str, default="USD", help="账户结算货币(默认 USD)") + ap.add_argument("--spread", type=float, default=1.0, help="点差(pips)") + ap.add_argument("--slip", type=float, default=0.2, help="滑点(pips)") + ap.add_argument("--comm", type=float, default=2.0, help="佣金($ per $1,000,000 名义)") + ap.add_argument("--sl", type=float, default=50.0, help="止损(pips)") + ap.add_argument("--tp", type=float, default=None, help="止盈(pips,可空)") + ap.add_argument("--atr-sl", type=float, default=None, help="ATR 止损倍数(k_SL),例如 2.0 表示 2×ATR") + ap.add_argument("--atr-tp", type=float, default=None, help="ATR 止盈倍数(m_TP),例如 3.0 表示 3×ATR;缺省表示不用 ATR 止盈") + ap.add_argument("--atr-window", type=int, default=14, help="ATR 窗口(默认 14)") + ap.add_argument("--regime-ema-window", dest="regime_ema_window", type=int, default=200, help="Regime 过滤使用的 EMA 窗口长度") + ap.add_argument("--regime-slope-min", dest="regime_slope_min", type=float, default=None, help="EMA 斜率阈值(价格单位)判定趋势 regime") + ap.add_argument("--regime-atr-min", dest="regime_atr_min", type=float, default=None, help="ATR 下限,用于判定趋势 regime") + ap.add_argument("--regime-atr-percentile-min", dest="regime_atr_percentile_min", type=float, default=None, help="ATR 百分位下限(0-1),用来过滤低波动段") + ap.add_argument("--regime-atr-percentile-window", dest="regime_atr_percentile_window", type=int, default=500, help="ATR 百分位计算窗口长度(条数)") + ap.add_argument("--regime-trend-min-bars", dest="regime_trend_min_bars", type=int, default=0, help="趋势 regime 需要至少持续多少根 K 才允许入场") + ap.add_argument("--htf-factor", dest="htf_factor", type=int, default=4, help="高时间框聚合倍数(例如 4 表示 4 根低频合成一根高频)") + ap.add_argument("--htf-ema-window", dest="htf_ema_window", type=int, default=None, help="高时间框 EMA 窗口") + ap.add_argument("--htf-rsi-period", dest="htf_rsi_period", type=int, default=None, help="高时间框 RSI 周期") + ap.add_argument("--rsi-period", type=int, default=14, help="RSI 窗口(默认 14)") + ap.add_argument("--rsi-long-thresh", type=float, default=None, help="做多入场最低 RSI(例如 55)") + ap.add_argument("--rsi-short-thresh", type=float, default=None, help="做空入场最高 RSI(例如 45)") + ap.add_argument("--enable-trailing", action="store_true", help="启用基于 ATR 的 trailing stop") + ap.add_argument("--trailing-enable-atr-mult", type=float, default=1.0, help="盈利达到多少倍 entry_atr 时启用 trailing(默认1.0)") + ap.add_argument("--trailing-atr-mult", type=float, default=0.5, help="trailing 步长,按 curr_atr 的倍数移动止损(默认0.5)") + ap.add_argument("--long-only-above-slow", action="store_true", help="仅当 close > SMA_slow 时允许做多") + ap.add_argument("--slope-lookback", type=int, default=0, help="fast SMA 斜率确认(>0 开启, 单位=bar)") + ap.add_argument("--cooldown", type=int, default=0, help="平仓后冷却 N 根bar 才允许再次进场") + ap.add_argument("--config", type=str, default=None, help="YAML 配置路径(命令行显式参数将覆盖配置)") + ap.add_argument( + "--fx-rate", + action="append", + default=None, + help="额外换汇报价(可重复),格式示例:GBPUSD=1.27 或 EUR/JPY=161.3", + ) + ap.add_argument("--no-short", action="store_true", help="禁用做空信号") + ap.add_argument("--short-only-below-slow", action="store_true", help="仅当 close < SMA_slow 时允许做空") + ap.add_argument("--skip-outlier-entries", action="store_true", help="标记为 outlier 的 bar 上禁止开新仓位") + ap.add_argument("--strategy-mode", choices=["first_hit", "weighted"], default="first_hit", help="多策略组合模式(first_hit 或 weighted)") + ap.add_argument("--strategy-vote-threshold", type=float, default=0.0, help="weighted 模式下投票阈值(默认 0)") + ap.add_argument("--cost-profile-file", type=str, help="JSON/YAML 文件路径,定义成本/点差 profile") + ap.add_argument("--slippage-model-file", type=str, help="JSON/YAML 文件路径,定义滑点模型") + ap.add_argument("--stress-cost-spread-mult", type=float, default=1.0, help="压力测试:点差乘子(默认1)") + ap.add_argument("--stress-cost-comm-mult", type=float, default=1.0, help="压力测试:佣金乘子(默认1)") + ap.add_argument("--stress-slippage-mult", type=float, default=1.0, help="压力测试:滑点乘子(默认1)") + ap.add_argument("--stress-price-vol-mult", type=float, default=1.0, help="压力测试:高低点范围乘子(默认1)") + ap.add_argument("--stress-skip-trade-pct", type=float, default=0.0, help="压力测试:随机跳过交易的概率(0-1)") + ap.add_argument("--risk-percent", type=float, default=None, help="每笔风险占当前权益比例(如 0.01 表示 1%)") + ap.add_argument("--max-drawdown", type=float, default=None, help="最大允许回撤(小数,如 0.2 表示 20%),超出后停止开仓") + ap.add_argument("--max-units", type=float, default=None, help="仓位上限(基准货币单位)") + ap.add_argument( + "--use-best", + action="store_true", + help="从 data/params/best_params_grid_{symbol}.json 读取最优参数并运行(命令行显式参数仍可覆盖)" + ) + args = ap.parse_args() + cli_fx_rates = _coerce_fx_rates(args.fx_rate) + + if args.grid: + grid_search( + symbol=args.symbol, + csv_path=args.csv, + qty=args.qty, + initial_cash=args.cash, + account_ccy=args.account_ccy, + fx_rates=cli_fx_rates, + spread=args.spread, + slip=args.slip, + comm=args.comm, + atr_window=args.atr_window, + skip_outlier_entries=args.skip_outlier_entries, + ) + else: + # ---- 加载 YAML 配置并与命令行合并(命令行显式参数优先) ---- + cfg = {} + if args.config: + with open(args.config, "r", encoding="utf-8") as f: + raw = yaml.safe_load(f) or {} + # 允许大小写/短名对齐 argparse 名称 + key_map = { + "symbol": "symbol", + "csv": "csv_path", + "cash": "cash", + "qty": "qty", + "account_ccy": "account_ccy", + "fast": "fast", + "slow": "slow", + "spread": "spread", + "slip": "slip", + "comm": "comm", + "sl": "sl", + "tp": "tp", + "atr_sl": "atr_sl", + "atr_tp": "atr_tp", + "atr_window": "atr_window", + "regime_ema_window": "regime_ema_window", + "regime_slope_min": "regime_slope_min", + "regime_atr_min": "regime_atr_min", + "regime_atr_percentile_min": "regime_atr_percentile_min", + "regime_atr_percentile_window": "regime_atr_percentile_window", + "regime_trend_min_bars": "regime_trend_min_bars", + "htf_factor": "htf_factor", + "htf_ema_window": "htf_ema_window", + "htf_rsi_period": "htf_rsi_period", + "rsi_period": "rsi_period", + "rsi_long_thresh": "rsi_long_thresh", + "rsi_short_thresh": "rsi_short_thresh", + "enable_trailing": "enable_trailing", + "trailing_enable_atr_mult": "trailing_enable_atr_mult", + "trailing_atr_mult": "trailing_atr_mult", + "long_only_above_slow": "long_only_above_slow", + "slope_lookback": "slope_lookback", + "cooldown": "cooldown", + "fx_rates": "fx_rates", + "allow_short": "allow_short", + "short_only_below_slow": "short_only_below_slow", + "risk_per_trade_pct": "risk_per_trade_pct", + "max_drawdown_pct": "max_drawdown_pct", + "max_position_units": "max_position_units", + "skip_outlier_entries": "skip_outlier_entries", + "strategies": "strategies", + "cost_profiles": "cost_profiles", + "slippage_model": "slippage_model", + } + # 规范化键名 + norm = {} + for k, v in raw.items(): + kk = k.strip() + if kk in key_map: + norm[key_map[kk]] = v + else: + norm[kk] = v + cfg = norm + else: + cfg = {} + cfg_fx_rates = _coerce_fx_rates(cfg.get("fx_rates")) if cfg else None + cfg_strategies = parse_strategy_specs(cfg.get("strategies")) if cfg else None + cfg_cost_profiles = cfg.get("cost_profiles") if cfg else None + cfg_slippage_model = cfg.get("slippage_model") if cfg else None + if args.cost_profile_file: + cfg_cost_profiles = _load_structured_data(args.cost_profile_file) + if args.slippage_model_file: + cfg_slippage_model = _load_structured_data(args.slippage_model_file) + # [PATCH B START] 载入 best_params_grid.json(若 --use-best),并做类型规范化 + best_cfg = {} + if args.use_best: + try: + import json, math + os.makedirs(PARAMS_DIR, exist_ok=True) + best_path = os.path.join(PARAMS_DIR, f"best_params_grid_{args.symbol}.json") + with open(best_path, "r", encoding="utf-8") as f: + best = json.load(f) or {} + + def _is_nan(x): + return isinstance(x, float) and math.isnan(x) + + def _to_int_or_none(x): + if x is None or _is_nan(x): + return None + if isinstance(x, (int, np.integer)): + return int(x) + if isinstance(x, (float, np.floating)): + return int(round(float(x))) + # 其他类型尝试转 + try: + return int(float(x)) + except Exception: + return None + + def _to_float_or_none(x): + if x is None or _is_nan(x): + return None + if isinstance(x, (int, float, np.integer, np.floating)): + return float(x) + try: + v = float(x) + return v if not math.isnan(v) else None + except Exception: + return None + + # 将网格结果列名映射为参数名,并做规范化 + best_cfg = { + "fast": _to_int_or_none(best.get("fast")), + "slow": _to_int_or_none(best.get("slow")), + "atr_sl": _to_float_or_none(best.get("atr_sl")), + "atr_tp": _to_float_or_none(best.get("atr_tp")), # NaN -> None + "atr_window": _to_int_or_none(best.get("atr_window")), + } + # 去掉 None 的键,避免覆盖有效默认值 + best_cfg = {k: v for k, v in best_cfg.items() if v is not None} + + logger.info(f"[BEST] 已载入 {args.symbol} 最优参数: {best_cfg}") + except Exception as e: + logger.warning(f"[BEST] 读取最优参数失败,忽略 --use-best:{e}") + + # 合并 best_cfg 到 cfg(优先级:命令行 > best_cfg > cfg > 默认) + for k, v in (best_cfg or {}).items(): + if k not in cfg: + cfg[k] = v + # [PATCH B END] + + + + # 构造 run_once 的最终参数(先用 cfg 的,若命令行显式传入则覆盖) + def override(val, default, cfg_val): + """ + 如果命令行传入值 != argparse 的 default,说明用户显式设置 => 用命令行;否则用 cfg;再否则用 default + """ + if val != default: + return val + return cfg_val if (cfg_val is not None) else default + + # 取 argparse 默认值(用于判断是否显式覆盖) + defaults = vars(ap.parse_args([])) # 空参解析拿到默认表 + + kwargs = dict( + symbol = override(args.symbol, defaults["symbol"], cfg.get("symbol")), + csv_path = override(args.csv, defaults["csv"], cfg.get("csv_path")), + initial_cash = override(args.cash, defaults["cash"], cfg.get("cash")), + qty = override(args.qty, defaults["qty"], cfg.get("qty")), + account_ccy = override(args.account_ccy, defaults["account_ccy"], cfg.get("account_ccy")), + fast_win = override(args.fast, defaults["fast"], cfg.get("fast")), + slow_win = override(args.slow, defaults["slow"], cfg.get("slow")), + spread_pips = override(args.spread, defaults["spread"], cfg.get("spread")), + commission_per_million = override(args.comm, defaults["comm"], cfg.get("comm")), + slippage_pips = override(args.slip, defaults["slip"], cfg.get("slip")), + stop_loss_pips = override(args.sl, defaults["sl"], cfg.get("sl")), + take_profit_pips = override(args.tp, defaults["tp"], cfg.get("tp")), + atr_sl = override(args.atr_sl, defaults["atr_sl"], cfg.get("atr_sl")), + atr_tp = override(args.atr_tp, defaults["atr_tp"], cfg.get("atr_tp")), + atr_window = override(args.atr_window, defaults["atr_window"], cfg.get("atr_window")), + regime_ema_window = override(args.regime_ema_window, defaults["regime_ema_window"], cfg.get("regime_ema_window")), + regime_slope_min = override(args.regime_slope_min, defaults["regime_slope_min"], cfg.get("regime_slope_min")), + regime_atr_min = override(args.regime_atr_min, defaults["regime_atr_min"], cfg.get("regime_atr_min")), + regime_atr_percentile_min = override(args.regime_atr_percentile_min, defaults["regime_atr_percentile_min"], cfg.get("regime_atr_percentile_min")), + regime_atr_percentile_window = override(args.regime_atr_percentile_window, defaults["regime_atr_percentile_window"], cfg.get("regime_atr_percentile_window")), + regime_trend_min_bars = override(args.regime_trend_min_bars, defaults["regime_trend_min_bars"], cfg.get("regime_trend_min_bars")), + htf_factor = override(args.htf_factor, defaults["htf_factor"], cfg.get("htf_factor")), + htf_ema_window = override(args.htf_ema_window, defaults["htf_ema_window"], cfg.get("htf_ema_window")), + htf_rsi_period = override(args.htf_rsi_period, defaults["htf_rsi_period"], cfg.get("htf_rsi_period")), + rsi_period = override(args.rsi_period, defaults["rsi_period"], cfg.get("rsi_period")), + rsi_long_thresh = override(args.rsi_long_thresh, defaults["rsi_long_thresh"], cfg.get("rsi_long_thresh")), + rsi_short_thresh = override(args.rsi_short_thresh, defaults["rsi_short_thresh"], cfg.get("rsi_short_thresh")), + enable_trailing = override(args.enable_trailing, defaults["enable_trailing"], cfg.get("enable_trailing")), + trailing_enable_atr_mult = override(args.trailing_enable_atr_mult, defaults["trailing_enable_atr_mult"], cfg.get("trailing_enable_atr_mult")), + trailing_atr_mult = override(args.trailing_atr_mult, defaults["trailing_atr_mult"], cfg.get("trailing_atr_mult")), + long_only_above_slow = override(args.long_only_above_slow, defaults["long_only_above_slow"], cfg.get("long_only_above_slow")), + slope_lookback = override(args.slope_lookback, defaults["slope_lookback"], cfg.get("slope_lookback")), + cooldown = override(args.cooldown, defaults["cooldown"], cfg.get("cooldown")), + short_only_below_slow = override(args.short_only_below_slow, defaults["short_only_below_slow"], cfg.get("short_only_below_slow")), + risk_per_trade_pct = override(args.risk_percent, defaults["risk_percent"], cfg.get("risk_per_trade_pct")), + max_drawdown_pct = override(args.max_drawdown, defaults["max_drawdown"], cfg.get("max_drawdown_pct")), + max_position_units = override(args.max_units, defaults["max_units"], cfg.get("max_position_units")), + skip_outlier_entries = override(args.skip_outlier_entries, defaults["skip_outlier_entries"], cfg.get("skip_outlier_entries")), + cost_profiles = cfg_cost_profiles, + slippage_model = cfg_slippage_model, + strategy_mode = override(args.strategy_mode, defaults["strategy_mode"], cfg.get("strategy_mode")), + strategy_vote_threshold = override(args.strategy_vote_threshold, defaults["strategy_vote_threshold"], cfg.get("strategy_vote_threshold")), + stress_cost_spread_mult = override(args.stress_cost_spread_mult, defaults["stress_cost_spread_mult"], cfg.get("stress_cost_spread_mult")), + stress_cost_comm_mult = override(args.stress_cost_comm_mult, defaults["stress_cost_comm_mult"], cfg.get("stress_cost_comm_mult")), + stress_slippage_mult = override(args.stress_slippage_mult, defaults["stress_slippage_mult"], cfg.get("stress_slippage_mult")), + stress_price_vol_mult = override(args.stress_price_vol_mult, defaults["stress_price_vol_mult"], cfg.get("stress_price_vol_mult")), + stress_skip_trade_pct = override(args.stress_skip_trade_pct, defaults["stress_skip_trade_pct"], cfg.get("stress_skip_trade_pct")), + ) + # [PATCH C START] 关键窗口参数安全转为 int + try: + kwargs["fast_win"] = int(kwargs["fast_win"]) + kwargs["slow_win"] = int(kwargs["slow_win"]) + kwargs["atr_window"] = int(kwargs["atr_window"]) + kwargs["rsi_period"] = int(kwargs.get("rsi_period", 14)) + except Exception as e: + logger.error(f"参数类型转换错误,请检查 fast/slow/atr_window:{e}") + raise + # [PATCH C END] + + kwargs["short_only_below_slow"] = bool(kwargs["short_only_below_slow"]) + if kwargs["risk_per_trade_pct"] is not None: + kwargs["risk_per_trade_pct"] = float(kwargs["risk_per_trade_pct"]) + if kwargs["max_drawdown_pct"] is not None: + kwargs["max_drawdown_pct"] = float(kwargs["max_drawdown_pct"]) + if kwargs["max_position_units"] is not None: + kwargs["max_position_units"] = float(kwargs["max_position_units"]) + # RSI / trailing types + if kwargs.get("rsi_long_thresh") is not None: + kwargs["rsi_long_thresh"] = float(kwargs["rsi_long_thresh"]) + if kwargs.get("rsi_short_thresh") is not None: + kwargs["rsi_short_thresh"] = float(kwargs["rsi_short_thresh"]) + kwargs["enable_trailing"] = bool(kwargs.get("enable_trailing", False)) + kwargs["trailing_enable_atr_mult"] = float(kwargs.get("trailing_enable_atr_mult", 1.0)) + kwargs["trailing_atr_mult"] = float(kwargs.get("trailing_atr_mult", 0.5)) + kwargs["regime_ema_window"] = int(kwargs.get("regime_ema_window") or 0) + if kwargs.get("regime_slope_min") is not None: + kwargs["regime_slope_min"] = float(kwargs["regime_slope_min"]) + if kwargs.get("regime_atr_min") is not None: + kwargs["regime_atr_min"] = float(kwargs["regime_atr_min"]) + if kwargs.get("regime_atr_percentile_min") is not None: + kwargs["regime_atr_percentile_min"] = float(kwargs["regime_atr_percentile_min"]) + kwargs["regime_atr_percentile_window"] = int(kwargs.get("regime_atr_percentile_window") or 0) + kwargs["regime_trend_min_bars"] = int(kwargs.get("regime_trend_min_bars") or 0) + kwargs["htf_factor"] = int(kwargs.get("htf_factor") or 1) + if kwargs.get("htf_ema_window") is not None: + kwargs["htf_ema_window"] = int(kwargs["htf_ema_window"]) + if kwargs.get("htf_rsi_period") is not None: + kwargs["htf_rsi_period"] = int(kwargs["htf_rsi_period"]) + + if args.no_short != defaults["no_short"]: + allow_short = not args.no_short + else: + cfg_allow = cfg.get("allow_short") if cfg else None + allow_short = bool(cfg_allow) if cfg_allow is not None else True + kwargs["allow_short"] = allow_short + + kwargs["fx_rates"] = _merge_fx_rates(cfg_fx_rates, cli_fx_rates) + kwargs["strategies"] = cfg_strategies + + run_once(**kwargs) diff --git a/Q Research/scripts/build_dataset_manifest.py b/Q Research/scripts/build_dataset_manifest.py new file mode 100644 index 0000000..f999488 --- /dev/null +++ b/Q Research/scripts/build_dataset_manifest.py @@ -0,0 +1,183 @@ +#!/usr/bin/env python3 +""" +Generate a manifest describing datasets stored under `data/`. + +For each supported file (CSV/Parquet/Feather) this script records: + - relative path + - file size, mtime, and SHA256 checksum + - row count, columns, dtypes + - detected time range (based on common timestamp columns) + +Usage: + python scripts/build_dataset_manifest.py \ + --dirs data/raw data/derived \ + --output data/_manifest.json +""" + +from __future__ import annotations + +import argparse +import json +import hashlib +import os +from datetime import datetime, timezone +from pathlib import Path +from typing import Iterable, List, Optional + +import pandas as pd +from loguru import logger + + +SUPPORTED_EXTS = {".csv", ".parquet", ".feather"} +TIME_COLUMNS = ["ts", "time", "timestamp", "datetime", "date"] + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Build dataset manifest.") + parser.add_argument( + "--dirs", + nargs="+", + default=["data/raw", "data/derived"], + help="Directories to scan for datasets.", + ) + parser.add_argument( + "--output", + default="data/_manifest.json", + help="Path to write manifest JSON.", + ) + parser.add_argument( + "--limit", + type=int, + default=None, + help="Optional limit on number of rows to load when summarizing large files.", + ) + parser.add_argument( + "--dry-run", + action="store_true", + help="Scan files but do not write manifest.", + ) + return parser.parse_args() + + +def sha256sum(path: Path, chunk_size: int = 1 << 20) -> str: + """Compute SHA256 hash for a file.""" + h = hashlib.sha256() + with path.open("rb") as f: + while True: + chunk = f.read(chunk_size) + if not chunk: + break + h.update(chunk) + return h.hexdigest() + + +def load_dataframe(path: Path, limit: Optional[int]) -> pd.DataFrame: + """Load a dataframe from supported formats with optional row limit.""" + if path.suffix == ".csv": + return pd.read_csv(path, nrows=limit) + if path.suffix == ".parquet": + df = pd.read_parquet(path) + return df.head(limit) if limit else df + if path.suffix == ".feather": + df = pd.read_feather(path) + return df.head(limit) if limit else df + raise ValueError(f"Unsupported file format: {path}") + + +def detect_time_range(df: pd.DataFrame) -> tuple[Optional[str], Optional[str]]: + """Try to detect timestamp column and return ISO8601 min/max.""" + col = next((c for c in TIME_COLUMNS if c in df.columns), None) + if col is None: + for candidate in df.columns: + if pd.api.types.is_datetime64_any_dtype(df[candidate]): + col = candidate + break + if col is None: + return None, None + ts = pd.to_datetime(df[col], utc=True, errors="coerce").dropna() + if ts.empty: + return None, None + return ts.min().isoformat(), ts.max().isoformat() + + +def summarize_file(path: Path, repo_root: Path, limit: Optional[int]) -> dict: + """Collect metadata for the given dataset file.""" + rel_path = path.relative_to(repo_root) + stat = path.stat() + size = stat.st_size + mtime = datetime.fromtimestamp(stat.st_mtime, tz=timezone.utc).isoformat() + checksum = sha256sum(path) + + try: + df = load_dataframe(path, limit) + except Exception as exc: + logger.error(f"Failed to load {rel_path}: {exc}") + df = None + + rows = int(df.shape[0]) if df is not None else None + columns: List[str] = df.columns.tolist() if df is not None else [] + dtypes = {col: str(dtype) for col, dtype in df.dtypes.items()} if df is not None else {} + t_start, t_end = detect_time_range(df) if df is not None else (None, None) + + return { + "path": str(rel_path), + "extension": path.suffix, + "size_bytes": size, + "modified_at": mtime, + "sha256": checksum, + "rows": rows, + "columns": columns, + "dtypes": dtypes, + "time_start": t_start, + "time_end": t_end, + } + + +def iter_dataset_files(dirs: Iterable[str], repo_root: Path) -> List[Path]: + files: List[Path] = [] + for d in dirs: + target = (repo_root / d).resolve() + if not target.exists(): + logger.warning(f"Skip missing directory: {target}") + continue + for path in target.rglob("*"): + if path.is_file() and path.suffix in SUPPORTED_EXTS: + files.append(path) + return files + + +def main() -> None: + args = parse_args() + repo_root = Path(__file__).resolve().parents[1] + files = iter_dataset_files(args.dirs, repo_root) + if not files: + logger.warning("No dataset files found.") + logger.info(f"Found {len(files)} dataset files.") + + entries = [] + for path in files: + logger.info(f"Summarizing {path}") + entry = summarize_file(path, repo_root, args.limit) + entries.append(entry) + + manifest = { + "generated_at": datetime.now(timezone.utc).isoformat(), + "inputs": args.dirs, + "file_count": len(entries), + "files": entries, + } + + if args.dry_run: + logger.info("Dry-run enabled; manifest not written.") + print(json.dumps(manifest, indent=2, ensure_ascii=False)) + return + + output_path = (repo_root / args.output).resolve() + output_path.parent.mkdir(parents=True, exist_ok=True) + with output_path.open("w", encoding="utf-8") as fh: + json.dump(manifest, fh, indent=2, ensure_ascii=False) + logger.info(f"Manifest written to {output_path}") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/check_data_integrity.py b/Q Research/scripts/check_data_integrity.py new file mode 100644 index 0000000..e1a6163 --- /dev/null +++ b/Q Research/scripts/check_data_integrity.py @@ -0,0 +1,110 @@ +#!/usr/bin/env python3 +""" +CI helper that ensures watched datasets keep their expected signatures and pass validation. + +Fail conditions: + 1. Dataset listed in baseline JSON is missing. + 2. Manifest entry hash differs from the baseline (meaning data changed without approval). + 3. Validation severity == "error". +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +from loguru import logger + +ROOT = Path(__file__).resolve().parents[1] + +sys.path.append(str(ROOT)) + +from scripts.validate_dataset import compute_report, DEFAULT_MANIFEST # noqa + + +def load_manifest(manifest_path: Path) -> dict: + if not manifest_path.exists(): + raise FileNotFoundError(f"Manifest not found: {manifest_path}") + with manifest_path.open("r", encoding="utf-8") as fh: + return json.load(fh) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Check dataset hashes & validation severity.") + parser.add_argument("--baseline", default="data/signature_baseline.json", help="Baseline hash file.") + parser.add_argument("--manifest", default=DEFAULT_MANIFEST, help="Manifest JSON to compare against.") + parser.add_argument("--max-outlier-z", type=float, default=5.0, help="Z-score threshold used for validation.") + return parser.parse_args() + + +def load_baseline_entries(path: Path) -> list[dict]: + if not path.exists(): + raise FileNotFoundError(f"Baseline file missing: {path}") + with path.open("r", encoding="utf-8") as fh: + data = json.load(fh) + if isinstance(data, dict): + return [{"path": k, "sha256": v, "allow_drift": False} for k, v in data.items()] + if isinstance(data, list): + return data + raise ValueError("Baseline must be a dict or list of entries.") + + +def main() -> None: + args = parse_args() + baseline_path = (ROOT / args.baseline).resolve() + manifest_path = (ROOT / args.manifest).resolve() + + baseline_entries = load_baseline_entries(baseline_path) + + manifest = load_manifest(manifest_path) + manifest_index = {entry["path"]: entry for entry in manifest.get("files", [])} + + errors = [] + warnings = [] + for item in baseline_entries: + rel_path = item["path"] + expected_hash = item.get("sha256") + allow_drift = bool(item.get("allow_drift", False)) + dataset_path = (ROOT / rel_path).resolve() + manifest_entry = manifest_index.get(rel_path) + label = rel_path + if manifest_entry is None: + errors.append(f"{label}: missing from manifest {manifest_path}") + continue + actual_hash = manifest_entry.get("sha256") + if actual_hash != expected_hash: + message = ( + f"{label}: hash drift detected (expected {expected_hash}, actual {actual_hash}). " + "Update the baseline with justification if intentional." + ) + if allow_drift: + warnings.append(message) + else: + errors.append(message) + if not dataset_path.exists(): + errors.append(f"{label}: dataset file not found on disk") + continue + report = compute_report(dataset_path, manifest_entry, args.max_outlier_z) + severity = report.get("severity") + msg = f"{label}: validation severity={severity} ({'; '.join(report.get('messages', []))})" + if severity == "error": + errors.append(msg) + elif severity == "warn": + warnings.append(msg) + + for warn in warnings: + logger.warning(warn) + + if errors: + logger.error("Data integrity check failed:") + for err in errors: + logger.error(" - {}", err) + sys.exit(1) + + logger.info(f"Data integrity checks passed for {len(baseline_entries)} dataset(s).") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/check_risk_report.py b/Q Research/scripts/check_risk_report.py new file mode 100644 index 0000000..aab4166 --- /dev/null +++ b/Q Research/scripts/check_risk_report.py @@ -0,0 +1,42 @@ +#!/usr/bin/env python3 +"""Fail CI if risk report exceeds thresholds.""" + +from __future__ import annotations + +import argparse +import csv +from pathlib import Path + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Validate risk report thresholds") + parser.add_argument("--report", default="results/risk/report.csv", help="CSV produced by risk_report.py") + parser.add_argument("--max-rejects", type=int, default=0) + parser.add_argument("--max-kill", type=int, default=0) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + report = Path(args.report) + if not report.exists(): + raise SystemExit(f"Report not found: {report}") + rejects = 0 + kills = 0 + with report.open("r", encoding="utf-8") as fh: + reader = csv.DictReader(fh) + for row in reader: + if row.get("event") == "reject": + rejects += 1 + elif row.get("event") == "kill_switch": + kills += 1 + if rejects > args.max_rejects or kills > args.max_kill: + raise SystemExit( + f"Risk report exceeded thresholds: rejects={rejects} (max {args.max_rejects}), " + f"kill_switch={kills} (max {args.max_kill})" + ) + print(f"Risk report OK (rejects={rejects}, kill_switch={kills})") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/cleanup_outputs.py b/Q Research/scripts/cleanup_outputs.py new file mode 100644 index 0000000..952df90 --- /dev/null +++ b/Q Research/scripts/cleanup_outputs.py @@ -0,0 +1,17 @@ +import os, time, glob +from datetime import datetime, timedelta + +base = "data/outputs" +days_to_keep = 3 +cutoff = time.time() - days_to_keep * 86400 + +def clean_dir(path, exts): + for ext in exts: + for f in glob.glob(os.path.join(path, f"*.{ext}")): + if os.path.getmtime(f) < cutoff: + os.remove(f) + print(f"🧹 Deleted old {ext}: {f}") + +clean_dir(os.path.join(base, "equity"), ["csv"]) +clean_dir(os.path.join(base, "trades"), ["csv"]) +print("✅ Cleanup complete.") \ No newline at end of file diff --git a/Q Research/scripts/compare_cost_scenarios.py b/Q Research/scripts/compare_cost_scenarios.py new file mode 100644 index 0000000..0e49d28 --- /dev/null +++ b/Q Research/scripts/compare_cost_scenarios.py @@ -0,0 +1,70 @@ +from __future__ import annotations + +import argparse +from pathlib import Path +from typing import List + +import pandas as pd +from loguru import logger + +METRIC_COLS = [ + "sharpe", + "ann_return", + "ann_vol", + "max_drawdown", + "expectancy", + "trades", + "win_rate", + "return_pct", + "rr", + "median_hold", + "symbol", + "source_file", +] + + +def main(): + parser = argparse.ArgumentParser(description="Compare zero-cost vs real-cost grid results.") + parser.add_argument("--zero", required=True, help="CSV from zero-cost grid run.") + parser.add_argument("--real", required=True, help="CSV from real-cost grid run.") + parser.add_argument("--out", default="data/grid/cost_comparison.csv", help="Output CSV path.") + parser.add_argument("--top", type=int, default=20, help="Print top-N rows with smallest Sharpe delta.") + args = parser.parse_args() + + df_zero = pd.read_csv(args.zero) + df_real = pd.read_csv(args.real) + + shared_cols = [c for c in df_zero.columns if c in df_real.columns] + if not shared_cols: + raise ValueError("No overlapping columns between zero and real CSVs.") + + param_cols: List[str] = [c for c in shared_cols if c not in METRIC_COLS] + if not param_cols: + raise ValueError("Unable to infer parameter columns for join; ensure CSVs contain metrics from METRIC_COLS.") + + z = df_zero.rename(columns={col: f"{col}_zero" for col in METRIC_COLS if col in df_zero.columns}) + r = df_real.rename(columns={col: f"{col}_real" for col in METRIC_COLS if col in df_real.columns}) + + merged = z.merge(r, on=param_cols, how="inner", suffixes=("_zero", "_real")) + if merged.empty: + raise RuntimeError("Join result is empty; ensure both CSVs share the same parameter combinations.") + + if "sharpe_zero" in merged.columns and "sharpe_real" in merged.columns: + merged["delta_sharpe"] = merged["sharpe_real"] - merged["sharpe_zero"] + if "ann_return_zero" in merged.columns and "ann_return_real" in merged.columns: + merged["delta_ann_return"] = merged["ann_return_real"] - merged["ann_return_zero"] + if "expectancy_zero" in merged.columns and "expectancy_real" in merged.columns: + merged["delta_expectancy"] = merged["expectancy_real"] - merged["expectancy_zero"] + + out_path = Path(args.out) + out_path.parent.mkdir(parents=True, exist_ok=True) + merged.to_csv(out_path, index=False) + logger.info(f"Cost comparison saved to {out_path} (rows={len(merged)})") + + if "delta_sharpe" in merged.columns: + top_df = merged.sort_values("delta_sharpe", ascending=False).head(args.top) + print(top_df[param_cols + ["sharpe_zero", "sharpe_real", "delta_sharpe"]].to_string(index=False)) + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/compare_fills.py b/Q Research/scripts/compare_fills.py new file mode 100644 index 0000000..19e80d4 --- /dev/null +++ b/Q Research/scripts/compare_fills.py @@ -0,0 +1,82 @@ +#!/usr/bin/env python3 +""" +Compare paper vs. live fills to produce a lightweight TCA summary. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import pandas as pd + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Compare paper vs live fills.") + parser.add_argument("--paper", required=True, help="Path to paper fills CSV.") + parser.add_argument("--live", required=True, help="Path to live fills CSV.") + parser.add_argument("--out", default="results/execution/tca_summary.json", help="Output JSON path.") + return parser.parse_args() + + +def load(path: str) -> pd.DataFrame: + csv = Path(path) + if not csv.exists(): + raise SystemExit(f"fills CSV not found: {csv}") + df = pd.read_csv(csv) + if df.empty: + raise SystemExit(f"{csv} is empty") + if "ts" in df.columns: + df["ts"] = pd.to_datetime(df["ts"]) + return df + + +def summary(df: pd.DataFrame, label: str) -> dict: + pnl = df["pnl"] if "pnl" in df.columns else pd.Series(dtype=float) + latency = df["adapter_latency_ms"] if "adapter_latency_ms" in df.columns else pd.Series(dtype=float) + return { + f"{label}_trade_count": int(len(df)), + f"{label}_total_pnl": float(pnl.sum()) if not pnl.empty else None, + f"{label}_avg_pnl": float(pnl.mean()) if not pnl.empty else None, + f"{label}_avg_latency_ms": float(latency.mean()) if not latency.empty else None, + } + + +def pnl_gap(paper: pd.DataFrame, live: pd.DataFrame) -> dict: + cols = [] + if "ts" in paper.columns and "ts" in live.columns: + merged = pd.merge( + paper[["ts", "pnl"]].rename(columns={"pnl": "paper_pnl"}), + live[["ts", "pnl"]].rename(columns={"pnl": "live_pnl"}), + on="ts", + how="outer", + ) + else: + merged = pd.DataFrame({"paper_pnl": paper.get("pnl"), "live_pnl": live.get("pnl")}) + merged = merged.fillna(0.0) + merged["pnl_diff"] = merged["live_pnl"] - merged["paper_pnl"] + return { + "pnl_diff_mean": float(merged["pnl_diff"].mean()), + "pnl_diff_std": float(merged["pnl_diff"].std(ddof=0)), + } + + +def main() -> None: + args = parse_args() + paper = load(args.paper) + live = load(args.live) + + report = {"paper_path": args.paper, "live_path": args.live} + report.update(summary(paper, "paper")) + report.update(summary(live, "live")) + report.update(pnl_gap(paper, live)) + + out_path = Path(args.out) + out_path.parent.mkdir(parents=True, exist_ok=True) + out_path.write_text(json.dumps(report, indent=2, default=float), encoding="utf-8") + print(json.dumps(report, indent=2, default=float)) + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/compute_indicators.py b/Q Research/scripts/compute_indicators.py new file mode 100644 index 0000000..d4fff63 --- /dev/null +++ b/Q Research/scripts/compute_indicators.py @@ -0,0 +1,50 @@ +# compute_indicators.py + +import os +import pandas as pd +from loguru import logger + +BASE_DIR = os.path.dirname(os.path.dirname(__file__)) +RAW_DATA_DIR = os.path.join(BASE_DIR, "data", "raw") +DERIVED_DATA_DIR = os.path.join(BASE_DIR, "data", "derived") +os.makedirs(DERIVED_DATA_DIR, exist_ok=True) + +def compute_indicators(df: pd.DataFrame) -> pd.DataFrame: + # SMA + df['SMA_20'] = df['close'].rolling(window=20).mean() + + # Bollinger Bands + df['BB_MID'] = df['close'].rolling(window=20).mean() + df['BB_STD'] = df['close'].rolling(window=20).std() + df['BB_UPPER'] = df['BB_MID'] + 2 * df['BB_STD'] + df['BB_LOWER'] = df['BB_MID'] - 2 * df['BB_STD'] + + # MACD + ema12 = df['close'].ewm(span=12, adjust=False).mean() + ema26 = df['close'].ewm(span=26, adjust=False).mean() + df['MACD'] = ema12 - ema26 + df['MACD_signal'] = df['MACD'].ewm(span=9, adjust=False).mean() + + # RSI + def compute_rsi(series, period=14): + delta = series.diff() + gain = delta.clip(lower=0) + loss = -delta.clip(upper=0) + avg_gain = gain.rolling(window=period).mean() + avg_loss = loss.rolling(window=period).mean() + rs = avg_gain / avg_loss + return 100 - (100 / (1 + rs)) + + df['RSI_14'] = compute_rsi(df['close']) + + return df + + +if __name__ == "__main__": + input_path = os.path.join(RAW_DATA_DIR, "EURUSD_H1.csv") + output_path = os.path.join(DERIVED_DATA_DIR, "EURUSD_H1_with_indicators.csv") + + df = pd.read_csv(input_path, parse_dates=["time"]) + df = compute_indicators(df) + df.to_csv(output_path, index=False) + logger.info(f"✅ Saved with indicators to {output_path}") diff --git a/Q Research/scripts/export_metrics_prom.py b/Q Research/scripts/export_metrics_prom.py new file mode 100644 index 0000000..c9ebeee --- /dev/null +++ b/Q Research/scripts/export_metrics_prom.py @@ -0,0 +1,65 @@ +#!/usr/bin/env python3 +""" +Export latest risk metrics row in Prometheus exposition format. + +Usage: + python scripts/export_metrics_prom.py --csv results/risk/metrics.csv --job risk_sim +""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +import pandas as pd + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Export risk metrics to Prometheus format.") + parser.add_argument("--csv", default="results/risk/metrics.csv") + parser.add_argument("--job", default="risk_sim") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + path = Path(args.csv) + if not path.exists(): + raise SystemExit(f"metrics CSV not found: {path}") + df = pd.read_csv(path) + if df.empty: + raise SystemExit("metrics CSV is empty.") + latest = df.tail(1).iloc[0] + run_id = latest.get("run_id", "unknown") + def value(key: str, default: float = 0.0) -> float: + val = latest.get(key, default) + if isinstance(val, str) and not val: + return default + try: + if pd.isna(val): + return default + except TypeError: + pass + return val + + metrics = { + "risk_rejects": value("rejects", 0), + "risk_kills": value("kills", 0), + "risk_latency_ms_avg": value("latency_ms_avg"), + "risk_latency_ms_p95": value("latency_ms_p95"), + "risk_total_pnl": value("total_pnl"), + "risk_max_gross_notional": value("max_gross_notional"), + "risk_max_symbol_exposure": value("max_symbol_exposure"), + "risk_max_drawdown_pct": value("max_drawdown_pct"), + "risk_live_sharpe_30d": value("rolling_sharpe_30d"), + "risk_live_drawdown_pct": value("live_drawdown_pct"), + "risk_live_latency_ms_p95": value("live_latency_ms_p95"), + "risk_slippage_bps": value("slippage_bps"), + } + labels = f'run="{run_id}",job="{args.job}"' + for name, value in metrics.items(): + print(f'{name}{{{labels}}} {value}') + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/get_candles.py b/Q Research/scripts/get_candles.py new file mode 100644 index 0000000..ce87e95 --- /dev/null +++ b/Q Research/scripts/get_candles.py @@ -0,0 +1,200 @@ +import math +import os +import sys +from datetime import datetime, timedelta, timezone + +import pandas as pd +from loguru import logger + +# 允许从项目根目录导入模块 +PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +REPO_ROOT = os.path.dirname(PROJECT_ROOT) +sys.path.append(PROJECT_ROOT) +sys.path.append(REPO_ROOT) + +from oandapyV20 import API +import oandapyV20.endpoints.instruments as instruments +from shared.utils.config import OANDA_TOKEN # 确保这里能拿到 token + +# 日志 +logger.remove() +logger.add(sys.stderr, level="INFO", format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {message}") + +def normalize_to_oanda(symbol: str) -> str: + """ + 将用户友好的交易对格式转换为OANDA格式,例如: + 'eurusd', 'EUR-USD', 'eur_usd' -> 'EUR_USD' + """ + s = symbol.upper().replace("-", "_").replace(" ", "").replace("/", "_") + # 如果已经是正确格式,直接返回 + if "_" in s and len(s) == 7: + return s + # 尝试拆分为两部分 + if len(s) == 6: + return s[:3] + "_" + s[3:] + return s + +def default_out_csv(project_root: str, instrument: str, granularity: str) -> str: + """ + 根据交易对和时间粒度生成默认输出路径,如: + data/raw/EURUSD_H1.csv + """ + fname = f"{instrument.replace('_','')}_{granularity}.csv" + raw_dir = os.path.join(project_root, "data", "raw") + os.makedirs(raw_dir, exist_ok=True) + return os.path.join(raw_dir, fname) + +def _bars_per_day(granularity: str) -> float: + """ + Rough estimate of bars per day for常见 OANDA 粒度。 + 用于根据 count 推算需要拉取的天数。 + """ + granularity = granularity.upper() + seconds_map = { + "S5": 5, + "S10": 10, + "S15": 15, + "S30": 30, + "M1": 60, + "M2": 120, + "M4": 240, + "M5": 300, + "M10": 600, + "M15": 900, + "M30": 1800, + "H1": 3600, + "H2": 7200, + "H3": 10800, + "H4": 14400, + "H6": 21600, + "H8": 28800, + "H12": 43200, + "D": 86400, + "W": 86400 * 5, + "M": 86400 * 21, + } + seconds = seconds_map.get(granularity, 3600) + if seconds <= 0: + return 24 + return max(86400 / seconds, 1) + + +def get_candles(symbol="EUR_USD", granularity="H1", start_days_ago=365, target_count: int | None = None) -> pd.DataFrame: + """ + 循环抓取 OANDA 历史K线(默认过去一年),自动分页拼接。 + """ + if not OANDA_TOKEN: + raise RuntimeError("OANDA_TOKEN 为空,请在 utils/config.py 配置或通过环境变量提供。") + + client = API(access_token=OANDA_TOKEN) + + end = datetime.now(timezone.utc) + start = end - timedelta(days=start_days_ago) + cur = start + all_rows = 0 + parts = [] + + logger.info(f"开始下载 {symbol} {granularity}(过去 {start_days_ago} 天)") + + # 每次抓 20 天(H1 ≈ 480 根),避免单次数据过大 + step = timedelta(days=20) + + while cur < end: + to_ts = min(cur + step, end) + params = { + "granularity": granularity, + "price": "M", + "from": cur.isoformat(), + "to": to_ts.isoformat(), + } + r = instruments.InstrumentsCandles(instrument=symbol, params=params) + try: + client.request(r) + except Exception as e: + logger.error(f"请求失败 {cur} ~ {to_ts}: {e}") + break + + candles = r.response.get("candles", []) + if not candles: + logger.warning(f"区间无数据:{cur} ~ {to_ts}") + cur = to_ts + continue + + data = [{ + "time": c["time"], + "open": float(c["mid"]["o"]), + "high": float(c["mid"]["h"]), + "low": float(c["mid"]["l"]), + "close":float(c["mid"]["c"]), + "volume": c["volume"] + } for c in candles if c.get("complete")] + + if data: + df_part = pd.DataFrame(data) + parts.append(df_part) + all_rows += len(df_part) + logger.info(f"抓取区间 {cur:%Y-%m-%d} ~ {to_ts:%Y-%m-%d} 行数={len(df_part)},累计={all_rows}") + + cur = to_ts # 推进窗口 + + if target_count and all_rows >= target_count: + logger.info(f"已满足目标条数 {target_count},停止抓取。") + break + + if not parts: + logger.warning("没有获取到任何数据。") + return pd.DataFrame() + + df = pd.concat(parts, ignore_index=True) + df["time"] = pd.to_datetime(df["time"]) + df = df.sort_values("time").drop_duplicates(subset=["time"]).reset_index(drop=True) + if target_count: + df = df.tail(target_count).reset_index(drop=True) + + logger.info(f"✅ 下载完成:总计 {len(df)} 行({df['time'].min()} ~ {df['time'].max()})") + return df + +def main(): + import argparse + parser = argparse.ArgumentParser() + parser.add_argument("--symbol", default="EUR_USD") + parser.add_argument("--granularity", default="H1") + parser.add_argument("--days", type=int, default=365, help="向前回溯天数") + parser.add_argument("--count", type=int, default=None, help="(可选)需要的 K 线数量,脚本会根据粒度估算天数,抓够后截断") + parser.add_argument("--out", "--output", dest="out", default=None) + args = parser.parse_args() + + # 规范化交易对格式 + args.symbol = normalize_to_oanda(args.symbol) + + # 自动生成输出路径(如果未指定) + if args.out is None: + args.out = default_out_csv(PROJECT_ROOT, args.symbol, args.granularity) + + # 确保输出目录存在 + os.makedirs(os.path.dirname(args.out), exist_ok=True) + + target_count = args.count if args.count and args.count > 0 else None + if target_count: + est_days = math.ceil(target_count / _bars_per_day(args.granularity)) + 5 + if est_days > args.days: + logger.info(f"根据 count={target_count} 估算需要 {est_days} 天数据(原 days={args.days}),已自动扩展。") + args.days = est_days + + try: + df = get_candles(args.symbol, args.granularity, args.days, target_count=target_count) + except Exception as e: + logger.exception(f"下载失败:{e}") + sys.exit(1) + + if df.empty: + logger.warning("结果为空,未保存。") + sys.exit(2) + + df.to_csv(args.out, index=False) + logger.info(f"📦 已保存到:{args.out}") + # 方便你肉眼确认 + logger.info(f"尾部预览:\n{df.tail(3).to_string(index=False)}") + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/grid_search_rsi_trailing.py b/Q Research/scripts/grid_search_rsi_trailing.py new file mode 100644 index 0000000..8c97cab --- /dev/null +++ b/Q Research/scripts/grid_search_rsi_trailing.py @@ -0,0 +1,490 @@ +from __future__ import annotations + +import argparse +import itertools +import os +import sys +from copy import deepcopy +from pathlib import Path +from queue import Empty, Queue +from typing import Dict, Iterable, List, Optional + +import pandas as pd +import yaml +from loguru import logger + +# 允许直接 import 项目内模块 +sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from core.backtest.strategy_engine import StrategyEngine, parse_strategy_specs, _coerce_fx_rates # noqa: E402 +from data.csv_feed import CSVFeed # noqa: E402 +from metrics.perf import trade_stats # noqa: E402 + + +BASE_CONFIG_PATH = Path("config/optimized_eurusd_v2_with_rsi.yaml") +DEFAULT_SYMBOL = "EURUSD" +DEFAULT_CSV = Path("data/raw/EURUSD_H1.csv") + +# 重点围绕 ATR / RSI / cooldown 进行调参 +PARAM_GRID = { + "fast": [20], + "slow": [120], + "atr_sl": [1.0, 1.3, 1.6], + "atr_tp": [None, 3.0, 4.5], + "rsi_long_thresh": [60], + "rsi_short_thresh": [40], + "cooldown": [12, 24, 36], + "trailing_enable_atr_mult": [0.5], + "trailing_atr_mult": [0.5], + "htf_factor": [4], + "size_tier_mode": ["base"], + "boll_window": [32, 48, 64], + "boll_enter_z": [1.0, 1.3, 1.6], + "boll_exit_z": [0.2, 0.4], + "boll_allow_short": [False], +} + +SIZE_TIER_PRESETS = { + "base": { + "base_size_mult": 1.0, + "size_tiers": [{"size_mult": 1.0}], + }, + "balanced": { + "base_size_mult": 1.0, + "size_tiers": [ + {"name": "strong", "min_atr_pct": 0.55, "min_trend_bars": 6, "size_mult": 1.3}, + {"name": "base", "size_mult": 1.0}, + ], + }, + "aggressive": { + "base_size_mult": 1.0, + "size_tiers": [ + {"name": "strong", "min_atr_pct": 0.6, "min_trend_bars": 8, "size_mult": 1.6}, + {"name": "mid", "min_atr_pct": 0.45, "min_trend_strength": 0.00008, "size_mult": 1.2}, + ], + }, +} + + +def _apply_size_tier_mode(cfg: Dict[str, object], mode: Optional[str]) -> None: + if not mode: + return + preset = SIZE_TIER_PRESETS.get(mode) + if not preset: + logger.warning(f"[GRID] 未知 size_tier 模式: {mode}") + return + strategies = cfg.get("strategies") + if not isinstance(strategies, list): + return + for strat in strategies: + if isinstance(strat, dict) and strat.get("name") == "regime_sma": + params = strat.setdefault("params", {}) + tiers = preset.get("size_tiers") + if tiers is not None: + params["size_tiers"] = deepcopy(tiers) + if "base_size_mult" in preset: + params["base_size_mult"] = preset["base_size_mult"] + + +def _apply_bollinger_params(cfg: Dict[str, object], overrides: Dict[str, object]) -> None: + strategies = cfg.get("strategies") + if not isinstance(strategies, list): + return + for strat in strategies: + if isinstance(strat, dict) and strat.get("name") == "bollinger_mean_revert": + params = strat.setdefault("params", {}) + if overrides.get("window") is not None: + params["window"] = int(overrides["window"]) + if overrides.get("enter_z") is not None: + params["enter_z"] = float(overrides["enter_z"]) + if overrides.get("exit_z") is not None: + params["exit_z"] = float(overrides["exit_z"]) + if overrides.get("allow_short") is not None: + params["allow_short"] = bool(overrides["allow_short"]) + + +def _to_float(value: Optional[object]) -> Optional[float]: + if value is None: + return None + if isinstance(value, str) and value.lower() in {"none", "null", ""}: + return None + return float(value) + + +def load_base_config(path: Path = BASE_CONFIG_PATH) -> Dict[str, object]: + with path.open("r", encoding="utf-8") as fh: + return yaml.safe_load(fh) or {} + + +def evaluate_config(base_cfg: Dict[str, object], overrides: Dict[str, object], symbol: str) -> Optional[Dict[str, object]]: + cfg = deepcopy(base_cfg) + local_overrides = dict(overrides) + size_mode = local_overrides.pop("size_tier_mode", None) + boll_window = local_overrides.pop("boll_window", None) + boll_enter = local_overrides.pop("boll_enter_z", None) + boll_exit = local_overrides.pop("boll_exit_z", None) + boll_allow_short = local_overrides.pop("boll_allow_short", None) + cfg.update(local_overrides) + if size_mode: + _apply_size_tier_mode(cfg, size_mode) + if any(v is not None for v in [boll_window, boll_enter, boll_exit, boll_allow_short]): + _apply_bollinger_params( + cfg, + { + "window": boll_window, + "enter_z": boll_enter, + "exit_z": boll_exit, + "allow_short": boll_allow_short, + }, + ) + + symbol = cfg.get("symbol", DEFAULT_SYMBOL) + csv_path = Path(cfg.get("csv", DEFAULT_CSV)) + if not csv_path.exists(): + logger.error(f"CSV 路径不存在: {csv_path}") + return None + + initial_cash = float(cfg.get("cash", 100_000)) + qty = float(cfg.get("qty", 10_000)) + account_ccy = cfg.get("account_ccy", "USD") + fast = int(cfg.get("fast", 20)) + slow = int(cfg.get("slow", 150)) + spread = float(cfg.get("spread", 1.0)) + slip = float(cfg.get("slip", 0.2)) + comm = float(cfg.get("comm", 2.0)) + stop_loss_pips = _to_float(cfg.get("sl")) + take_profit_pips = _to_float(cfg.get("tp")) + atr_sl = _to_float(cfg.get("atr_sl")) + atr_tp = _to_float(cfg.get("atr_tp")) + atr_window = int(cfg.get("atr_window", 14)) + rsi_period = int(cfg.get("rsi_period", 14)) + rsi_long = _to_float(cfg.get("rsi_long_thresh")) + rsi_short = _to_float(cfg.get("rsi_short_thresh")) + enable_trailing = bool(cfg.get("enable_trailing", False)) + trailing_enable = float(cfg.get("trailing_enable_atr_mult", 1.0)) + trailing_mult = float(cfg.get("trailing_atr_mult", 0.5)) + slope_lookback = int(cfg.get("slope_lookback", 0)) + cooldown = int(cfg.get("cooldown", 0)) + allow_short = bool(cfg.get("allow_short", True)) + long_only_above_slow = bool(cfg.get("long_only_above_slow", False)) + short_only_below_slow = bool(cfg.get("short_only_below_slow", False)) + risk_per_trade_pct = _to_float(cfg.get("risk_per_trade_pct")) + max_drawdown_pct = _to_float(cfg.get("max_drawdown_pct")) + max_position_units = _to_float(cfg.get("max_position_units")) + + cfg_fx_rates = _coerce_fx_rates(cfg.get("fx_rates")) + fx_rates = cfg_fx_rates if cfg_fx_rates else None + + strategy_specs = parse_strategy_specs(cfg.get("strategies")) + + engine = StrategyEngine( + symbol=symbol, + fast_win=fast, + slow_win=slow, + spread_pips=spread, + commission_per_million=comm, + slippage_pips=slip, + stop_loss_pips=stop_loss_pips, + take_profit_pips=take_profit_pips, + atr_sl=atr_sl, + atr_tp=atr_tp, + atr_window=atr_window, + rsi_period=rsi_period, + rsi_long_thresh=rsi_long, + rsi_short_thresh=rsi_short, + enable_trailing=enable_trailing, + trailing_enable_atr_mult=trailing_enable, + trailing_atr_mult=trailing_mult, + long_only_above_slow=long_only_above_slow, + slope_lookback=slope_lookback, + cooldown=cooldown, + qty=qty, + account_ccy=account_ccy, + fx_rates=fx_rates, + strategy_specs=strategy_specs, + allow_short=allow_short, + short_only_below_slow=short_only_below_slow, + risk_per_trade_pct=risk_per_trade_pct, + max_drawdown_pct=max_drawdown_pct, + max_position_units=max_position_units, + ) + engine.set_initial_cash(initial_cash) + + q: Queue = Queue() + feed = CSVFeed(q, path=str(csv_path), symbol=symbol) + feed.start() + + try: + while True: + try: + event = q.get(timeout=0.05) + except Empty: + if hasattr(feed, "pump"): + feed.pump(n=100) + if getattr(feed, "finished", False): + break + continue + + if event.get("type") != "bar": + continue + engine.handle_bar(event) + finally: + engine.finalize() + + summary = engine.summary(fast, slow) + stats = trade_stats(engine.trade_log) if engine.trade_log else {} + + final_equity = summary.get("final_equity", engine.cash) + ret_pct = (final_equity / initial_cash - 1.0) if initial_cash else None + + return { + "params": overrides, + "summary": summary, + "stats": stats, + "final_equity": final_equity, + "return_pct": ret_pct, + "symbol": symbol, + } + + +def param_product(grid: Dict[str, Iterable[object]]) -> Iterable[Dict[str, object]]: + keys = list(grid.keys()) + for combo in itertools.product(*(grid[k] for k in keys)): + params = dict(zip(keys, combo)) + fast_v = params.get("fast") + slow_v = params.get("slow") + if fast_v is not None and slow_v is not None: + if float(fast_v) >= float(slow_v): + continue + long_v = params.get("rsi_long_thresh") + short_v = params.get("rsi_short_thresh") + if long_v is not None and short_v is not None and long_v <= short_v: + continue + yield params + + +def run_grid(base_cfg: Dict[str, object], save_suffix: Optional[str] = None, param_grid: Optional[Dict[str, List[object]]] = None) -> Path: + # 降低日志噪音 + logger.remove() + logger.add(sys.stderr, level="WARNING") + + results: List[Dict[str, object]] = [] + grid_def = deepcopy(param_grid or PARAM_GRID) + combos = list(param_product(grid_def)) + total = len(combos) + print(f"将测试 {total} 种参数组合") + + suffix = f"_{save_suffix}" if save_suffix else "" + out_dir = Path("data/grid") + out_dir.mkdir(parents=True, exist_ok=True) + out_filename = f"grid_rsi_trailing_diagnostics{suffix}.csv" + + def fmt_pct(value: Optional[float]) -> str: + return f"{value:.2%}" if value is not None else "NA" + + def fmt_float(value: Optional[float], digits: int = 3) -> str: + return f"{value:.{digits}f}" if value is not None else "NA" + + for idx, params in enumerate(combos, 1): + print(f"\n[{idx}/{total}] 评估参数: {params}") + res = evaluate_config(base_cfg, params, base_cfg.get("symbol", DEFAULT_SYMBOL)) + if not res: + print(" -> 运行失败") + continue + summary = res["summary"] + stats = res["stats"] + sharpe = summary.get("sharpe") + win_rate = stats.get("win_rate") + exp = stats.get("expectancy") + trades = summary.get("trades") + ret_pct = res["return_pct"] + dd = summary.get("max_drawdown") + print( + " -> Sharpe={} 回撤={} 胜率={} 期望={} 交易数={}".format( + fmt_float(sharpe), + fmt_pct(dd), + fmt_pct(win_rate), + fmt_float(exp, digits=2), + trades if trades is not None else "NA", + ) + ) + res_record = { + "sharpe": sharpe, + "ann_return": summary.get("ann_return"), + "ann_vol": summary.get("ann_vol"), + "max_drawdown": summary.get("max_drawdown"), + "trades": trades, + "return_pct": ret_pct, + "win_rate": win_rate, + "rr": stats.get("rr"), + "expectancy": exp, + "median_hold": stats.get("median_hold"), + "symbol": res.get("symbol", base_cfg.get("symbol", DEFAULT_SYMBOL)), + **params, + "source_file": out_filename, + } + results.append(res_record) + + if not results: + print("\n未得到任何有效结果。") + return + + df = pd.DataFrame(results) + df["sharpe"] = pd.to_numeric(df["sharpe"], errors="coerce") + df["expectancy"] = pd.to_numeric(df["expectancy"], errors="coerce") + df = df.sort_values(by="sharpe", ascending=False) + out_path = out_dir / out_filename + df.to_csv(out_path, index=False) + + print(f"\n已保存全部结果 -> {out_path}") + top_n = df.head(5) + print("\nTop 5 组合:") + for _, row in top_n.iterrows(): + win_rate_val = row.get("win_rate") + win_rate_val = None if pd.isna(win_rate_val) else win_rate_val + cooldown_val = row.get("cooldown") + cooldown_disp = int(cooldown_val) if cooldown_val is not None and not pd.isna(cooldown_val) else "NA" + atr_sl_disp = fmt_float(row.get("atr_sl"), digits=2) + atr_tp_val = row.get("atr_tp") + atr_tp_disp = "None" if atr_tp_val is None or (isinstance(atr_tp_val, float) and pd.isna(atr_tp_val)) else fmt_float(atr_tp_val, digits=2) + rsi_short = row.get("rsi_short_thresh") + rsi_long = row.get("rsi_long_thresh") + fast_val = row.get("fast") + slow_val = row.get("slow") + print( + " Sharpe={} 回撤={} 胜率={} fast/slow={}/{} cooldown={} atr_sl={} atr_tp={} RSI=({}/{})".format( + fmt_float(row.get("sharpe")), + fmt_pct(row.get("max_drawdown")), + fmt_pct(win_rate_val), + fmt_float(fast_val, digits=0) if fast_val is not None and not pd.isna(fast_val) else "NA", + fmt_float(slow_val, digits=0) if slow_val is not None and not pd.isna(slow_val) else "NA", + cooldown_disp, + atr_sl_disp, + atr_tp_disp, + fmt_float(rsi_short, digits=1) if rsi_short is not None and not pd.isna(rsi_short) else "NA", + fmt_float(rsi_long, digits=1) if rsi_long is not None and not pd.isna(rsi_long) else "NA", + ) + ) + + usdjpy_mask = ( + df["symbol"].astype(str).str.upper().eq("USDJPY") + & df["sharpe"].gt(1.5) + & df["expectancy"].gt(2.0) + ) + top_usdjpy = df.loc[usdjpy_mask].head(3) + if not top_usdjpy.empty: + print("\nUSDJPY Sharpe>1.5 & Expectancy>$2 (Top 3):") + for _, row in top_usdjpy.iterrows(): + print( + " Sharpe={:.3f} Expectancy=${:.2f} Trades={} atr_sl={} atr_tp={} cooldown={}".format( + row["sharpe"], + row["expectancy"], + row.get("trades", "NA"), + fmt_float(row.get("atr_sl"), digits=2), + "None" if pd.isna(row.get("atr_tp")) or row.get("atr_tp") is None else fmt_float(row.get("atr_tp"), digits=2), + int(row.get("cooldown")) if row.get("cooldown") is not None and not pd.isna(row.get("cooldown")) else "NA", + ) + ) + best_out = out_dir / "grid_usdjpy_top3.csv" + top_usdjpy.to_csv(best_out, index=False) + print(f"\n已保存 USDJPY 筛选结果 -> {best_out}") + else: + print("\nUSDJPY 暂无满足 Sharpe>1.5 & Expectancy>$2 的组合。") + + return out_path + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Grid search for ATR/RSI/trailing parameters.") + parser.add_argument( + "--config", + type=str, + default=str(BASE_CONFIG_PATH), + help="YAML 配置路径(默认使用 optimized_eurusd_v2_with_rsi.yaml)", + ) + parser.add_argument("--symbol", type=str, default=None, help="覆盖配置中的 symbol(可选)") + parser.add_argument("--csv", type=str, default=None, help="覆盖配置中的 csv 路径(可选)") + parser.add_argument("--suffix", type=str, default=None, help="输出文件名后缀(默认取 symbol)") + parser.add_argument("--atr-sl", type=str, default=None, help="自定义 atr_sl 列表,例如 '1.0,1.3,1.6'") + parser.add_argument("--atr-tp", type=str, default=None, help="自定义 atr_tp 列表,例如 'None,3.0,4.0'") + parser.add_argument("--cooldown-list", type=str, default=None, help="自定义 cooldown 列表,例如 '12,24,36'") + parser.add_argument("--htf-factor-list", type=str, default=None, help="自定义 htf_factor 列表,例如 '2,4,6'") + parser.add_argument("--size-tier-mode-list", type=str, default=None, help="size_tier 模式列表,例如 'base,aggressive'") + parser.add_argument("--boll-window-list", type=str, default=None, help="Bollinger 窗口列表,例如 '32,48,64'") + parser.add_argument("--boll-enter-list", type=str, default=None, help="Bollinger 入场 Z 值列表,例如 '1.0,1.3'") + parser.add_argument("--boll-exit-list", type=str, default=None, help="Bollinger 退出 Z 值列表,例如 '0.2,0.4'") + parser.add_argument("--boll-allow-short", type=str, default=None, help="Bollinger 是否允许做空,示例 'true,false'") + return parser.parse_args() + + +if __name__ == "__main__": + args = parse_args() + cfg_path = Path(args.config) + base_cfg = load_base_config(cfg_path) + if args.symbol: + base_cfg["symbol"] = args.symbol + if args.csv: + base_cfg["csv"] = args.csv + + suffix = args.suffix or base_cfg.get("symbol") + + grid_override = deepcopy(PARAM_GRID) + + def _parse_float_list(raw: Optional[str]) -> Optional[List[Optional[float]]]: + if raw is None: + return None + values: List[Optional[float]] = [] + for token in raw.split(","): + token = token.strip() + if not token: + continue + if token.lower() in {"none", "null"}: + values.append(None) + else: + values.append(float(token)) + return values or None + + atr_sl_list = _parse_float_list(args.atr_sl) + atr_tp_list = _parse_float_list(args.atr_tp) + cooldown_list = None + if args.cooldown_list: + cooldown_list = [int(item.strip()) for item in args.cooldown_list.split(",") if item.strip()] + htf_factor_list = None + if args.htf_factor_list: + htf_factor_list = [int(item.strip()) for item in args.htf_factor_list.split(",") if item.strip()] + size_mode_list = None + if args.size_tier_mode_list: + size_mode_list = [item.strip() for item in args.size_tier_mode_list.split(",") if item.strip()] + boll_window_list = None + if args.boll_window_list: + boll_window_list = [int(item.strip()) for item in args.boll_window_list.split(",") if item.strip()] + boll_enter_list = _parse_float_list(args.boll_enter_list) + boll_exit_list = _parse_float_list(args.boll_exit_list) + boll_allow_short_list = None + if args.boll_allow_short: + mapping = {"true": True, "false": False, "1": True, "0": False} + boll_allow_short_list = [ + mapping.get(item.strip().lower(), item.strip().lower() in {"true", "1"}) for item in args.boll_allow_short.split(",") if item.strip() + ] + + if atr_sl_list: + grid_override["atr_sl"] = atr_sl_list + if atr_tp_list: + grid_override["atr_tp"] = atr_tp_list + if cooldown_list: + grid_override["cooldown"] = cooldown_list + if htf_factor_list: + grid_override["htf_factor"] = htf_factor_list + if size_mode_list: + grid_override["size_tier_mode"] = size_mode_list + if boll_window_list: + grid_override["boll_window"] = boll_window_list + if boll_enter_list: + grid_override["boll_enter_z"] = boll_enter_list + if boll_exit_list: + grid_override["boll_exit_z"] = boll_exit_list + if boll_allow_short_list: + grid_override["boll_allow_short"] = boll_allow_short_list + + out_csv = run_grid(base_cfg, suffix, grid_override) + print(f"结果已写入: {out_csv}") diff --git a/Q Research/scripts/ingest_oanda.py b/Q Research/scripts/ingest_oanda.py new file mode 100644 index 0000000..39cae3c --- /dev/null +++ b/Q Research/scripts/ingest_oanda.py @@ -0,0 +1,218 @@ +#!/usr/bin/env python3 +""" +Automated OANDA ingestion with retries, logging, and manifest refresh. + +- Accepts direct CLI arguments (symbol/granularity/days/target count/output). +- Or accepts a schedule YAML describing multiple jobs: + - symbol: EUR_USD + - granularity: H1 + - days: 365 + - target_count: 8000 + - output: data/raw/EURUSD_H1.csv (optional; defaults to auto path) + +Each ingestion: + * retries on failure with exponential backoff + * appends a JSON line to metrics/ingest.log + * updates metrics/ingest_status.json with latest state per (symbol, granularity) + * triggers build_dataset_manifest.py to refresh hashes. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import subprocess +import sys +import time +from datetime import datetime, timezone +from pathlib import Path +from typing import Dict, List, Optional + +import yaml +from loguru import logger + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +sys.path.append(str(PROJECT_ROOT)) + +from scripts.get_candles import get_candles, default_out_csv, normalize_to_oanda # noqa + +METRICS_DIR = PROJECT_ROOT / "metrics" +METRICS_DIR.mkdir(exist_ok=True) +INGEST_LOG = METRICS_DIR / "ingest.log" +INGEST_STATUS = METRICS_DIR / "ingest_status.json" +INGEST_METRICS = METRICS_DIR / "ingest_metrics.csv" + +MANIFEST_SCRIPT = PROJECT_ROOT / "scripts" / "build_dataset_manifest.py" + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Ingest OANDA candles into data/raw.") + parser.add_argument("--symbol", help="Instrument symbol (e.g., EUR_USD)") + parser.add_argument("--granularity", default="H1", help="OANDA granularity (default H1)") + parser.add_argument("--days", type=int, default=365, help="Lookback days for ingestion") + parser.add_argument("--target-count", type=int, default=None, help="Optional target bar count") + parser.add_argument("--output", help="Explicit CSV output path") + parser.add_argument("--retries", type=int, default=3, help="Number of retries on failure") + parser.add_argument("--backoff", type=float, default=15.0, help="Base backoff seconds between retries") + parser.add_argument("--schedule", help="YAML file listing ingestion jobs") + parser.add_argument("--manifest-output", default="data/_manifest.json", help="Manifest path to refresh") + return parser.parse_args() + + +def load_schedule(path: str) -> List[Dict]: + schedule_path = PROJECT_ROOT / path + if not schedule_path.exists(): + raise FileNotFoundError(f"Schedule file not found: {schedule_path}") + with schedule_path.open("r", encoding="utf-8") as fh: + data = yaml.safe_load(fh) or [] + if not isinstance(data, list): + raise ValueError("Schedule YAML must be a list of jobs.") + return data + + +def resolve_output(symbol: str, granularity: str, explicit: Optional[str]) -> Path: + if explicit: + return (PROJECT_ROOT / explicit).resolve() + norm = normalize_to_oanda(symbol) + return Path(default_out_csv(str(PROJECT_ROOT), norm, granularity)) + + +def log_ingest(entry: Dict) -> None: + line = json.dumps(entry, ensure_ascii=False) + with INGEST_LOG.open("a", encoding="utf-8") as fh: + fh.write(line + "\n") + + existing = {} + if INGEST_STATUS.exists(): + try: + existing = json.load(INGEST_STATUS.open("r", encoding="utf-8")) or {} + except Exception: + existing = {} + key = f"{entry['symbol']}::{entry['granularity']}" + existing[key] = entry + with INGEST_STATUS.open("w", encoding="utf-8") as fh: + json.dump(existing, fh, indent=2, ensure_ascii=False) + + append_ingest_metric(entry) + + +def append_ingest_metric(entry: Dict) -> None: + headers = ["timestamp", "symbol", "granularity", "status", "rows", "duration_sec"] + INGEST_METRICS.parent.mkdir(exist_ok=True) + write_header = not INGEST_METRICS.exists() + with INGEST_METRICS.open("a", newline="", encoding="utf-8") as fh: + writer = csv.writer(fh) + if write_header: + writer.writerow(headers) + writer.writerow([ + entry.get("timestamp"), + entry.get("symbol"), + entry.get("granularity"), + entry.get("status"), + entry.get("rows"), + entry.get("duration_sec"), + ]) + + +def run_job(job: Dict, retries: int, backoff: float) -> bool: + symbol = job["symbol"] + granularity = job.get("granularity", "H1") + days = int(job.get("days", 365)) + target_count = job.get("target_count") + target_count = int(target_count) if target_count else None + output_path = resolve_output(symbol, granularity, job.get("output")) + output_path.parent.mkdir(parents=True, exist_ok=True) + + attempt = 0 + start_time = time.time() + errors = [] + while attempt <= retries: + attempt += 1 + try: + logger.info(f"[INGEST] {symbol} {granularity} attempt {attempt}/{retries+1} (days={days})") + df = get_candles(symbol=symbol, granularity=granularity, start_days_ago=days, target_count=target_count) + if df.empty: + raise RuntimeError("Fetched DataFrame is empty.") + df.to_csv(output_path, index=False) + duration = time.time() - start_time + entry = { + "timestamp": datetime.now(timezone.utc).isoformat(), + "symbol": symbol, + "granularity": granularity, + "rows": int(len(df)), + "output": str(output_path.relative_to(PROJECT_ROOT)), + "status": "success", + "duration_sec": round(duration, 2), + } + log_ingest(entry) + logger.info(f"[INGEST] Success {symbol} {granularity}: {len(df)} rows -> {output_path}") + return True + except Exception as exc: + err_msg = f"{type(exc).__name__}: {exc}" + errors.append(err_msg) + logger.error(f"[INGEST] Failed {symbol} {granularity} attempt {attempt}: {err_msg}") + if attempt <= retries: + sleep_time = backoff * attempt + logger.info(f"[INGEST] Sleeping {sleep_time:.1f}s before retry...") + time.sleep(sleep_time) + + duration = time.time() - start_time + entry = { + "timestamp": datetime.now(timezone.utc).isoformat(), + "symbol": symbol, + "granularity": granularity, + "rows": None, + "output": str(output_path.relative_to(PROJECT_ROOT)), + "status": "failure", + "duration_sec": round(duration, 2), + "errors": errors[-retries:], + } + log_ingest(entry) + return False + + +def refresh_manifest(manifest_output: str) -> None: + logger.info("[MANIFEST] Refreshing dataset manifest…") + cmd = [ + sys.executable, + str(MANIFEST_SCRIPT), + "--dirs", + "data/raw", + "data/derived", + "--output", + manifest_output, + ] + subprocess.run(cmd, cwd=PROJECT_ROOT, check=True) + + +def main() -> None: + args = parse_args() + jobs: List[Dict] + if args.schedule: + schedule_jobs = load_schedule(args.schedule) + jobs = schedule_jobs + elif args.symbol: + jobs = [{ + "symbol": args.symbol, + "granularity": args.granularity, + "days": args.days, + "target_count": args.target_count, + "output": args.output, + }] + else: + raise ValueError("Either --symbol or --schedule must be provided.") + + successes = 0 + for job in jobs: + if run_job(job, args.retries, args.backoff): + successes += 1 + + if successes > 0: + refresh_manifest(args.manifest_output) + else: + logger.warning("No successful ingestions; manifest not updated.") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/notify_risk_metrics.sh b/Q Research/scripts/notify_risk_metrics.sh new file mode 100644 index 0000000..6201ed4 --- /dev/null +++ b/Q Research/scripts/notify_risk_metrics.sh @@ -0,0 +1,25 @@ +#!/usr/bin/env bash +set -euo pipefail +if [ -z "${SLACK_RISK_WEBHOOK:-}" ]; then + echo "SLACK_RISK_WEBHOOK not set" >&2 + exit 1 +fi +cd "$(dirname "$0")/.." +set +e +output=$(python scripts/watch_ops_metrics.py --csv results/risk/metrics.csv 2>&1) +status=$? +set -e +export OUTPUT="$output" +if [ $status -ne 0 ]; then + payload=$(python - <<'PY' +import json, os +msg = os.environ['OUTPUT'] +print(json.dumps({"text": f"[Risk Alert] {msg}"})) +PY +) + curl -s -X POST -H 'Content-type: application/json' --data "$payload" "$SLACK_RISK_WEBHOOK" + echo "$output" + exit 1 +else + echo "$output" +fi diff --git a/Q Research/scripts/patch_missing_summaries.py b/Q Research/scripts/patch_missing_summaries.py new file mode 100644 index 0000000..18dbf1e --- /dev/null +++ b/Q Research/scripts/patch_missing_summaries.py @@ -0,0 +1,119 @@ +#!/usr/bin/env python3 +""" +Generate minimal summary.json and placeholder trades for runs that lack metadata. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import yaml +from datetime import datetime, timezone +from pathlib import Path +from typing import Dict, List + +ROOT = Path(__file__).resolve().parents[1] +RESULTS_DIR = ROOT / "results" +TRADES_DIR = ROOT / "data" / "outputs" / "trades" + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Patch runs missing summary.json artifacts.") + parser.add_argument("--runs", required=True, help="Comma-separated run IDs (e.g. 20251106_142210,20251106_142422)") + parser.add_argument("--symbol", default="EURUSD", help="Fallback symbol when unknown.") + parser.add_argument("--csv-path", default=None, help="Optional CSV path override (defaults to data/raw/_H1.csv).") + return parser.parse_args() + + +def ensure_placeholder_trades(run_id: str, symbol: str) -> str: + TRADES_DIR.mkdir(parents=True, exist_ok=True) + path = TRADES_DIR / f"trades_PLACEHOLDER_{run_id}.csv" + if path.exists(): + return str(path.as_posix()) + timestamp = datetime.strptime(run_id, "%Y%m%d_%H%M%S").replace(tzinfo=timezone.utc) + row = [ + timestamp.isoformat(), + timestamp.isoformat(), + symbol, + "long", + "0", + "1.0", + "1.0", + "0.0", + "default", + ] + with path.open("w", newline="", encoding="utf-8") as fh: + writer = csv.writer(fh) + writer.writerow(["ts_entry", "ts_exit", "symbol", "direction", "qty", "price_entry", "exit", "pnl", "strategy"]) + writer.writerow(row) + return str(path.as_posix()) + + +def load_performance(run_dir: Path) -> Dict[str, float]: + perf_path = run_dir / "performance.yml" + if not perf_path.exists(): + return {} + data = yaml.safe_load(perf_path.read_text(encoding="utf-8")) or {} + mapping = { + "annualized_return": "ann_return", + "volatility": "ann_vol", + "sharpe_ratio": "sharpe", + "max_drawdown": "max_drawdown", + "total_return": "total_return", + } + metrics: Dict[str, float] = {} + for src, dest in mapping.items(): + value = data.get(src) + if value is not None: + metrics[dest] = value + return metrics + + +def write_summary(run_id: str, symbol: str, csv_path: str) -> None: + run_dir = RESULTS_DIR / run_id + run_dir.mkdir(parents=True, exist_ok=True) + summary_path = run_dir / "summary.json" + metrics_path = run_dir / "metrics.json" + equity_path = run_dir / "equity_curve.csv" + trades_path = ensure_placeholder_trades(run_id, symbol) + + perf_metrics = load_performance(run_dir) + timestamp = datetime.strptime(run_id, "%Y%m%d_%H%M%S").replace(tzinfo=timezone.utc).isoformat() + + summary = { + "run_id": run_id, + "timestamp": timestamp, + "symbol": symbol, + "csv_path": csv_path, + "parameters": {"note": "Placeholder summary generated via patch_missing_summaries.py"}, + "metrics": perf_metrics, + "data_report": {"severity": "unknown", "messages": []}, + "artifacts": { + "equity": str(equity_path.relative_to(ROOT)) if equity_path.exists() else "", + "trades": trades_path, + "trade_stats": "", + }, + } + + summary_path.write_text(json.dumps(summary, indent=2), encoding="utf-8") + metrics_path.write_text(json.dumps(perf_metrics, indent=2), encoding="utf-8") + print(f"Patched summary for {run_id}") + + +def main() -> None: + args = parse_args() + runs: List[str] = [item.strip() for item in args.runs.split(",") if item.strip()] + if not runs: + raise SystemExit("No runs provided.") + csv_path = args.csv_path or f"data/raw/{args.symbol}_H1.csv" + for run_id in runs: + summary_file = RESULTS_DIR / run_id / "summary.json" + if summary_file.exists(): + print(f"{run_id}: summary already exists, skipping.") + continue + write_summary(run_id, args.symbol, csv_path) + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/plot_backtest_diagnostics.py b/Q Research/scripts/plot_backtest_diagnostics.py new file mode 100644 index 0000000..4096c6d --- /dev/null +++ b/Q Research/scripts/plot_backtest_diagnostics.py @@ -0,0 +1,282 @@ +#!/usr/bin/env python3 +"""Generate Phase 2 diagnostics charts (batch heatmap, Monte Carlo box, walk-forward timeline, equity plots).""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path +from typing import Dict, Optional, Tuple + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +from loguru import logger + +DEFAULT_OUT = Path("charts") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Plot diagnostics from batch/Monte Carlo/walk-forward outputs.") + parser.add_argument("--batch-csv", help="Path to batch_backtests_*.csv") + parser.add_argument("--walkforward-csv", help="Path to walkforward/metrics.csv") + parser.add_argument("--mc-summary", help="Path to stress/mc_summary.json") + parser.add_argument("--mc-iterations", help="Path to stress/mc_iterations.csv") + parser.add_argument("--equity-csv", help="Path to equity curve CSV (ts,equity)") + parser.add_argument("--underwater-csv", help="Path to underwater/drawdown CSV (ts,drawdown)") + parser.add_argument("--out", default=str(DEFAULT_OUT), help="Output directory (default charts/)") + parser.add_argument("--x-col", default="fast_win", help="Batch heatmap X-axis column") + parser.add_argument("--y-col", default="slow_win", help="Batch heatmap Y-axis column") + parser.add_argument("--metric", default="sharpe", help="Batch heatmap metric column") + parser.add_argument( + "--facet-scenario", + action="store_true", + help="If set, creates one heatmap per scenario column when available.", + ) + parser.add_argument("--format", default="png", choices=["png", "pdf"], help="Image format") + return parser.parse_args() + + +def ensure_dir(path: Path) -> Path: + path.mkdir(parents=True, exist_ok=True) + return path + + +def plot_heatmap( + batch_csv: Path, + out_dir: Path, + x_col: str, + y_col: str, + metric: str, + fmt: str, + facet: bool, +) -> Tuple[Dict[str, Optional[str]], Dict[str, Dict]]: + outputs: Dict[str, Optional[str]] = {} + data: Dict[str, Dict] = {} + if not batch_csv: + return outputs, data + df = pd.read_csv(batch_csv) + if not {x_col, y_col, metric}.issubset(df.columns): + logger.warning("Batch CSV missing columns required for heatmap (%s, %s, %s)", x_col, y_col, metric) + return outputs, data + + def _render(sub_df: pd.DataFrame, suffix: str) -> Optional[Path]: + pivot = sub_df.pivot_table(index=y_col, columns=x_col, values=metric, aggfunc="mean") + if pivot.empty: + return None + fig, ax = plt.subplots(figsize=(6, 4)) + c = ax.imshow(pivot.values, origin="lower", aspect="auto", cmap="viridis") + ax.set_xticks(np.arange(len(pivot.columns))) + ax.set_yticks(np.arange(len(pivot.index))) + ax.set_xticklabels(pivot.columns) + ax.set_yticklabels(pivot.index) + ax.set_xlabel(x_col) + ax.set_ylabel(y_col) + ax.set_title(f"{metric.title()} heatmap ({suffix})") + fig.colorbar(c, ax=ax, label=metric) + for i in range(len(pivot.index)): + for j in range(len(pivot.columns)): + value = pivot.values[i, j] + if not np.isnan(value): + ax.text(j, i, f"{value:.2f}", ha="center", va="center", color="white", fontsize=8) + filename = f"heatmap_{metric}_{suffix}.{fmt}" if suffix else f"heatmap_{metric}.{fmt}" + out_path = out_dir / filename + fig.tight_layout() + fig.savefig(out_path) + plt.close(fig) + data[suffix or "all"] = pivot.to_dict() + return out_path + + if facet and "scenario" in df.columns: + for scenario, sub in df.groupby("scenario"): + outputs[scenario] = str(_render(sub, scenario) or "") + else: + outputs["all"] = str(_render(df, "all") or "") + return outputs, data + + +def plot_mc_box(itr_csv: Path, summary_json: Path, out_dir: Path, fmt: str) -> Tuple[Optional[str], Dict[str, float]]: + if not itr_csv or not summary_json: + return None, {} + if not itr_csv.exists() or not summary_json.exists(): + logger.warning("Monte Carlo files missing: %s or %s", itr_csv, summary_json) + return None, {} + df = pd.read_csv(itr_csv) + if "sharpe" not in df.columns: + logger.warning("Monte Carlo iterations missing 'sharpe' column; skipping box plot") + return None, {} + series = df["sharpe"].dropna() + summary = json.loads(summary_json.read_text(encoding="utf-8")) + scenario = summary.get("scenario", "unknown") + fig, ax = plt.subplots(figsize=(5, 4)) + ax.boxplot(series, labels=[scenario]) + ax.set_ylabel("Sharpe") + ax.set_title("Monte Carlo Sharpe distribution") + out_path = out_dir / f"monte_carlo_box.{fmt}" + fig.tight_layout() + fig.savefig(out_path) + plt.close(fig) + logger.info("Saved Monte Carlo box plot to %s", out_path) + stats = { + "mean": float(series.mean()) if not series.empty else None, + "p05": float(series.quantile(0.05)) if not series.empty else None, + "p50": float(series.quantile(0.5)) if not series.empty else None, + "p95": float(series.quantile(0.95)) if not series.empty else None, + "count": int(series.size), + } + return str(out_path), stats + + +def plot_walkforward(wf_csv: Path, out_dir: Path, fmt: str) -> Tuple[Optional[str], list]: + if not wf_csv: + return None, [] + if not wf_csv.exists(): + logger.warning("Walk-forward metrics file missing: %s", wf_csv) + return None, [] + df = pd.read_csv(wf_csv) + if "window" not in df.columns or "sharpe" not in df.columns: + logger.warning("Walk-forward metrics missing 'window' or 'sharpe'; skipping timeline") + return None, [] + status = df.get("status", pd.Series(["unknown"] * len(df))) + colors = status.map({"pass": "#2ca02c", "fail": "#d62728"}).fillna("#1f77b4") + fig, ax = plt.subplots(figsize=(6, 3)) + ax.scatter(df["window"], df["sharpe"], c=colors) + ax.plot(df["window"], df["sharpe"], color="#cccccc", linewidth=1, alpha=0.5) + ax.set_xlabel("Window") + ax.set_ylabel("Sharpe") + ax.set_title("Walk-forward Sharpe timeline") + out_path = out_dir / f"walkforward_timeline.{fmt}" + fig.tight_layout() + fig.savefig(out_path) + plt.close(fig) + logger.info("Saved walk-forward timeline to %s", out_path) + return str(out_path), df[["window", "sharpe", *([ "status"] if "status" in df.columns else [])]].to_dict(orient="records") + + +def _resolve_column(df: pd.DataFrame, candidates) -> Optional[str]: + for col in candidates: + if col in df.columns: + return col + return None + + +def plot_equity(equity_csv: Path, out_dir: Path, fmt: str) -> Optional[str]: + if not equity_csv: + return None + if not equity_csv.exists(): + logger.warning("Equity CSV missing: %s", equity_csv) + return None + df = pd.read_csv(equity_csv) + ts_col = _resolve_column(df, ["ts", "timestamp", "time", "datetime"]) + equity_col = _resolve_column(df, ["equity", "capital", "balance"]) + if not ts_col or not equity_col: + logger.warning( + "Equity CSV missing timestamp/equity columns (looked for %s / %s)", + ["ts", "timestamp", "time", "datetime"], + ["equity", "capital", "balance"], + ) + return None + df[ts_col] = pd.to_datetime(df[ts_col]) + fig, ax = plt.subplots(figsize=(6, 3)) + ax.plot(df[ts_col], df[equity_col], color="#1f77b4") + ax.set_xlabel("Time") + ax.set_ylabel("Equity") + ax.set_title("Equity Curve") + fig.autofmt_xdate() + out_path = out_dir / f"equity_curve.{fmt}" + fig.tight_layout() + fig.savefig(out_path) + plt.close(fig) + logger.info("Saved equity curve to %s", out_path) + return str(out_path) + + +def plot_underwater(underwater_csv: Path, out_dir: Path, fmt: str) -> Optional[str]: + if not underwater_csv: + return None + if not underwater_csv.exists(): + logger.warning("Underwater CSV missing: %s", underwater_csv) + return None + df = pd.read_csv(underwater_csv) + required = {"ts", "drawdown"} + if not required.issubset(df.columns): + logger.warning("Underwater CSV missing columns %s", required) + return None + df["ts"] = pd.to_datetime(df["ts"]) + fig, ax = plt.subplots(figsize=(6, 3)) + ax.fill_between(df["ts"], df["drawdown"], color="#d62728", alpha=0.6) + ax.set_xlabel("Time") + ax.set_ylabel("Drawdown") + ax.set_title("Underwater Curve") + fig.autofmt_xdate() + out_path = out_dir / f"underwater_curve.{fmt}" + fig.tight_layout() + fig.savefig(out_path) + plt.close(fig) + logger.info("Saved underwater curve to %s", out_path) + return str(out_path) + + +def main() -> None: + args = parse_args() + out_dir = ensure_dir(Path(args.out).expanduser()) / Path( + f"diagnostics_{pd.Timestamp.now().strftime('%Y%m%d_%H%M%S')}" + ) + ensure_dir(out_dir) + + artifacts: Dict[str, Optional[str]] = {} + data_dump: Dict[str, Dict] = {} + + if args.batch_csv: + path = Path(args.batch_csv).expanduser() + heatmap_paths, heatmap_data = plot_heatmap( + path, out_dir, args.x_col, args.y_col, args.metric, args.format, args.facet_scenario + ) + artifacts["batch_heatmap"] = heatmap_paths.get("all") + if args.facet_scenario and "scenario" in heatmap_paths: + artifacts["batch_heatmap_scenarios"] = heatmap_paths + data_dump["batch_heatmap"] = heatmap_data + if args.mc_iterations and args.mc_summary: + itr = Path(args.mc_iterations).expanduser() + summary = Path(args.mc_summary).expanduser() + box_path, mc_stats = plot_mc_box(itr, summary, out_dir, args.format) + artifacts["mc_boxplot"] = box_path + data_dump["monte_carlo"] = mc_stats + if args.walkforward_csv: + wf = Path(args.walkforward_csv).expanduser() + timeline_path, wf_records = plot_walkforward(wf, out_dir, args.format) + artifacts["walkforward_timeline"] = timeline_path + data_dump["walkforward"] = wf_records + if args.equity_csv: + equity_path = plot_equity(Path(args.equity_csv).expanduser(), out_dir, args.format) + artifacts["equity_curve"] = equity_path + if args.underwater_csv: + underwater_path = plot_underwater(Path(args.underwater_csv).expanduser(), out_dir, args.format) + artifacts["underwater_curve"] = underwater_path + + meta = { + "generated_at": pd.Timestamp.now(tz="UTC").isoformat(), + "inputs": { + "batch_csv": args.batch_csv, + "walkforward_csv": args.walkforward_csv, + "mc_summary": args.mc_summary, + "mc_iterations": args.mc_iterations, + "equity_csv": args.equity_csv, + "underwater_csv": args.underwater_csv, + }, + "artifacts": artifacts, + } + meta_path = out_dir / "diagnostics_metadata.json" + meta_path.write_text(json.dumps(meta, indent=2), encoding="utf-8") + logger.info("Diagnostics metadata written to %s", meta_path) + + data_path = out_dir / "diagnostics_data.json" + data_path.write_text(json.dumps(data_dump, indent=2, default=str), encoding="utf-8") + logger.info("Diagnostics data written to %s", data_path) + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/plot_candles.py b/Q Research/scripts/plot_candles.py new file mode 100644 index 0000000..4108914 --- /dev/null +++ b/Q Research/scripts/plot_candles.py @@ -0,0 +1,36 @@ +# scripts/plot_candles.py + +import os +import pandas as pd +import plotly.graph_objects as go + +BASE_DIR = os.path.dirname(os.path.dirname(__file__)) +RAW_DATA_DIR = os.path.join(BASE_DIR, "data", "raw") + +def plot_csv_candles(file_path, symbol, granularity): + df = pd.read_csv(file_path, parse_dates=["time"]) + + fig = go.Figure(data=[ + go.Candlestick( + x=df["time"], + open=df["open"], + high=df["high"], + low=df["low"], + close=df["close"] + ) + ]) + fig.update_layout( + title=f"{symbol} {granularity} Candlestick Chart", + xaxis_title="Time", + yaxis_title="Price", + xaxis_rangeslider_visible=False + ) + + # 创建 charts 文件夹(如果不存在) + os.makedirs("charts", exist_ok=True) + output_file = f"charts/{symbol.replace('/', '')}_{granularity}.html" + fig.write_html(output_file) + print(f"✅ Saved chart to {output_file}") + +if __name__ == "__main__": + plot_csv_candles(os.path.join(RAW_DATA_DIR, "EURUSD_H1.csv"), "EUR/USD", "H1") diff --git a/Q Research/scripts/plot_indicators.py b/Q Research/scripts/plot_indicators.py new file mode 100644 index 0000000..9554cc7 --- /dev/null +++ b/Q Research/scripts/plot_indicators.py @@ -0,0 +1,138 @@ +# scripts/plot_indicators.py +import os +from datetime import datetime + +import numpy as np +import pandas as pd +import plotly.graph_objects as go +from plotly.subplots import make_subplots + +BASE_DIR = os.path.dirname(os.path.dirname(__file__)) +DERIVED_DATA_DIR = os.path.join(BASE_DIR, "data", "derived") +CHARTS_DIR = os.path.join(BASE_DIR, "charts") +os.makedirs(CHARTS_DIR, exist_ok=True) + +DATA_PATH = os.path.join(DERIVED_DATA_DIR, "EURUSD_H1_with_indicators.csv") + +# 自动生成带时间戳的输出文件名 +timestamp = datetime.now().strftime("%Y%m%d_%H%M") +OUTPUT_PATH = os.path.join(CHARTS_DIR, f"EURUSD_H1_with_indicators_{timestamp}.html") + +df = pd.read_csv(DATA_PATH, parse_dates=["time"]) + +# 计算均线(若已存在将覆盖为最新计算) +if {"close"}.issubset(df.columns): + df["SMA_20"] = df["close"].rolling(20).mean() + df["SMA_200"] = df["close"].rolling(200).mean() + +# 创建两行子图,共享x轴 +fig = make_subplots(rows=2, cols=1, shared_xaxes=True, + vertical_spacing=0.1, + row_heights=[0.7, 0.3], + specs=[[{"type": "candlestick"}], + [{"secondary_y": True}]]) + +# 第一个子图:蜡烛图 + 均线 +fig.add_trace(go.Candlestick( + x=df['time'], + open=df['open'], + high=df['high'], + low=df['low'], + close=df['close'], + name='Candlestick' +), row=1, col=1) + +if 'SMA_20' in df.columns: + fig.add_trace(go.Scatter( + x=df['time'], y=df['SMA_20'], + line=dict(color='blue', width=1), + name='SMA 20' + ), row=1, col=1) + +if 'EMA_50' in df.columns: + fig.add_trace(go.Scatter( + x=df['time'], y=df['EMA_50'], + line=dict(color='orange', width=1), + name='EMA 50' + ), row=1, col=1) + +# 绘制 SMA 200 +if 'SMA_200' in df.columns: + fig.add_trace(go.Scatter( + x=df['time'], y=df['SMA_200'], + line=dict(width=1), + name='SMA 200' + ), row=1, col=1) + +# 金叉/死叉标记:SMA20 与 SMA200 交叉 +if 'SMA_20' in df.columns and 'SMA_200' in df.columns: + sign = np.sign(df["SMA_20"] - df["SMA_200"]) + cross = sign.diff().fillna(0).ne(0) & df["SMA_20"].notna() & df["SMA_200"].notna() + golden = cross & (sign > 0) + dead = cross & (sign < 0) + + # 金叉 + fig.add_trace(go.Scatter( + x=df.loc[golden, "time"], y=df.loc[golden, "SMA_20"], + mode="markers", name="Golden Cross", + marker_symbol="triangle-up", marker_size=9 + ), row=1, col=1) + + # 死叉 + fig.add_trace(go.Scatter( + x=df.loc[dead, "time"], y=df.loc[dead, "SMA_20"], + mode="markers", name="Dead Cross", + marker_symbol="triangle-down", marker_size=9 + ), row=1, col=1) + +# 第二个子图:RSI 和 MACD +rsi_exists = 'RSI' in df.columns +macd_exists = 'MACD' in df.columns and 'MACD_signal' in df.columns and 'MACD_hist' in df.columns + +if rsi_exists: + fig.add_trace(go.Scatter( + x=df['time'], y=df['RSI'], + line=dict(color='purple', width=1), + name='RSI' + ), row=2, col=1, secondary_y=False) + +if macd_exists: + # MACD 柱状图 + fig.add_trace(go.Bar( + x=df['time'], y=df['MACD_hist'], + marker_color='grey', + name='MACD Hist' + ), row=2, col=1, secondary_y=True) + # MACD 线 + fig.add_trace(go.Scatter( + x=df['time'], y=df['MACD'], + line=dict(color='blue', width=1), + name='MACD' + ), row=2, col=1, secondary_y=True) + # MACD 信号线 + fig.add_trace(go.Scatter( + x=df['time'], y=df['MACD_signal'], + line=dict(color='orange', width=1, dash='dot'), + name='MACD Signal' + ), row=2, col=1, secondary_y=True) + +# 布局设置 +fig.update_layout( + title="EUR/USD with Technical Indicators", + xaxis_title="Time", + yaxis_title="Price", + xaxis_rangeslider_visible=False, + legend=dict(orientation="h", yanchor="bottom", y=1.02, xanchor="right", x=1) +) + +# RSI y轴范围限制 +if rsi_exists: + fig.update_yaxes(title_text="RSI", row=2, col=1, secondary_y=False, range=[0, 100]) + +# MACD y轴标题 +if macd_exists: + fig.update_yaxes(title_text="MACD", row=2, col=1, secondary_y=True) + +# 保存图表 +fig.write_html(OUTPUT_PATH) +print(f"✅ 图表已保存至 {OUTPUT_PATH}") diff --git a/Q Research/scripts/plot_risk_metrics.py b/Q Research/scripts/plot_risk_metrics.py new file mode 100644 index 0000000..ef60c0d --- /dev/null +++ b/Q Research/scripts/plot_risk_metrics.py @@ -0,0 +1,63 @@ +#!/usr/bin/env python3 +""" +Quick visualization/report for results/risk/metrics.csv. + +Usage: + python scripts/plot_risk_metrics.py --csv results/risk/metrics.csv --out charts/risk_metrics.png +""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +import matplotlib.pyplot as plt +import pandas as pd + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Plot risk metrics history (reject counts, status).") + parser.add_argument("--csv", default="results/risk/metrics.csv", help="Path to metrics CSV.") + parser.add_argument("--out", default="charts/risk_metrics.png", help="Output image path.") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + csv_path = Path(args.csv) + if not csv_path.exists(): + raise SystemExit(f"metrics CSV not found: {csv_path}") + df = pd.read_csv(csv_path) + if df.empty: + raise SystemExit("metrics CSV is empty.") + df["timestamp"] = pd.to_datetime(df["timestamp"]) + df.sort_values("timestamp", inplace=True) + df["status"] = df["status"].fillna("unknown").str.lower() + + fig, ax1 = plt.subplots(figsize=(10, 4)) + ax1.plot(df["timestamp"], df["rejects"], marker="o", label="Rejects") + ax1.set_ylabel("Reject count") + ax1.set_xlabel("Timestamp") + ax1.set_title("Risk simulation rejects over time") + + fail_mask = df["status"] == "fail" + ax1.scatter(df.loc[fail_mask, "timestamp"], df.loc[fail_mask, "rejects"], color="red", label="Fail", zorder=5) + ax1.legend(loc="upper left") + + ax2 = ax1.twinx() + status_numeric = df["status"].map({"pass": 1, "fail": 0}).fillna(0.5) + ax2.plot(df["timestamp"], status_numeric, color="gray", alpha=0.3, label="Status (1=pass,0=fail)") + ax2.set_ylim(-0.1, 1.1) + ax2.set_yticks([0, 0.5, 1]) + ax2.set_yticklabels(["fail", "unknown", "pass"]) + + fig.tight_layout() + out_path = Path(args.out) + out_path.parent.mkdir(parents=True, exist_ok=True) + fig.savefig(out_path) + plt.close(fig) + print(f"Saved risk metrics plot to {out_path}") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/risk_report.py b/Q Research/scripts/risk_report.py new file mode 100644 index 0000000..f30f4fd --- /dev/null +++ b/Q Research/scripts/risk_report.py @@ -0,0 +1,107 @@ +#!/usr/bin/env python3 +"""Summarize risk events logged by simulate_execution or live adapters.""" + +from __future__ import annotations + +import argparse +import json +from collections import Counter +from datetime import datetime, timezone +import os +from pathlib import Path + +import pandas as pd + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Aggregate risk event logs") + parser.add_argument("--log", default="results/risk/events.jsonl", help="Path to risk events JSONL") + parser.add_argument("--out", default="results/risk/report.csv", help="CSV summary output") + parser.add_argument("--run-id", help="Run identifier for metrics logging (defaults to $RUN)") + parser.add_argument("--metrics-path", default="results/risk/metrics.csv", help="Aggregate metrics CSV") + parser.add_argument( + "--status", + default="unknown", + help="Execution status recorded in metrics (e.g., pending/pass/fail).", + ) + parser.add_argument("--skip-report", action="store_true", help="Do not write the detailed CSV report.") + parser.add_argument("--skip-metrics", action="store_true", help="Do not append to metrics CSV.") + return parser.parse_args() + + +def load_events(path: Path) -> list[dict]: + if not path.exists(): + return [] + events = [] + with path.open("r", encoding="utf-8") as fh: + for line in fh: + try: + events.append(json.loads(line)) + except json.JSONDecodeError: + continue + return events + + +def main() -> None: + args = parse_args() + path = Path(args.log) + events = load_events(path) + if events: + df = pd.DataFrame(events) + counts = Counter(df["event"].fillna("unknown")) + print("Event counts:") + for event, count in counts.items(): + print(f" {event}: {count}") + else: + print(f"No events found in {path}") + df = pd.DataFrame(columns=["event", "ts", "symbol", "strategy", "reason"]) + counts = Counter() + if not args.skip_report: + df.to_csv(args.out, index=False) + if events: + print(f"Detailed report saved to {args.out}") + else: + print(f"Empty report written to {args.out}") + + run_id = args.run_id or os.environ.get("RUN") + if not run_id: + print("Run ID not provided; skipping metrics append.") + return + if args.skip_metrics: + print("skip-metrics enabled; metrics append skipped.") + return + reject_count = counts.get("reject", 0) + kill_count = counts.get("kill_switch", 0) + exec_dir = Path(f"results/execution/{run_id}") + append_metrics(Path(args.metrics_path), run_id, reject_count, kill_count, status=args.status, exec_dir=exec_dir) + + +def append_metrics(path: Path, run_id: str, rejects: int, kills: int, status: str = "unknown", exec_dir: Path | None = None) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + is_new = not path.exists() + timestamp = datetime.now(timezone.utc).isoformat() + latency_avg = latency_p95 = total_pnl = max_drawdown = 0.0 + max_gross = max_symbol = 0.0 + if exec_dir and (exec_dir / "fills.csv").exists(): + df = pd.read_csv(exec_dir / "fills.csv") + if "adapter_latency_ms" in df.columns and not df["adapter_latency_ms"].dropna().empty: + series = df["adapter_latency_ms"].fillna(0.0) + latency_avg = float(series.mean()) + latency_p95 = float(series.quantile(0.95)) + if "pnl" in df.columns: + total_pnl = float(df["pnl"].sum()) + if exec_dir and (exec_dir / "sim_results.json").exists(): + data = json.loads((exec_dir / "sim_results.json").read_text(encoding="utf-8")) + max_gross = float(data.get("max_gross_notional", 0.0)) + symbol_peaks = data.get("max_symbol_exposure", {}) + if isinstance(symbol_peaks, dict) and symbol_peaks: + max_symbol = float(max(abs(v) for v in symbol_peaks.values())) + max_drawdown = float(data.get("max_drawdown_pct", 0.0)) + with path.open("a", encoding="utf-8") as fh: + if is_new: + fh.write("timestamp,run_id,rejects,kills,status,latency_ms_avg,latency_ms_p95,total_pnl,max_gross_notional,max_symbol_exposure,max_drawdown_pct\n") + fh.write(f"{timestamp},{run_id},{rejects},{kills},{status},{latency_avg},{latency_p95},{total_pnl},{max_gross},{max_symbol},{max_drawdown}\n") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/run_batch_backtests.py b/Q Research/scripts/run_batch_backtests.py new file mode 100644 index 0000000..4b49ac6 --- /dev/null +++ b/Q Research/scripts/run_batch_backtests.py @@ -0,0 +1,268 @@ +from __future__ import annotations + +"""Batch backtests for multiple configs/symbols.""" + +import argparse +import json +import sys +from datetime import datetime +from pathlib import Path +from typing import Any, Dict, List, Optional, Sequence + +import pandas as pd +import yaml +from loguru import logger + +sys.path.append(str(Path(__file__).resolve().parents[1])) +from core.backtest.strategy_engine import parse_strategy_specs # type: ignore +from scripts.backtest_strategy import run_once # type: ignore +from scripts.scenario_utils import load_scenarios + + +KEY_MAP = { + "csv": "csv_path", + "cash": "initial_cash", + "qty": "qty", + "account_ccy": "account_ccy", + "fast": "fast_win", + "slow": "slow_win", + "spread": "spread_pips", + "slip": "slippage_pips", + "comm": "commission_per_million", + "sl": "stop_loss_pips", + "tp": "take_profit_pips", + "atr_sl": "atr_sl", + "atr_tp": "atr_tp", + "atr_window": "atr_window", + "rsi_period": "rsi_period", + "rsi_long_thresh": "rsi_long_thresh", + "rsi_short_thresh": "rsi_short_thresh", + "enable_trailing": "enable_trailing", + "trailing_enable_atr_mult": "trailing_enable_atr_mult", + "trailing_atr_mult": "trailing_atr_mult", + "long_only_above_slow": "long_only_above_slow", + "slope_lookback": "slope_lookback", + "cooldown": "cooldown", + "allow_short": "allow_short", + "short_only_below_slow": "short_only_below_slow", + "risk_per_trade_pct": "risk_per_trade_pct", + "max_drawdown_pct": "max_drawdown_pct", + "max_position_units": "max_position_units", + "regime_ema_window": "regime_ema_window", + "regime_slope_min": "regime_slope_min", + "regime_atr_min": "regime_atr_min", + "strategies": "strategies", + "htf_factor": "htf_factor", + "htf_ema_window": "htf_ema_window", + "htf_rsi_period": "htf_rsi_period", + "cost_profiles": "cost_profiles", + "slippage_model": "slippage_model", + "strategy_mode": "strategy_mode", + "strategy_vote_threshold": "strategy_vote_threshold", + "stress_cost_spread_mult": "stress_cost_spread_mult", + "stress_cost_comm_mult": "stress_cost_comm_mult", + "stress_slippage_mult": "stress_slippage_mult", + "stress_price_vol_mult": "stress_price_vol_mult", + "stress_skip_trade_pct": "stress_skip_trade_pct", +} + +SCENARIO_FIELDS = [ + "stress_cost_spread_mult", + "stress_cost_comm_mult", + "stress_slippage_mult", + "stress_price_vol_mult", + "stress_skip_trade_pct", +] + +PARAM_COLUMNS = [ + "fast_win", + "slow_win", + "atr_sl", + "atr_tp", + "atr_window", + "cooldown", + "allow_short", + "long_only_above_slow", + "short_only_below_slow", + "slope_lookback", + "strategy_vote_threshold", + "spread_pips", + "slippage_pips", + "commission_per_million", + "risk_per_trade_pct", +] + + +def normalize_params(raw: Dict[str, Any]) -> Dict[str, Any]: + params: Dict[str, Any] = {} + for k, v in (raw or {}).items(): + key = KEY_MAP.get(k, k) + params[key] = v + return params + + +def load_params(cfg_path: Path) -> Dict[str, Any]: + with cfg_path.open("r", encoding="utf-8") as fh: + raw = yaml.safe_load(fh) or {} + params = normalize_params(raw) + params["symbol"] = params.get("symbol", cfg_path.stem.upper()) + params["config_name"] = cfg_path.name + return params + + +def run_job( + params: Dict[str, Any], + scenario_name: Optional[str], + scenario_cfg: Optional[Dict[str, Any]], + extra_cols: Optional[Sequence[str]] = None, +) -> Dict[str, Any]: + label = params.get("label") + config_name = params.get("config_name") + job_params = {k: v for k, v in params.items() if k not in ("label", "config_name")} + if isinstance(job_params.get("strategies"), (list, dict)): + job_params["strategies"] = parse_strategy_specs(job_params["strategies"]) + symbol = job_params.get("symbol", "UNKNOWN").upper() + logger.info(f"Running {symbol} ({label or config_name})") + result = run_once(**job_params) + data_validation = result.get("data_validation") or {} + summary = { + "symbol": symbol, + "label": label, + "config_name": config_name, + "sharpe": result.get("sharpe"), + "ann_return": result.get("ann_return"), + "ann_vol": result.get("ann_vol"), + "max_drawdown": result.get("max_drawdown"), + "trades": result.get("trades"), + "final_equity": result.get("final_equity"), + "sortino": result.get("sortino"), + "calmar": result.get("calmar"), + "run_id": result.get("run_id"), + "summary_path": result.get("summary_path"), + "data_severity": data_validation.get("severity"), + "strategy_mode": job_params.get("strategy_mode"), + "stress_cost_spread_mult": job_params.get("stress_cost_spread_mult"), + "stress_cost_comm_mult": job_params.get("stress_cost_comm_mult"), + "stress_slippage_mult": job_params.get("stress_slippage_mult"), + "stress_price_vol_mult": job_params.get("stress_price_vol_mult"), + "stress_skip_trade_pct": job_params.get("stress_skip_trade_pct"), + "scenario": scenario_name, + "scenario_overrides": scenario_cfg, + } + for col in PARAM_COLUMNS: + value = job_params.get(col) + if isinstance(value, (int, float, bool, str)) or value is None: + summary[col] = value + if extra_cols: + for col in extra_cols: + if col in summary: + continue + if col in job_params: + summary[col] = job_params.get(col) + elif col in result: + summary[col] = result.get(col) + return summary + + +def load_jobs(symbols: List[str], config_dir: Path, schedule_path: Optional[str]) -> List[Dict[str, Any]]: + jobs: List[Dict[str, Any]] = [] + if schedule_path: + schedule_file = Path(schedule_path).expanduser() + data = yaml.safe_load(schedule_file.open("r", encoding="utf-8")) or [] + if not isinstance(data, list): + raise ValueError("Schedule YAML must be a list of job definitions") + for idx, raw in enumerate(data): + job = normalize_params(dict(raw or {})) + job["label"] = job.get("label") or f"job_{idx}" + if "symbol" not in job: + raise ValueError(f"Job {job['label']} missing 'symbol'") + jobs.append(job) + else: + for sym in symbols: + cfg_path = config_dir / f"{sym.lower()}_regime.yaml" + if not cfg_path.exists(): + logger.error(f"Config not found for {sym}: {cfg_path}") + continue + jobs.append(load_params(cfg_path)) + return jobs + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--symbols", type=str, default="EURUSD", help="Comma separated symbols (ignored if --schedule provided)") + parser.add_argument("--config-dir", type=str, default="config", help="Config directory") + parser.add_argument("--out", type=str, default="data/results", help="Output dir for summaries") + parser.add_argument("--schedule", type=str, help="YAML schedule of jobs") + parser.add_argument("--scenario", type=str, default=None, help="Default stress scenario for all jobs") + parser.add_argument( + "--scenario-file", + type=str, + default="config/stress_scenarios.yaml", + help="Scenario definition file (default: config/stress_scenarios.yaml)", + ) + parser.add_argument( + "--extra-cols", + type=str, + default="", + help="Comma-separated list of additional columns to include in the output (copied from params/result).", + ) + args = parser.parse_args() + + config_dir = Path(args.config_dir) + symbols = [s.strip().upper() for s in args.symbols.split(",")] + jobs = load_jobs(symbols, config_dir, args.schedule) + if not jobs: + logger.error("No jobs to run") + sys.exit(1) + + scenario_path = Path(args.scenario_file).expanduser() + scenario_required = bool(args.scenario) or any(job.get("scenario") for job in jobs) + scenarios = load_scenarios(scenario_path) if scenario_required else {} + + extra_cols = [c.strip() for c in (args.extra_cols or "").split(",") if c.strip()] + rows = [] + for job in jobs: + scenario_name = job.get("scenario") or args.scenario + scenario_cfg = None + if scenario_name: + if scenario_name not in scenarios: + raise ValueError(f"Scenario '{scenario_name}' not found in {scenario_path}") + scenario_cfg = scenarios[scenario_name] + job_copy = dict(job) + job_params = {k: v for k, v in job_copy.items() if k not in ("scenario",)} + if scenario_cfg: + for field in SCENARIO_FIELDS: + if field in scenario_cfg and field not in job_params: + job_params[field] = scenario_cfg[field] + try: + rows.append(run_job(job_params, scenario_name, scenario_cfg, extra_cols=extra_cols)) + except Exception as exc: + label = job.get("label") or job.get("symbol") + logger.exception(f"Backtest failed for job {label}: {exc}") + + if not rows: + logger.error("All backtests failed") + sys.exit(1) + + df = pd.DataFrame(rows) + out_dir = Path(args.out) + out_dir.mkdir(parents=True, exist_ok=True) + ts = datetime.now().strftime("%Y%m%d_%H%M%S") + csv_path = out_dir / f"batch_backtests_{ts}.csv" + df.to_csv(csv_path, index=False) + json_path = out_dir / f"batch_backtests_{ts}.json" + json_path.write_text(df.to_json(orient="records", indent=2), encoding="utf-8") + stats = { + "timestamp": ts, + "job_count": len(rows), + "symbols": sorted(df["symbol"].unique()), + "sharpe_max": df["sharpe"].max(), + "sharpe_min": df["sharpe"].min(), + } + (out_dir / f"batch_backtests_{ts}_stats.json").write_text(json.dumps(stats, indent=2), encoding="utf-8") + print(df.to_string(index=False)) + print(f"Saved summary to {csv_path} / {json_path}") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/run_ci_diagnostics.sh b/Q Research/scripts/run_ci_diagnostics.sh new file mode 100644 index 0000000..ce1f14c --- /dev/null +++ b/Q Research/scripts/run_ci_diagnostics.sh @@ -0,0 +1,35 @@ +#!/usr/bin/env bash +set -euo pipefail +cd "$(dirname "$0")/.." + +# Auto-select newest artifacts unless explicit paths supplied. +TARGET=${1:-latest} +if [ "$TARGET" = "latest" ]; then + BATCH=$(ls -t data/results/batch_backtests_*.csv 2>/dev/null | head -n1 || true) + WALKFORWARD=$(ls -t results/*/walkforward/metrics.csv 2>/dev/null | head -n1 || true) + MC_SUMMARY=$(ls -t results/*/stress/mc_summary.json 2>/dev/null | head -n1 || true) + MC_ITER=$(ls -t results/*/stress/mc_iterations.csv 2>/dev/null | head -n1 || true) + EQUITY=$(ls -t results/*/equity.csv 2>/dev/null | head -n1 || true) + UNDERWATER=$(ls -t results/*/stats/underwater.csv 2>/dev/null | head -n1 || true) +fi +BATCH=${BATCH:-tests/fixtures/batch_results_sample.csv} +WALKFORWARD=${WALKFORWARD:-tests/fixtures/walkforward_metrics_sample.csv} +MC_SUMMARY=${MC_SUMMARY:-tests/fixtures/mc_summary_sample.json} +MC_ITER=${MC_ITER:-tests/fixtures/mc_iterations_sample.csv} +EQUITY=${EQUITY:-tests/fixtures/equity_sample.csv} +UNDERWATER=${UNDERWATER:-tests/fixtures/underwater_sample.csv} + +OUT_DIR="charts/ci/diagnostics_$(date +%Y%m%d_%H%M%S)" + +python scripts/plot_backtest_diagnostics.py \ + --batch-csv "$BATCH" \ + --walkforward-csv "$WALKFORWARD" \ + --mc-summary "$MC_SUMMARY" \ + --mc-iterations "$MC_ITER" \ + --equity-csv "$EQUITY" \ + --underwater-csv "$UNDERWATER" \ + --facet-scenario \ + --out "$OUT_DIR" \ + --format png + +echo "Diagnostics artifacts saved to $OUT_DIR" diff --git a/Q Research/scripts/run_monte_carlo.py b/Q Research/scripts/run_monte_carlo.py new file mode 100644 index 0000000..2f1515b --- /dev/null +++ b/Q Research/scripts/run_monte_carlo.py @@ -0,0 +1,181 @@ +#!/usr/bin/env python3 +""" +Monte Carlo / stress test for an existing backtest run. +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path +from typing import List, Optional + +import numpy as np +import pandas as pd +from loguru import logger + +BASE_DIR = Path(__file__).resolve().parents[1] +if str(BASE_DIR) not in sys.path: + sys.path.insert(0, str(BASE_DIR)) + +from metrics.perf import compute_metrics +from scripts.scenario_utils import get_scenario + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Run Monte Carlo stress on a backtest result.") + parser.add_argument("--run", help="Path to results/ directory.") + parser.add_argument("--equity", help="Explicit equity CSV (ts,equity). Overrides --run artifacts.") + parser.add_argument("--iterations", type=int, default=500, help="Number of bootstrap iterations.") + parser.add_argument("--method", choices=["bootstrap", "block"], default="bootstrap", help="Resampling method.") + parser.add_argument("--block-size", type=int, default=None, help="Block size for block bootstrap (overrides scenario).") + parser.add_argument("--return-scale", type=float, default=None, help="Scale factor applied to resampled returns (overrides scenario).") + parser.add_argument("--ruin-threshold", type=float, default=0.8, help="Final equity / initial equity threshold to count as ruin.") + parser.add_argument("--seed", type=int, default=None, help="Random seed.") + parser.add_argument("--scenario", default=None, help="Optional stress scenario label stored in outputs.") + parser.add_argument( + "--scenario-file", + default="config/stress_scenarios.yaml", + help="Scenario definitions file (default: config/stress_scenarios.yaml).", + ) + return parser.parse_args() + + +def load_equity_series(run_path: Path | None, explicit_csv: str | None) -> pd.Series: + if explicit_csv: + path = Path(explicit_csv).expanduser() + else: + if not run_path: + raise ValueError("Either --run or --equity must be provided.") + summary_file = run_path / "summary.json" + if not summary_file.exists(): + raise FileNotFoundError(f"summary.json not found in {run_path}") + summary = json.load(summary_file.open("r", encoding="utf-8")) + artifacts = summary.get("artifacts") or {} + equity_path = artifacts.get("equity") + if not equity_path: + raise FileNotFoundError("Equity artifact missing in summary; rerun backtest after upgrading.") + path = (BASE_DIR / equity_path).resolve() + if not path.exists(): + raise FileNotFoundError(f"Equity file not found: {path}") + df = pd.read_csv(path) + if "equity" not in df.columns: + raise ValueError(f"Equity CSV missing 'equity' column: {path}") + return df["equity"].astype(float) + + +def block_bootstrap(returns: np.ndarray, size: int, block_size: int, rng: np.random.Generator) -> np.ndarray: + if returns.size == 0: + raise ValueError("Cannot run block bootstrap on an empty return series.") + if block_size <= 0: + raise ValueError("Block size must be a positive integer.") + if block_size > len(returns): + raise ValueError(f"Block size {block_size} exceeds return series length {len(returns)}.") + out = [] + while len(out) < size: + start = rng.integers(0, len(returns) - block_size + 1) + block = returns[start : start + block_size] + out.extend(block.tolist()) + return np.array(out[:size]) + + +def run_iteration(returns: np.ndarray, args, rng: np.random.Generator) -> dict: + if returns.size == 0: + raise ValueError("Return series is empty; cannot run Monte Carlo.") + if args.method == "bootstrap": + draw = rng.choice(returns, size=returns.size, replace=True) + else: + draw = block_bootstrap(returns, returns.size, args.block_size, rng) + draw = draw * args.return_scale + equity = np.cumprod(1.0 + draw) + equity_series = list(enumerate(equity, start=1)) + metrics = compute_metrics(equity_series) + return metrics + + +def summarize(metrics_list: List[dict], ruin_threshold: float, initial_equity: float) -> dict: + df = pd.DataFrame(metrics_list) + summary = {} + for column in ["sharpe", "sortino", "calmar", "ann_return", "max_drawdown"]: + if column in df.columns: + summary[column] = { + "mean": float(df[column].mean()), + "std": float(df[column].std()), + "p05": float(df[column].quantile(0.05)), + "p50": float(df[column].quantile(0.5)), + "p95": float(df[column].quantile(0.95)), + } + ruin = (df["final_equity"] <= initial_equity * ruin_threshold).mean() if "final_equity" in df else None + summary["p_ruin"] = float(ruin) if ruin is not None else None + return summary + + +def _load_scenario(args: argparse.Namespace) -> Optional[dict]: + if not args.scenario: + return None + scenario_path = Path(args.scenario_file).expanduser() + try: + scenario = get_scenario(args.scenario, scenario_path) + except KeyError as exc: + raise ValueError(f"Scenario '{args.scenario}' not found in {scenario_path}") from exc + return scenario + + +def _apply_scenario(args: argparse.Namespace, scenario_cfg: Optional[dict]) -> None: + if not scenario_cfg: + args.return_scale = args.return_scale if args.return_scale is not None else 1.0 + args.block_size = args.block_size if args.block_size is not None else 20 + return + if args.return_scale is None and scenario_cfg.get("return_scale") is not None: + args.return_scale = float(scenario_cfg["return_scale"]) + if args.block_size is None and scenario_cfg.get("block_size") is not None: + args.block_size = int(scenario_cfg["block_size"]) + args.return_scale = args.return_scale if args.return_scale is not None else scenario_cfg.get("return_scale", 1.0) + args.block_size = args.block_size if args.block_size is not None else scenario_cfg.get("block_size", 20) + + +def main(): + args = parse_args() + scenario_cfg = _load_scenario(args) + _apply_scenario(args, scenario_cfg) + run_path = Path(args.run).expanduser().resolve() if args.run else None + equity = load_equity_series(run_path, args.equity) + returns = np.diff(equity.values) / equity.values[:-1] + if returns.size == 0: + raise ValueError("Equity series must contain at least two points.") + initial_equity = float(equity.iloc[0]) + rng = np.random.default_rng(args.seed) + metrics_list: List[dict] = [] + + logger.info( + "Running Monte Carlo | method={method} iterations={iterations} scenario={scenario} seed={seed}", + method=args.method, + iterations=args.iterations, + scenario=args.scenario or "default", + seed=args.seed, + ) + for _ in range(args.iterations): + metrics = run_iteration(returns, args, rng) + metrics_list.append(metrics) + + summary = summarize(metrics_list, args.ruin_threshold, initial_equity) + summary["iterations"] = args.iterations + summary["method"] = args.method + summary["return_scale"] = args.return_scale + summary["scenario"] = args.scenario or "default" + summary["seed"] = args.seed + summary["scenario_overrides"] = scenario_cfg + + output_dir = run_path / "stress" if run_path else BASE_DIR / "results" / "stress" + output_dir.mkdir(parents=True, exist_ok=True) + iterations_csv = output_dir / "mc_iterations.csv" + pd.DataFrame(metrics_list).to_csv(iterations_csv, index=False) + summary_json = output_dir / "mc_summary.json" + with summary_json.open("w", encoding="utf-8") as fh: + json.dump(summary, fh, indent=2, ensure_ascii=False) + logger.info(f"Monte Carlo summary saved to {summary_json}") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/run_risk_sim.sh b/Q Research/scripts/run_risk_sim.sh new file mode 100644 index 0000000..75187c0 --- /dev/null +++ b/Q Research/scripts/run_risk_sim.sh @@ -0,0 +1,55 @@ +#!/usr/bin/env bash +set -euo pipefail +cd "$(dirname "$0")/.." + +RUN_ID=${RUN:-$(date +"risk_%Y%m%d_%H%M%S")} +RISK_LOG="results/risk/events.jsonl" + +# 每次运行前清理旧的风险事件,避免历史数据干扰统计 +mkdir -p "$(dirname "$RISK_LOG")" +: > "$RISK_LOG" + +# 自动发现 trades.csv(优先使用显式 TRADES,未设置则读取 summary.json 中的 artifacts.trades) +if [ -z "${TRADES:-}" ]; then + SUMMARY="results/${RUN_ID}/summary.json" + if [ -f "$SUMMARY" ]; then + TRADES=$(python - <<'PY' "$SUMMARY" +import json, sys +with open(sys.argv[1], encoding="utf-8") as fh: + summary = json.load(fh) +print(summary.get("artifacts", {}).get("trades", "")) +PY +) + fi +fi + +# 如果 summary 中给的是相对路径,补全为绝对路径 +if [ -n "${TRADES:-}" ] && [ ! -f "$TRADES" ] && [ -f "$PWD/$TRADES" ]; then + TRADES="$PWD/$TRADES" +fi + +if [ -z "${TRADES:-}" ] || [ ! -f "$TRADES" ]; then + echo "Trades CSV not found. Provide TRADES env or ensure results/${RUN_ID}/summary.json.artifacts.trades 指向有效文件。" >&2 + exit 1 +fi + +python scripts/simulate_execution.py --trades-csv "$TRADES" \ + --risk-limits-yaml ../QuantTrader/config/risk_limits_sim.yaml \ + --adapter paper \ + --paper-latency-ms 25 \ + --paper-slippage-pips 0.05 \ + --run-id "$RUN_ID" \ + --risk-log "$RISK_LOG" +python scripts/risk_report.py --log "$RISK_LOG" --out results/risk/report.csv --run-id "$RUN_ID" --skip-metrics +if python scripts/check_risk_report.py --report results/risk/report.csv --max-rejects 0 --max-kill 0; then + STATUS="pass" +else + STATUS="fail" +fi +python scripts/risk_report.py --log "$RISK_LOG" --out results/risk/report.csv --run-id "$RUN_ID" --status "$STATUS" --skip-report +if [ -s "$RISK_LOG" ]; then + cp "$RISK_LOG" "results/risk/events_${RUN_ID}.jsonl" +fi +if [ "$STATUS" != "pass" ]; then + exit 1 +fi diff --git a/Q Research/scripts/run_walkforward.py b/Q Research/scripts/run_walkforward.py new file mode 100644 index 0000000..58901e2 --- /dev/null +++ b/Q Research/scripts/run_walkforward.py @@ -0,0 +1,359 @@ +#!/usr/bin/env python3 +""" +Rolling walk-forward runner that slices a dataset into train/test windows +and executes the upgraded backtest pipeline for each slice. +""" + +from __future__ import annotations + +import argparse +import copy +import hashlib +import inspect +import json +import sys +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Dict, List, Optional, Tuple + +import pandas as pd +import yaml +from loguru import logger + +BASE_DIR = Path(__file__).resolve().parents[1] +if str(BASE_DIR) not in sys.path: + sys.path.insert(0, str(BASE_DIR)) + +from core.backtest.strategy_engine import parse_strategy_specs # type: ignore +from scripts import validate_dataset as dq # type: ignore +from scripts.backtest_strategy import run_once # type: ignore + +TIME_COLUMNS = ["ts", "time", "timestamp", "datetime", "date"] +DEFAULT_MANIFEST = "data/_manifest.json" +DEFAULT_RESULTS = "results" + +RUN_ONCE_PARAMS = set(inspect.signature(run_once).parameters.keys()) + +KEY_MAP = { + "csv": "csv_path", + "cash": "initial_cash", + "qty": "qty", + "account_ccy": "account_ccy", + "fast": "fast_win", + "slow": "slow_win", + "spread": "spread_pips", + "slip": "slippage_pips", + "comm": "commission_per_million", + "sl": "stop_loss_pips", + "tp": "take_profit_pips", + "atr_sl": "atr_sl", + "atr_tp": "atr_tp", + "atr_window": "atr_window", + "rsi_period": "rsi_period", + "rsi_long_thresh": "rsi_long_thresh", + "rsi_short_thresh": "rsi_short_thresh", + "enable_trailing": "enable_trailing", + "trailing_enable_atr_mult": "trailing_enable_atr_mult", + "trailing_atr_mult": "trailing_atr_mult", + "long_only_above_slow": "long_only_above_slow", + "slope_lookback": "slope_lookback", + "cooldown": "cooldown", + "allow_short": "allow_short", + "short_only_below_slow": "short_only_below_slow", + "risk_per_trade_pct": "risk_per_trade_pct", + "max_drawdown_pct": "max_drawdown_pct", + "max_position_units": "max_position_units", + "regime_ema_window": "regime_ema_window", + "regime_slope_min": "regime_slope_min", + "regime_atr_min": "regime_atr_min", + "regime_atr_percentile_min": "regime_atr_percentile_min", + "regime_atr_percentile_window": "regime_atr_percentile_window", + "regime_trend_min_bars": "regime_trend_min_bars", + "strategies": "strategies", + "htf_factor": "htf_factor", + "htf_ema_window": "htf_ema_window", + "htf_rsi_period": "htf_rsi_period", + "cost_profiles": "cost_profiles", + "slippage_model": "slippage_model", + "strategy_mode": "strategy_mode", +} + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Run walk-forward analysis across rolling windows.") + parser.add_argument("--config", required=True, help="YAML config with base strategy parameters.") + parser.add_argument("--csv", help="Override CSV path (defaults to config csv_path).") + parser.add_argument("--train-bars", type=int, default=3000, help="Number of bars in each training window.") + parser.add_argument("--test-bars", type=int, default=1000, help="Number of bars in each test window.") + parser.add_argument( + "--step-bars", + type=int, + default=None, + help="Step size between windows (defaults to test-bars).", + ) + parser.add_argument("--max-windows", type=int, default=None, help="Optional cap on number of windows.") + parser.add_argument( + "--output-root", + default=DEFAULT_RESULTS, + help="Directory where aggregated walk-forward artifacts will be stored.", + ) + parser.add_argument("--manifest", default=DEFAULT_MANIFEST, help="Manifest path for optional validation.") + parser.add_argument("--label", default=None, help="Optional label recorded in summary.json.") + parser.add_argument( + "--sharpe-threshold", + type=float, + default=1.0, + help="Minimum Sharpe to mark a window as pass.", + ) + parser.add_argument( + "--max-dd-threshold", + type=float, + default=0.1, + help="Maximum allowed drawdown magnitude (positive value).", + ) + parser.add_argument( + "--validate-base-data", + action="store_true", + help="Run data validation once on the source CSV before slicing.", + ) + parser.add_argument( + "--keep-train-csv", + action="store_true", + help="Export the train slices alongside test slices for auditing (default: only test).", + ) + return parser.parse_args() + + +def normalize_params(raw: Dict[str, Any]) -> Dict[str, Any]: + normalized: Dict[str, Any] = {} + for key, value in (raw or {}).items(): + canon_key = KEY_MAP.get(key, key) + if canon_key in RUN_ONCE_PARAMS: + normalized[canon_key] = value + return normalized + + +def load_config(cfg_path: Path) -> Dict[str, Any]: + if not cfg_path.exists(): + raise FileNotFoundError(f"Config not found: {cfg_path}") + with cfg_path.open("r", encoding="utf-8") as fh: + raw = yaml.safe_load(fh) or {} + cfg = normalize_params(raw) + strategies = cfg.get("strategies") + if strategies: + cfg["strategies"] = parse_strategy_specs(strategies) + return cfg + + +def detect_time_column(df: pd.DataFrame) -> str: + for col in TIME_COLUMNS: + if col in df.columns: + return col + for col in df.columns: + if pd.api.types.is_datetime64_any_dtype(df[col]): + return col + raise ValueError(f"No timestamp column found in dataset; expected any of {TIME_COLUMNS}") + + +def prepare_dataset(csv_path: Path) -> Tuple[pd.DataFrame, str]: + df = pd.read_csv(csv_path) + time_col = detect_time_column(df) + df["ts"] = pd.to_datetime(df[time_col], utc=True, errors="coerce") + df = df.dropna(subset=["ts"]).sort_values("ts").reset_index(drop=True) + return df, time_col + + +def compute_windows( + df: pd.DataFrame, + train: int, + test: int, + step: int, + max_windows: Optional[int] = None, +) -> List[Tuple[int, slice, slice]]: + total = len(df) + if train <= 0 or test <= 0: + raise ValueError("train-bars and test-bars must be positive.") + if train + test > total: + raise ValueError(f"Dataset length ({total}) insufficient for a single window of train+test={train + test}.") + windows: List[Tuple[int, slice, slice]] = [] + idx = 0 + win_id = 0 + while idx + train + test <= total: + train_slice = slice(idx, idx + train) + test_slice = slice(idx + train, idx + train + test) + windows.append((win_id, train_slice, test_slice)) + win_id += 1 + if max_windows is not None and win_id >= max_windows: + break + idx += step + return windows + + +def file_sha256(path: Path) -> str: + h = hashlib.sha256() + with path.open("rb") as fh: + for chunk in iter(lambda: fh.read(65536), b""): + h.update(chunk) + return h.hexdigest() + + +def params_fingerprint(params: Dict[str, Any]) -> str: + ignore = {"csv_path", "results_dir", "manifest_path", "validate_data", "write_summary"} + filtered = {k: v for k, v in params.items() if k not in ignore} + blob = json.dumps(filtered, sort_keys=True, default=str) + return hashlib.sha256(blob.encode("utf-8")).hexdigest() + + +def export_slice(df: pd.DataFrame, indices: slice, path: Path) -> None: + subset = df.iloc[indices] + subset.to_csv(path, index=False) + + +def maybe_validate_dataset(csv_path: Path, manifest: Path) -> Optional[Dict[str, Any]]: + if not csv_path.exists(): + raise FileNotFoundError(f"Dataset for validation not found: {csv_path}") + manifest_entry = dq.load_manifest_entry(Path(manifest), csv_path) if manifest else None + report = dq.compute_report(csv_path, manifest_entry, z_threshold=5.0) + severity = report.get("severity") + logger.info( + "Base dataset validation: severity=%s rows=%s gap_ratio=%.6f", + severity, + report.get("total_rows"), + report.get("gap_ratio", 0.0), + ) + if severity == "error": + raise RuntimeError(f"Dataset validation failed for {csv_path}: {report.get('messages')}") + return report + + +def build_summary(stats: List[Dict[str, Any]]) -> Dict[str, Any]: + df = pd.DataFrame(stats) + aggregates: Dict[str, Dict[str, float]] = {} + for metric in ["sharpe", "ann_return", "ann_vol", "max_drawdown", "sortino", "calmar"]: + if metric in df.columns and not df[metric].dropna().empty: + series = df[metric].dropna() + aggregates[metric] = { + "mean": float(series.mean()), + "median": float(series.median()), + "std": float(series.std(ddof=0)), + "p05": float(series.quantile(0.05)), + "p95": float(series.quantile(0.95)), + } + summary = { + "windows": len(stats), + "aggregates": aggregates, + "run_ids": [row.get("run_id") for row in stats], + "passes": int((df["status"] == "pass").sum()) if "status" in df.columns else None, + "fails": int((df["status"] == "fail").sum()) if "status" in df.columns else None, + } + return summary + + +def main(): + args = parse_args() + cfg_path = Path(args.config).expanduser().resolve() + cfg = load_config(cfg_path) + csv_path = Path(args.csv or cfg.get("csv_path") or cfg.get("csv", "")).expanduser() + if not csv_path: + raise ValueError("CSV path must be provided via --csv or config file.") + if not csv_path.exists(): + raise FileNotFoundError(f"CSV file not found: {csv_path}") + + df, _ = prepare_dataset(csv_path) + train = args.train_bars + test = args.test_bars + step = args.step_bars or test + windows = compute_windows(df, train, test, step, args.max_windows) + if not windows: + raise RuntimeError("No walk-forward windows could be generated with the provided parameters.") + + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + label = args.label or cfg_path.stem + session_dir = Path(args.output_root).expanduser().resolve() / f"walkforward_{label}_{timestamp}" + wf_dir = session_dir / "walkforward" + wf_dir.mkdir(parents=True, exist_ok=True) + + base_validation_report = None + if args.validate_base_data: + base_validation_report = maybe_validate_dataset(csv_path, Path(args.manifest).expanduser().resolve()) + + base_params = copy.deepcopy(cfg) + base_params["validate_data"] = False # slices derive from validated dataset + base_params.setdefault("symbol", label.upper()) + + stats: List[Dict[str, Any]] = [] + for win_id, train_slice, test_slice in windows: + test_csv_path = wf_dir / f"window_{win_id:03d}_test.csv" + export_slice(df, test_slice, test_csv_path) + if args.keep_train_csv: + train_csv_path = wf_dir / f"window_{win_id:03d}_train.csv" + export_slice(df, train_slice, train_csv_path) + + params = copy.deepcopy(base_params) + params["csv_path"] = str(test_csv_path) + params["results_dir"] = str(session_dir) + params["manifest_path"] = args.manifest + + logger.info( + "Walk-forward window {win} | train={train} bars test={test} bars ({start} → {end})", + win=win_id, + train=train, + test=test, + start=df.iloc[test_slice.start]["ts"], + end=df.iloc[test_slice.stop - 1]["ts"], + ) + + result = run_once(**params) + + record = { + "window": win_id, + "train_rows": train_slice.stop - train_slice.start, + "test_rows": test_slice.stop - test_slice.start, + "train_start": df.iloc[train_slice.start]["ts"].isoformat(), + "train_end": df.iloc[train_slice.stop - 1]["ts"].isoformat(), + "test_start": df.iloc[test_slice.start]["ts"].isoformat(), + "test_end": df.iloc[test_slice.stop - 1]["ts"].isoformat(), + "run_id": result.get("run_id"), + "summary_path": result.get("summary_path"), + "sharpe": result.get("sharpe"), + "ann_return": result.get("ann_return"), + "ann_vol": result.get("ann_vol"), + "max_drawdown": result.get("max_drawdown"), + "sortino": result.get("sortino"), + "calmar": result.get("calmar"), + "final_equity": result.get("final_equity"), + "trades": result.get("trades"), + "data_hash": file_sha256(test_csv_path), + "data_path": str(test_csv_path.relative_to(BASE_DIR)) if test_csv_path.is_relative_to(BASE_DIR) else str(test_csv_path), + "param_fingerprint": params_fingerprint(params), + } + max_dd = abs(record.get("max_drawdown") or 0.0) + sharpe = record.get("sharpe") or 0.0 + record["status"] = "pass" if sharpe >= args.sharpe_threshold and max_dd <= args.max_dd_threshold else "fail" + stats.append(record) + + metrics_csv = wf_dir / "metrics.csv" + pd.DataFrame(stats).to_csv(metrics_csv, index=False) + + summary = { + "label": label, + "session_dir": str(session_dir.relative_to(BASE_DIR)) if session_dir.is_relative_to(BASE_DIR) else str(session_dir), + "source_csv": str(csv_path.relative_to(BASE_DIR)) if csv_path.is_relative_to(BASE_DIR) else str(csv_path), + "train_bars": train, + "test_bars": test, + "step_bars": step, + "created_at": datetime.now(timezone.utc).isoformat(), + "sharpe_threshold": args.sharpe_threshold, + "max_dd_threshold": args.max_dd_threshold, + "base_validation": base_validation_report, + } + summary.update(build_summary(stats)) + summary_path = wf_dir / "summary.json" + summary_path.write_text(json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8") + + logger.info("Walk-forward run complete: %s", summary_path) + logger.info("Metrics CSV saved to %s", metrics_csv) + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/scenario_utils.py b/Q Research/scripts/scenario_utils.py new file mode 100644 index 0000000..6343ee6 --- /dev/null +++ b/Q Research/scripts/scenario_utils.py @@ -0,0 +1,99 @@ +"""Helpers for loading and validating stress scenarios.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any, Dict, List, Optional + +import yaml + +ALLOWED_FIELDS = { + "description", + "notes", + "stress_cost_spread_mult", + "stress_cost_comm_mult", + "stress_slippage_mult", + "stress_price_vol_mult", + "stress_skip_trade_pct", + "return_scale", + "block_size", +} + +NUMERIC_FIELDS = { + "stress_cost_spread_mult": (0.0, None), + "stress_cost_comm_mult": (0.0, None), + "stress_slippage_mult": (0.0, None), + "stress_price_vol_mult": (0.0, None), + "stress_skip_trade_pct": (0.0, 1.0), + "return_scale": (0.0, None), +} + +INT_FIELDS = { + "block_size": (1, None), +} + + +def _validate_mapping(name: str, config: Any) -> List[str]: + errors: List[str] = [] + if not isinstance(config, dict): + errors.append(f"Scenario '{name}' must be a mapping of field -> value.") + return errors + unknown = set(config.keys()) - ALLOWED_FIELDS + if unknown: + errors.append(f"Scenario '{name}' has unknown fields: {sorted(unknown)}") + for field, (min_val, max_val) in NUMERIC_FIELDS.items(): + if field in config and config[field] is not None: + try: + value = float(config[field]) + except (TypeError, ValueError): + errors.append(f"Scenario '{name}' field '{field}' must be numeric.") + continue + if min_val is not None and value < min_val: + errors.append(f"Scenario '{name}' field '{field}' must be >= {min_val}. Got {value}.") + if max_val is not None and value > max_val: + errors.append(f"Scenario '{name}' field '{field}' must be <= {max_val}. Got {value}.") + for field, (min_val, max_val) in INT_FIELDS.items(): + if field in config and config[field] is not None: + try: + value = int(config[field]) + except (TypeError, ValueError): + errors.append(f"Scenario '{name}' field '{field}' must be an integer.") + continue + if min_val is not None and value < min_val: + errors.append(f"Scenario '{name}' field '{field}' must be >= {min_val}. Got {value}.") + if max_val is not None and value > max_val: + errors.append(f"Scenario '{name}' field '{field}' must be <= {max_val}. Got {value}.") + return errors + + +def validate_data(data: Any) -> List[str]: + if not isinstance(data, dict): + return ["Scenario file must contain a mapping of scenario_name -> config"] + errors: List[str] = [] + for name, config in data.items(): + errors.extend(_validate_mapping(name, config)) + return errors + + +def load_scenarios(path: Path) -> Dict[str, Dict[str, Any]]: + if not path.exists(): + raise FileNotFoundError(f"Scenario file not found: {path}") + raw = yaml.safe_load(path.read_text(encoding="utf-8")) or {} + errors = validate_data(raw) + if errors: + raise ValueError("Invalid scenario definitions:\n" + "\n".join(errors)) + return {name: dict(config or {}) for name, config in raw.items()} + + +def get_scenario(name: str, path: Path) -> Dict[str, Any]: + scenarios = load_scenarios(path) + if name not in scenarios: + raise KeyError(f"Scenario '{name}' not found in {path}") + return scenarios[name] + + +def apply_defaults(values: Dict[str, Any], defaults: Dict[str, Any]) -> Dict[str, Any]: + merged = dict(defaults) + merged.update({k: v for k, v in values.items() if v is not None}) + return merged + diff --git a/Q Research/scripts/simulate_execution.py b/Q Research/scripts/simulate_execution.py new file mode 100644 index 0000000..129a805 --- /dev/null +++ b/Q Research/scripts/simulate_execution.py @@ -0,0 +1,277 @@ +#!/usr/bin/env python3 +"""Replay orders through ExecutionAdapter + RiskEngine with per-strategy configs.""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path +from typing import Dict, List, Optional + +import pandas as pd +import yaml + +BASE_DIR = Path(__file__).resolve().parents[1] +ROOT_DIR = BASE_DIR.parent +for path in (BASE_DIR, ROOT_DIR): + if str(path) not in sys.path: + sys.path.insert(0, str(path)) + +from QuantTrader.core.risk.risk_engine import RiskEngine, RiskLimits +from QuantTrader.execution.adapter import MockAdapter, OrderParams +from QuantTrader.execution.paper_adapter import PaperAdapter + + +def infer_symbol_from_path(path: Path) -> Optional[str]: + stem = path.stem + for token in stem.split("_"): + clean = "".join(ch for ch in token if ch.isalpha()) + if not clean: + continue + if clean.lower() in {"trade", "trades"}: + continue + if 3 <= len(clean) <= 10: + return clean.upper() + return None + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Simulate execution with risk checks.") + parser.add_argument("--orders", help="CSV with columns ts,symbol,side,qty,price,notional,strategy(optional)") + parser.add_argument("--trades-csv", help="Optional trades.csv (ts_entry,symbol,direction,qty,price_entry,pnl,strategy)") + parser.add_argument("--symbol", help="Fallback symbol when trades CSV lacks column") + parser.add_argument("--risk-config", help="JSON risk config (single engine)") + parser.add_argument("--risk-limits-yaml", help="YAML mapping strategies -> limits") + parser.add_argument("--run-id", help="Run identifier (defaults to timestamp)") + parser.add_argument("--output", default=None, help="Summary output path (auto from run-id if omitted)") + parser.add_argument("--risk-log", default="results/risk/events.jsonl", help="Risk event log path") + parser.add_argument("--adapter", choices=["mock", "paper"], default="mock", help="Execution adapter to use.") + parser.add_argument("--paper-latency-ms", type=float, default=50.0, help="Paper adapter latency (ms).") + parser.add_argument("--paper-slippage-pips", type=float, default=0.1, help="Paper adapter slippage (pips).") + return parser.parse_args() + + +def load_risk_engine(config_path: Path) -> RiskEngine: + cfg = json.loads(config_path.read_text(encoding="utf-8")) + limits = RiskLimits( + max_position_notional=cfg["max_position_notional"], + max_gross_leverage=cfg["max_gross_leverage"], + max_daily_loss=cfg["max_daily_loss"], + max_drawdown=cfg["max_drawdown"], + ) + return RiskEngine(limits=limits, starting_equity=cfg["starting_equity"]) + + +def load_strategy_engines(yaml_path: Path) -> Dict[str, RiskEngine]: + data = yaml.safe_load(yaml_path.read_text(encoding="utf-8")) or {} + engines: Dict[str, RiskEngine] = {} + global_cfg = data.get("global", {}) + global_limits = global_cfg.get("limits", {}) + default_engine = RiskEngine( + limits=RiskLimits( + max_position_notional=global_limits.get("max_position_notional", 0), + max_gross_leverage=global_limits.get("max_gross_leverage", 0), + max_daily_loss=global_limits.get("max_daily_loss", 0), + max_drawdown=global_limits.get("max_drawdown", 0), + ), + starting_equity=global_cfg.get("starting_equity", 0), + ) + engines["default"] = default_engine + for name, cfg in (data.get("strategies") or {}).items(): + limits_cfg = cfg.get("limits", {}) + engines[name] = RiskEngine( + limits=RiskLimits( + max_position_notional=limits_cfg.get("max_position_notional", global_limits.get("max_position_notional", 0)), + max_gross_leverage=limits_cfg.get("max_gross_leverage", global_limits.get("max_gross_leverage", 0)), + max_daily_loss=limits_cfg.get("max_daily_loss", global_limits.get("max_daily_loss", 0)), + max_drawdown=limits_cfg.get("max_drawdown", global_limits.get("max_drawdown", 0)), + ), + starting_equity=cfg.get("starting_equity", global_cfg.get("starting_equity", 0)), + ) + return engines + + +def append_event(path: Path, event: Dict) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("a", encoding="utf-8") as fh: + fh.write(json.dumps(event) + "\n") + + +def load_orders(args: argparse.Namespace) -> pd.DataFrame: + if args.trades_csv: + trades_path = Path(args.trades_csv) + trades = pd.read_csv(trades_path) + + if "direction" not in trades.columns and "side" in trades.columns: + trades["direction"] = trades["side"] + if "price_entry" not in trades.columns: + if "entry_price" in trades.columns: + trades["price_entry"] = trades["entry_price"] + elif "entry" in trades.columns: + trades["price_entry"] = trades["entry"] + + if "symbol" not in trades.columns: + inferred = infer_symbol_from_path(trades_path) + symbol = args.symbol or inferred + if symbol: + trades["symbol"] = symbol + else: + raise ValueError("trades_csv missing 'symbol' and no --symbol fallback provided") + + required = {"ts_entry", "ts_exit", "symbol", "direction", "qty", "price_entry"} + missing = required - set(trades.columns) + if missing: + raise ValueError(f"trades_csv missing columns: {missing}") + + trades["direction"] = trades["direction"].str.lower() + trades["price_entry"] = trades["price_entry"].astype(float) + trades["qty"] = trades["qty"].astype(float) + + records: List[Dict] = [] + for _, row in trades.iterrows(): + direction = str(row["direction"]).lower() + symbol = row["symbol"] + qty = float(row["qty"]) + price_entry = float(row["price_entry"]) + notional = qty * price_entry + strategy = row.get("strategy") if isinstance(row.get("strategy"), str) else "default" + pnl_value = float(row["pnl"]) if not pd.isna(row.get("pnl")) else 0.0 + + ts_entry = row.get("ts_entry") + ts_exit = row.get("ts_exit") + exit_price = row.get("exit") + + if pd.notna(ts_entry): + side = "buy" if direction == "long" else "sell" + records.append( + { + "ts": ts_entry, + "symbol": symbol, + "side": side, + "qty": qty, + "price": price_entry, + "notional": notional, + "pnl": 0.0, + "strategy": strategy, + } + ) + + if pd.notna(ts_exit): + side = "sell" if direction == "long" else "buy" + price_use = float(exit_price) if exit_price and not pd.isna(exit_price) else price_entry + records.append( + { + "ts": ts_exit, + "symbol": symbol, + "side": side, + "qty": qty, + "price": price_use, + "notional": notional, + "pnl": pnl_value, + "strategy": strategy, + } + ) + + if not records: + raise ValueError("No usable rows found in trades CSV.") + return pd.DataFrame.from_records(records) + if not args.orders: + raise ValueError("Either --orders or --trades-csv must be provided") + return pd.read_csv(Path(args.orders)) + + +def simulate() -> None: + args = parse_args() + orders_df = load_orders(args) + + if args.risk_limits_yaml: + engines = load_strategy_engines(Path(args.risk_limits_yaml)) + default_engine = engines.get("default") + elif args.risk_config: + engine = load_risk_engine(Path(args.risk_config)) + engines = {"default": engine} + default_engine = engine + else: + raise ValueError("Provide either --risk-config or --risk-limits-yaml") + + latency_ms = args.paper_latency_ms if args.adapter == "paper" else 0.0 + if args.adapter == "paper": + adapter = PaperAdapter(latency_ms=latency_ms, slippage_pips=args.paper_slippage_pips) + else: + adapter = MockAdapter() + + symbol_exposure_peaks: Dict[str, float] = {} + max_gross_notional = 0.0 + total_pnl_sum = 0.0 + run_id = args.run_id or pd.Timestamp.now().strftime("exec_%Y%m%d_%H%M%S") + run_dir = Path(f"results/execution/{run_id}") + run_dir.mkdir(parents=True, exist_ok=True) + risk_log_path = Path(args.risk_log) + + fills: List[Dict] = [] + rejects: List[Dict] = [] + kill_events: List[Dict] = [] + + for _, row in orders_df.iterrows(): + strategy = row.get("strategy", "default") + risk_engine = engines.get(strategy, default_engine) + ok, reason = risk_engine.evaluate_order(row["symbol"], row["side"], row["notional"]) + if not ok: + reject = {"ts": row["ts"], "symbol": row["symbol"], "strategy": strategy, "reason": reason} + rejects.append(reject) + append_event(risk_log_path, {"event": "reject", **reject}) + continue + ack = adapter.submit( + OrderParams( + symbol=row["symbol"], + side=row["side"], + quantity=row["qty"], + price=row.get("price"), + metadata={"strategy": strategy}, + ) + ) + pnl = row.get("pnl", 0.0) + risk_engine.record_fill(row["symbol"], row["side"], row["notional"], pnl) + exposures = risk_engine.state.exposures + current_gross = sum(abs(v) for v in exposures.values()) + max_gross_notional = max(max_gross_notional, current_gross) + for sym, val in exposures.items(): + peak = symbol_exposure_peaks.get(sym, 0.0) + symbol_exposure_peaks[sym] = max(peak, abs(val)) + total_pnl_sum += pnl + ok_loss, reason_loss = risk_engine.check_loss_limits() + if not ok_loss: + event = {"event": "kill_switch", "strategy": strategy, "reason": reason_loss} + kill_events.append(event) + append_event(risk_log_path, event) + lat = getattr(ack, "latency_ms", latency_ms) + fills.append({ + "order_id": ack.order_id, + "ts": row["ts"], + "symbol": row["symbol"], + "pnl": pnl, + "strategy": strategy, + "adapter_latency_ms": lat, + }) + + summary = { + "fills": fills, + "rejects": rejects, + "kill_switch_events": kill_events, + "run_id": run_id, + "max_symbol_exposure": symbol_exposure_peaks, + "max_gross_notional": max_gross_notional, + "total_pnl": total_pnl_sum, + "max_drawdown_pct": risk_engine.max_drawdown_pct(), + } + out_path = Path(args.output) if args.output else run_dir / "sim_results.json" + out_path.parent.mkdir(parents=True, exist_ok=True) + out_path.write_text(json.dumps(summary, indent=2), encoding="utf-8") + pd.DataFrame(fills).to_csv(run_dir / "fills.csv", index=False) + pd.DataFrame(rejects).to_csv(run_dir / "rejects.csv", index=False) + print(f"Simulation summary saved to {out_path} (run_id={run_id})") + + +if __name__ == "__main__": + simulate() diff --git a/Q Research/scripts/train_xgb_usdjpy.py b/Q Research/scripts/train_xgb_usdjpy.py new file mode 100644 index 0000000..33c570a --- /dev/null +++ b/Q Research/scripts/train_xgb_usdjpy.py @@ -0,0 +1,234 @@ +#!/usr/bin/env python3 +""" +Train an XGBoost classifier on USDJPY H1 and export model artifacts. + +Artifacts are written to: QuantResearch/artifacts/models/usdjpy_h1_xgb// + - model.json (xgboost Booster) + - feature_list.json (ordered feature names) + - thresholds.json (p_long, p_exit, val/test metrics) + - meta.json (dataset/costs/params/seed/etc.) + +Also updates: QuantResearch/artifacts/models/usdjpy_h1_xgb_latest.json + {"model_dir": "QuantResearch/artifacts/models/usdjpy_h1_xgb/"} +""" + +from __future__ import annotations + +import argparse +import json +import math +from datetime import datetime, timezone +from pathlib import Path +from typing import Dict, List, Tuple + +import numpy as np +import pandas as pd + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser(description="Train XGB (long-only) on USDJPY H1") + p.add_argument("--csv", default="QuantResearch/data/raw/USDJPY_H1.csv") + p.add_argument("--symbol", default="USDJPY") + p.add_argument("--horizon", type=int, default=6) + p.add_argument("--train-ratio", type=float, default=0.6) + p.add_argument("--val-ratio", type=float, default=0.2) + p.add_argument("--seed", type=int, default=42) + # Unified costs (match backtest/YAML) + p.add_argument("--spread-pips", type=float, default=2.0) + p.add_argument("--slip-pips", type=float, default=0.3) + p.add_argument("--comm-per-million", type=float, default=0.25) + # Model params (conservative defaults) + p.add_argument("--max-depth", type=int, default=4) + p.add_argument("--n-estimators", type=int, default=300) + p.add_argument("--learning-rate", type=float, default=0.05) + p.add_argument("--min-child-weight", type=float, default=1.0) + p.add_argument("--reg-lambda", type=float, default=1.0) + # Output base dir + p.add_argument("--out", default="QuantResearch/artifacts/models/usdjpy_h1_xgb") + return p.parse_args() + + +def _pip_value(symbol: str) -> float: + return 0.01 if symbol.upper().endswith("JPY") else 0.0001 + + +def compute_cost_return(symbol: str, price: pd.Series, spread_pips: float, slip_pips: float, comm_per_million: float) -> pd.Series: + pip = _pip_value(symbol) + # Approx trade cost (fractional): spread + 2*slippage in price terms + commission fraction + frac_px = (spread_pips + 2.0 * slip_pips) * pip / price + frac_comm = (comm_per_million / 1_000_000.0) + return frac_px + frac_comm + + +def rsi(series: pd.Series, period: int = 14) -> pd.Series: + delta = series.diff() + up = delta.clip(lower=0) + down = -delta.clip(upper=0) + roll_up = up.rolling(period).mean() + roll_down = down.rolling(period).mean() + rs = roll_up / roll_down + return 100 - (100 / (1 + rs)) + + +def atr(high: pd.Series, low: pd.Series, close: pd.Series, period: int = 14) -> pd.Series: + prev_close = close.shift(1) + tr = pd.concat([(high - low).abs(), (high - prev_close).abs(), (low - prev_close).abs()], axis=1).max(axis=1) + return tr.rolling(period).mean() + + +def add_time_features(df: pd.DataFrame) -> pd.DataFrame: + ts = pd.to_datetime(df["time"], utc=True) + hour = ts.dt.hour.astype(float) + dow = ts.dt.dayofweek.astype(float) + df["hour_sin"], df["hour_cos"] = np.sin(2 * np.pi * hour / 24.0), np.cos(2 * np.pi * hour / 24.0) + df["dow_sin"], df["dow_cos"] = np.sin(2 * np.pi * dow / 7.0), np.cos(2 * np.pi * dow / 7.0) + return df + + +def build_features(df: pd.DataFrame, fast: int = 20, slow: int = 80, rsi_p: int = 14, atr_p: int = 14) -> Tuple[pd.DataFrame, List[str]]: + out = df.copy() + out = add_time_features(out) + out["ret_1"] = out["close"].pct_change(1) + out["ret_3"] = out["close"].pct_change(3) + out["ret_6"] = out["close"].pct_change(6) + out["vol_24"] = out["close"].pct_change(1).rolling(24).std() + sma_f = out["close"].rolling(fast).mean() + sma_s = out["close"].rolling(slow).mean() + out["sma_diff"] = (sma_f - sma_s) / out["close"] + out["rsi"] = rsi(out["close"], rsi_p) + out["atr_norm"] = atr(out["high"], out["low"], out["close"], atr_p) / out["close"] + feats = [ + "ret_1", "ret_3", "ret_6", "vol_24", + "sma_diff", "rsi", "atr_norm", + "hour_sin", "hour_cos", "dow_sin", "dow_cos", + ] + return out, feats + + +def forward_return(close: pd.Series, horizon: int) -> pd.Series: + return close.shift(-horizon) / close - 1.0 + + +def pick_thresholds(p_val: np.ndarray, fwd_val: np.ndarray, cost_val: np.ndarray) -> Dict[str, float]: + best = {"thr_long": 0.6, "thr_exit": 0.5, "sharpe": -np.inf, "trades": 0} + for thr in np.linspace(0.55, 0.8, 11): + for thr_exit in np.linspace(0.45, 0.6, 16): + chosen = p_val >= thr + if not np.any(chosen): + continue + net = (fwd_val - cost_val)[chosen] + if net.size < 50: + continue + mu = float(np.nanmean(net)) + sd = float(np.nanstd(net, ddof=1)) + sharpe = (mu / sd) * math.sqrt(252 * 24) if sd > 0 else -np.inf + if sharpe > best["sharpe"]: + best = {"thr_long": float(thr), "thr_exit": float(thr_exit), "sharpe": sharpe, "trades": int(net.size)} + return best + + +def main() -> None: + try: + import xgboost as xgb # requires xgboost==1.7.6 per requirements + except Exception as exc: + raise SystemExit("xgboost is required. Please install xgboost==1.7.6.") from exc + + args = parse_args() + np.random.seed(args.seed) + + df = pd.read_csv(args.csv) + if "time" not in df.columns: + raise SystemExit("CSV must include 'time' column") + df = df[["time", "open", "high", "low", "close", "volume"]].copy() + + df_feat, feat_list = build_features(df) + fwd = forward_return(df_feat["close"], args.horizon) + costs = compute_cost_return(args.symbol, df_feat["close"], args.spread_pips, args.slip_pips, args.comm_per_million) + # Label with cost margin + y = pd.Series(np.where(fwd > costs, 1, np.where(fwd < -costs, 0, np.nan)), index=df_feat.index) + + data = df_feat.assign(y=y, cost=costs).dropna(subset=feat_list + ["y"]).reset_index(drop=True) + X = data[feat_list].to_numpy() + y_arr = data["y"].astype(int).to_numpy() + fwd_arr = fwd.loc[data.index].to_numpy() + cost_arr = data["cost"].to_numpy() + + n = len(data) + i_train = int(n * args.train_ratio) + i_val = int(n * (args.train_ratio + args.val_ratio)) + if i_val >= n: + i_val = n - max(1, n // 10) + + X_train, y_train = X[:i_train], y_arr[:i_train] + X_val, y_val = X[i_train:i_val], y_arr[i_train:i_val] + X_test, y_test = X[i_val:], y_arr[i_val:] + fwd_val, cost_val = fwd_arr[i_train:i_val], cost_arr[i_train:i_val] + fwd_test, cost_test = fwd_arr[i_val:], cost_arr[i_val:] + + dtrain = xgb.DMatrix(X_train, label=y_train, feature_names=feat_list) + dval = xgb.DMatrix(X_val, label=y_val, feature_names=feat_list) + dtest = xgb.DMatrix(X_test, label=y_test, feature_names=feat_list) + + params = { + "objective": "binary:logistic", + "eval_metric": "logloss", + "max_depth": args.max_depth, + "eta": args.learning_rate, + "lambda": args.reg_lambda, + "min_child_weight": args.min_child_weight, + "subsample": 0.9, + "colsample_bytree": 0.9, + "seed": args.seed, + } + evals = [(dtrain, "train"), (dval, "val")] + booster = xgb.train(params, dtrain, num_boost_round=args.n_estimators, evals=evals, early_stopping_rounds=50, verbose_eval=False) + + p_val = booster.predict(dval) + th = pick_thresholds(p_val, fwd_val, cost_val) + p_test = booster.predict(dtest) + chosen = p_test >= th["thr_long"] + net = (fwd_test - cost_test)[chosen] + test_trades = int(net.size) + mu = float(np.nanmean(net)) if net.size else 0.0 + sd = float(np.nanstd(net, ddof=1)) if net.size > 1 else 0.0 + test_sharpe = (mu / sd) * math.sqrt(252 * 24) if sd > 0 else 0.0 + + ts = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S") + model_dir = Path(args.out).with_suffix("") / ts + model_dir.mkdir(parents=True, exist_ok=True) + # Save artifacts + booster.save_model(str(model_dir / "model.json")) + (model_dir / "feature_list.json").write_text(json.dumps(feat_list, indent=2), encoding="utf-8") + thresholds = { + "p_long": float(th["thr_long"]), + "p_exit": float(th["thr_exit"]), + "val_sharpe": float(th["sharpe"]), + "val_trades": int(th["trades"]), + "test_sharpe": float(test_sharpe), + "test_trades": int(test_trades), + } + (model_dir / "thresholds.json").write_text(json.dumps(thresholds, indent=2), encoding="utf-8") + meta = { + "generated_at": datetime.now(timezone.utc).isoformat(), + "symbol": args.symbol, + "csv": args.csv, + "horizon": args.horizon, + "splits": {"train": int(i_train), "val": int(i_val - i_train), "test": int(n - i_val)}, + "seed": int(args.seed), + "features": feat_list, + "xgb_params": params, + "best_iteration": int(getattr(booster, "best_iteration", 0) or 0), + "thresholds": thresholds, + "costs": {"spread_pips": args.spread_pips, "slip_pips": args.slip_pips, "comm_per_million": args.comm_per_million}, + "git_commit": None, + } + (model_dir / "meta.json").write_text(json.dumps(meta, indent=2), encoding="utf-8") + + latest_path = Path(args.out).with_suffix("").parent / "usdjpy_h1_xgb_latest.json" + latest_path.parent.mkdir(parents=True, exist_ok=True) + latest_path.write_text(json.dumps({"model_dir": str(model_dir)}, indent=2), encoding="utf-8") + print(f"Saved model artifacts to {model_dir}") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/update_metrics_from_tca.py b/Q Research/scripts/update_metrics_from_tca.py new file mode 100644 index 0000000..ec2d19e --- /dev/null +++ b/Q Research/scripts/update_metrics_from_tca.py @@ -0,0 +1,90 @@ +#!/usr/bin/env python3 +""" +Append a metrics row using TCA summary + KPI overrides. + +Example: + python scripts/update_metrics_from_tca.py \ + --tca QuantTrader/results/execution/tca_summary.json \ + --run-id 20251112_parallel_demo_live \ + --status pass --latency-avg 28 --latency-p95 45 \ + --total-pnl 27.3 --max-exposure 500000 --max-drawdown 0.02 \ + --rolling-sharpe 1.5 --live-drawdown 0.03 \ + --live-latency-p95 45 --slippage-bps 1.2 +""" + +from __future__ import annotations + +import argparse +import csv +import json +from datetime import datetime, timezone +from pathlib import Path + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Append metrics row from TCA summary.") + parser.add_argument("--tca", required=True, help="Path to tca_summary.json") + parser.add_argument("--metrics", default="results/risk/metrics.csv", help="Metrics CSV path") + parser.add_argument("--run-id", required=True) + parser.add_argument("--status", default="pass") + parser.add_argument("--latency-avg", type=float, required=True) + parser.add_argument("--latency-p95", type=float, required=True) + parser.add_argument("--total-pnl", type=float, required=True) + parser.add_argument("--max-exposure", type=float, required=True) + parser.add_argument("--max-drawdown", type=float, required=True) + parser.add_argument("--rolling-sharpe", type=float, required=True) + parser.add_argument("--live-drawdown", type=float, required=True) + parser.add_argument("--live-latency-p95", type=float, required=True) + parser.add_argument("--slippage-bps", type=float, required=True) + parser.add_argument("--rejects", type=int, default=0) + parser.add_argument("--kills", type=int, default=0) + parser.add_argument("--max-gross-notional", type=float, default=0.0) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + tca_path = Path(args.tca) + if not tca_path.exists(): + raise SystemExit(f"TCA summary not found: {tca_path}") + with tca_path.open("r", encoding="utf-8") as f: + tca = json.load(f) + + timestamp = datetime.now(timezone.utc).isoformat() + row = { + "timestamp": timestamp, + "run_id": args.run_id, + "rejects": args.rejects, + "kills": args.kills, + "status": args.status, + "latency_ms_avg": args.latency_avg, + "latency_ms_p95": args.latency_p95, + "total_pnl": args.total_pnl, + "max_gross_notional": args.max_gross_notional or args.max_exposure, + "max_symbol_exposure": args.max_exposure, + "max_drawdown_pct": args.max_drawdown, + "rolling_sharpe_30d": args.rolling_sharpe, + "live_drawdown_pct": args.live_drawdown, + "live_latency_ms_p95": args.live_latency_p95, + "slippage_bps": args.slippage_bps, + "paper_trade_count": tca.get("paper_trade_count"), + "paper_total_pnl": tca.get("paper_total_pnl"), + "live_trade_count": tca.get("live_trade_count"), + "live_total_pnl": tca.get("live_total_pnl"), + "pnl_diff_mean": tca.get("pnl_diff_mean"), + "pnl_diff_std": tca.get("pnl_diff_std"), + } + + metrics_path = Path(args.metrics) + metrics_path.parent.mkdir(parents=True, exist_ok=True) + write_header = not metrics_path.exists() + with metrics_path.open("a", newline="", encoding="utf-8") as f: + writer = csv.DictWriter(f, fieldnames=row.keys()) + if write_header: + writer.writeheader() + writer.writerow(row) + print(f"Appended metrics row for run={args.run_id}") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/validate_dataset.py b/Q Research/scripts/validate_dataset.py new file mode 100644 index 0000000..902e585 --- /dev/null +++ b/Q Research/scripts/validate_dataset.py @@ -0,0 +1,206 @@ +#!/usr/bin/env python3 +""" +Validate dataset quality (missing data, duplicate timestamps, gaps, outliers). + +Example: + python scripts/validate_dataset.py --path data/raw/EURUSD_H1.csv +""" + +from __future__ import annotations + +import argparse +import json +from datetime import datetime, timezone +from pathlib import Path +from typing import Dict, Optional, Tuple + +import numpy as np +import pandas as pd +from loguru import logger + + +TIME_COLUMNS = ["ts", "time", "timestamp", "datetime", "date"] +DEFAULT_MANIFEST = "data/_manifest.json" +REPORT_DIR = Path("results/data_quality") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Validate dataset quality.") + parser.add_argument("--path", required=True, help="Dataset file path (relative or absolute).") + parser.add_argument( + "--manifest", + default=DEFAULT_MANIFEST, + help="Optional manifest JSON to enrich report.", + ) + parser.add_argument( + "--output", + default=None, + help="Optional explicit path for JSON report.", + ) + parser.add_argument( + "--max-outlier-z", + type=float, + default=5.0, + help="Z-score threshold to flag numeric outliers.", + ) + return parser.parse_args() + + +def detect_time_column(df: pd.DataFrame) -> Optional[str]: + for col in TIME_COLUMNS: + if col in df.columns: + return col + for col in df.columns: + if pd.api.types.is_datetime64_any_dtype(df[col]): + return col + return None + + +def infer_interval_seconds(series: pd.Series) -> Optional[float]: + ts = pd.to_datetime(series, utc=True, errors="coerce").dropna() + if ts.size < 3: + return None + diffs = np.diff(ts.values.astype("datetime64[ns]").astype(np.int64) // 1_000_000_000) + diffs = diffs[diffs > 0] + if diffs.size == 0: + return None + return float(np.median(diffs)) + + +def load_manifest_entry(manifest_path: Path, dataset_path: Path) -> Optional[Dict]: + if not manifest_path.exists(): + return None + try: + with manifest_path.open("r", encoding="utf-8") as fh: + manifest = json.load(fh) + except Exception as exc: + logger.warning(f"Failed to parse manifest {manifest_path}: {exc}") + return None + rel = str(dataset_path) + rel_alt = str(dataset_path.resolve().relative_to(Path.cwd())) + for entry in manifest.get("files", []): + if entry.get("path") in (rel, rel_alt): + return entry + return None + + +def summarize_numeric_outliers(df: pd.DataFrame, threshold: float) -> Dict[str, int]: + outliers: Dict[str, int] = {} + numeric_cols = df.select_dtypes(include=[np.number]).columns + if numeric_cols.empty: + return outliers + for col in numeric_cols: + series = df[col].dropna() + if series.empty: + continue + zscores = np.abs((series - series.mean()) / (series.std(ddof=1) or 1)) + count = int((zscores > threshold).sum()) + if count: + outliers[col] = count + return outliers + + +def compute_report(path: Path, manifest_entry: Optional[Dict], z_threshold: float) -> Dict: + df = pd.read_csv(path) + total_rows = int(df.shape[0]) + null_counts = df.isna().sum().to_dict() + + time_col = detect_time_column(df) + duplicates = 0 + gap_ratio = 0.0 + inferred_interval = None + if time_col: + ts = pd.to_datetime(df[time_col], utc=True, errors="coerce") + duplicates = int(ts.duplicated().sum()) + valid_ts = ts.dropna().sort_values() + inferred_interval = infer_interval_seconds(valid_ts) + if inferred_interval: + diffs = np.diff(valid_ts.values.astype("datetime64[ns]").astype(np.int64) // 1_000_000_000) + expected = inferred_interval + gaps = diffs[diffs > expected * 1.5] + gap_ratio = float(len(gaps) / max(len(diffs), 1)) + outliers = summarize_numeric_outliers(df, z_threshold) + + messages = [] + severity = "pass" + if total_rows == 0: + severity = "error" + messages.append("Dataset is empty.") + if duplicates > 0: + severity = "error" + messages.append(f"Found {duplicates} duplicate timestamps in column '{time_col}'.") + if gap_ratio > 0.01: + severity = "error" + messages.append(f"Timestamp gaps exceed 1% of intervals (ratio={gap_ratio:.4f}).") + elif gap_ratio > 0: + severity = "warn" + messages.append(f"Non-zero timestamp gaps detected (ratio={gap_ratio:.4f}).") + null_ratio = max((count / total_rows) if total_rows else 0 for count in null_counts.values() or [0]) + if null_ratio > 0.05: + severity = "error" + messages.append(f"Columns contain >5% nulls (max ratio={null_ratio:.4f}).") + elif null_ratio > 0: + severity = "warn" + messages.append(f"Columns contain nulls (max ratio={null_ratio:.4f}).") + if outliers: + if severity == "pass": + severity = "warn" + messages.append(f"Detected numeric outliers (threshold z>{z_threshold}).") + + report = { + "generated_at": datetime.now(timezone.utc).isoformat(), + "dataset_path": str(path), + "total_rows": total_rows, + "time_column": time_col, + "inferred_interval_seconds": inferred_interval, + "duplicate_timestamps": duplicates, + "gap_ratio": gap_ratio, + "null_counts": null_counts, + "numeric_outliers": outliers, + "severity": severity, + "messages": messages, + } + if manifest_entry: + report["manifest"] = { + "path": manifest_entry.get("path"), + "sha256": manifest_entry.get("sha256"), + "rows": manifest_entry.get("rows"), + "time_start": manifest_entry.get("time_start"), + "time_end": manifest_entry.get("time_end"), + } + return report + + +def save_report(report: Dict, output_path: Optional[Path], dataset_path: Path) -> Path: + if output_path is None: + REPORT_DIR.mkdir(parents=True, exist_ok=True) + ts = datetime.now().strftime("%Y%m%d_%H%M%S") + name = dataset_path.stem + output_path = REPORT_DIR / f"{ts}_{name}.json" + else: + output_path = output_path.resolve() + output_path.parent.mkdir(parents=True, exist_ok=True) + + with output_path.open("w", encoding="utf-8") as fh: + json.dump(report, fh, indent=2, ensure_ascii=False) + return output_path + + +def main() -> None: + args = parse_args() + dataset_path = Path(args.path).expanduser().resolve() + if not dataset_path.exists(): + raise FileNotFoundError(f"Dataset not found: {dataset_path}") + + manifest_entry = None + if args.manifest: + manifest_entry = load_manifest_entry(Path(args.manifest), dataset_path) + + report = compute_report(dataset_path, manifest_entry, args.max_outlier_z) + output_path = save_report(report, Path(args.output) if args.output else None, dataset_path) + logger.info(f"Validation severity: {report['severity']}") + logger.info(f"Report written to {output_path}") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/validate_results.py b/Q Research/scripts/validate_results.py new file mode 100644 index 0000000..7c3b084 --- /dev/null +++ b/Q Research/scripts/validate_results.py @@ -0,0 +1,101 @@ +#!/usr/bin/env python3 +""" +Validate a backtest result folder produced by run_once/grid/batch. + +Checks: + - summary.json / metrics.json exist + - core KPI fields present and not None + - data report metadata available + referenced file存在(可选) +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path +from typing import List + +BASE_DIR = Path(__file__).resolve().parents[1] + +REQUIRED_FILES = ["summary.json", "metrics.json"] +KPI_KEYS = [ + "final_equity", + "ann_return", + "ann_vol", + "sharpe", + "sortino", + "calmar", + "max_drawdown", + "max_drawdown_duration_bars", + "recovery_time_bars", +] + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Validate contents of a results/ directory.") + parser.add_argument("path", help="Path to results/ directory.") + parser.add_argument( + "--require-data-report", + action="store_true", + help="Fail if data_report metadata or referenced file is missing.", + ) + return parser.parse_args() + + +def load_json(path: Path): + with path.open("r", encoding="utf-8") as fh: + return json.load(fh) + + +def main() -> None: + args = parse_args() + run_path = Path(args.path).expanduser().resolve() + if not run_path.exists(): + print(f"[ERROR] Run directory not found: {run_path}", file=sys.stderr) + sys.exit(1) + + errors: List[str] = [] + + files = {} + for name in REQUIRED_FILES: + file_path = run_path / name + if not file_path.exists(): + errors.append(f"Missing required file: {file_path}") + else: + files[name] = load_json(file_path) + + summary = files.get("summary.json") or {} + metrics = files.get("metrics.json") or summary.get("metrics") or {} + + for key in KPI_KEYS: + if key not in metrics or metrics[key] is None: + errors.append(f"Metric '{key}' missing or null in metrics.json") + + if args.require_data_report: + data_meta = summary.get("data_report") + if not data_meta: + errors.append("data_report metadata missing in summary.json") + else: + if data_meta.get("severity") is None: + errors.append("data_report.severity missing in summary.json") + report_path = metrics.get("data_report") + if not report_path: + errors.append("metrics.data_report missing (expected relative path to JSON)") + else: + report_path_obj = Path(report_path) + report_full = report_path_obj if report_path_obj.is_absolute() else (BASE_DIR / report_path_obj).resolve() + if not report_full.exists(): + errors.append(f"Referenced data report not found: {report_path}") + + if errors: + print("[FAIL] Result validation failed:") + for err in errors: + print(f" - {err}") + sys.exit(1) + + print(f"[OK] Result folder {run_path.name} passed validation.") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/validate_stress_scenarios.py b/Q Research/scripts/validate_stress_scenarios.py new file mode 100644 index 0000000..57c1f16 --- /dev/null +++ b/Q Research/scripts/validate_stress_scenarios.py @@ -0,0 +1,49 @@ +#!/usr/bin/env python3 +"""Validate config/stress_scenarios.yaml definitions.""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +BASE_DIR = Path(__file__).resolve().parents[1] +import sys + +if str(BASE_DIR) not in sys.path: + sys.path.insert(0, str(BASE_DIR)) + +from scripts.scenario_utils import ALLOWED_FIELDS, load_scenarios, validate_data + +DEFAULT_PATH = Path("config/stress_scenarios.yaml") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Validate stress scenario catalog") + parser.add_argument( + "--path", + default=str(DEFAULT_PATH), + help="Path to stress_scenarios YAML (default: config/stress_scenarios.yaml)", + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + path = Path(args.path).expanduser() + try: + raw = load_scenarios(path) + except Exception as exc: + print(f"Validation failed: {exc}") + sys.exit(1) + errors = validate_data(raw) + if errors: + print("Validation failed:") + for err in errors: + print(f" - {err}") + sys.exit(1) + print(f"{path} OK ({len(ALLOWED_FIELDS)} allowed fields validated, {len(raw)} scenarios).") + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/watch_ops_metrics.py b/Q Research/scripts/watch_ops_metrics.py new file mode 100644 index 0000000..fecf400 --- /dev/null +++ b/Q Research/scripts/watch_ops_metrics.py @@ -0,0 +1,79 @@ +#!/usr/bin/env python3 +""" +Run multiple metric validators (risk, latency, pnl) in a single command. +""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import pandas as pd + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Watch aggregated risk/ops metrics.") + parser.add_argument("--csv", default="results/risk/metrics.csv", help="Metrics CSV path.") + parser.add_argument("--max-rejects", type=int, default=0) + parser.add_argument("--max-latency-ms", type=float, default=500.0) + parser.add_argument("--min-pnl", type=float, default=-3000.0) + parser.add_argument("--max-exposure", type=float, default=2_000_000.0) + parser.add_argument("--max-drawdown", type=float, default=0.1) + parser.add_argument("--min-live-sharpe", type=float, default=1.4) + parser.add_argument("--max-live-drawdown", type=float, default=0.05) + parser.add_argument("--max-live-latency-ms", type=float, default=500.0) + parser.add_argument("--max-slippage-bps", type=float, default=2.0) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + csv_path = Path(args.csv) + if not csv_path.exists(): + raise SystemExit(f"metrics CSV not found: {csv_path}") + df = pd.read_csv(csv_path) + if df.empty: + raise SystemExit("metrics CSV is empty.") + latest = df.tail(1).iloc[0] + rejects = latest.get("rejects", 0) + status = str(latest.get("status", "unknown")).lower() + latency = latest.get("latency_ms_avg", 0.0) + pnl = latest.get("total_pnl", 0.0) + max_exposure = latest.get("max_symbol_exposure", 0.0) + drawdown = latest.get("max_drawdown_pct", 0.0) + live_sharpe = latest.get("rolling_sharpe_30d", float("nan")) + live_drawdown = latest.get("live_drawdown_pct", float("nan")) + live_latency = latest.get("live_latency_ms_p95", float("nan")) + slippage_bps = latest.get("slippage_bps", float("nan")) + + errors = [] + if rejects > args.max_rejects or status != "pass": + errors.append(f"Rejects/status violation (rejects={rejects}, status={status})") + if latency > args.max_latency_ms: + errors.append(f"Latency {latency:.1f}ms > threshold {args.max_latency_ms}") + if pnl < args.min_pnl: + errors.append(f"Total PnL {pnl:.2f} < min {args.min_pnl}") + if max_exposure > args.max_exposure: + errors.append(f"Exposure {max_exposure:.2f} > max {args.max_exposure}") + if drawdown > args.max_drawdown: + errors.append(f"Drawdown {drawdown:.3f} > max {args.max_drawdown}") + if not pd.isna(live_sharpe) and live_sharpe < args.min_live_sharpe: + errors.append(f"Live Sharpe {live_sharpe:.2f} < min {args.min_live_sharpe}") + if not pd.isna(live_drawdown) and live_drawdown > args.max_live_drawdown: + errors.append(f"Live drawdown {live_drawdown:.3f} > max {args.max_live_drawdown}") + if not pd.isna(live_latency) and live_latency > args.max_live_latency_ms: + errors.append(f"Live latency p95 {live_latency:.1f}ms > max {args.max_live_latency_ms}") + if not pd.isna(slippage_bps) and slippage_bps > args.max_slippage_bps: + errors.append(f"Slippage {slippage_bps:.2f}bps > max {args.max_slippage_bps}") + print( + f"[watch_ops_metrics] run={latest.get('run_id')} status={status} " + f"rejects={rejects} latency_avg={latency} pnl={pnl} exposure={max_exposure} drawdown={drawdown} " + f"live_sharpe={live_sharpe} live_drawdown={live_drawdown} live_latency_p95={live_latency} slippage_bps={slippage_bps}" + ) + if errors: + raise SystemExit("; ".join(errors)) + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/watch_quality.py b/Q Research/scripts/watch_quality.py new file mode 100644 index 0000000..cedea0f --- /dev/null +++ b/Q Research/scripts/watch_quality.py @@ -0,0 +1,80 @@ +#!/usr/bin/env python3 +""" +Scan recent data-quality reports and raise alerts when severity >= warn. +""" + +from __future__ import annotations + +import argparse +import json +from datetime import datetime, timezone, timedelta +from pathlib import Path +from typing import List, Dict +from urllib import request, error + +REPORT_DIR = Path("results/data_quality") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Watch data-quality reports.") + parser.add_argument("--hours", type=float, default=24.0, help="Lookback window in hours (default 24).") + parser.add_argument("--report-dir", default=str(REPORT_DIR), help="Directory with data quality JSON files.") + parser.add_argument("--webhook-url", help="Optional webhook URL for POST notifications.") + return parser.parse_args() + + +def load_recent_reports(report_dir: Path, min_ts: datetime) -> List[Dict]: + rows: List[Dict] = [] + if not report_dir.exists(): + return rows + for path in report_dir.glob("*.json"): + try: + data = json.load(path.open("r", encoding="utf-8")) + except Exception: + continue + ts_raw = data.get("generated_at") + if not ts_raw: + continue + ts = datetime.fromisoformat(ts_raw.replace("Z", "+00:00")) + if ts >= min_ts: + data["__file"] = str(path) + rows.append(data) + return rows + + +def notify(webhook: str, message: str) -> None: + payload = json.dumps({"text": message}).encode("utf-8") + req = request.Request(webhook, data=payload, headers={"Content-Type": "application/json"}) + try: + request.urlopen(req, timeout=10) + except error.URLError as exc: + print(f"Failed to deliver webhook notification: {exc}") + + +def main() -> None: + args = parse_args() + cutoff = datetime.now(timezone.utc) - timedelta(hours=args.hours) + reports = load_recent_reports(Path(args.report_dir), cutoff) + alerts = [ + r for r in reports + if r.get("severity") in {"warn", "error"} + ] + if not alerts: + print("No alerts in the selected window.") + return + lines = [] + for r in alerts: + manifest = r.get("manifest") or {} + lines.append( + f"{r.get('generated_at')} | {manifest.get('path')} | severity={r.get('severity')} " + f"gap_ratio={r.get('gap_ratio')} file={r.get('__file')}" + ) + message = "\n".join(["Data quality alerts detected:"] + lines) + print(message) + if args.webhook_url: + notify(args.webhook_url, message) + raise SystemExit(1) + + +if __name__ == "__main__": + main() diff --git a/Q Research/scripts/watch_risk_metrics.py b/Q Research/scripts/watch_risk_metrics.py new file mode 100644 index 0000000..be580eb --- /dev/null +++ b/Q Research/scripts/watch_risk_metrics.py @@ -0,0 +1,56 @@ +#!/usr/bin/env python3 +""" +Watchdog script for risk metrics. + +If the latest entry in results/risk/metrics.csv has status=fail or rejects>0, +exit with non-zero status (can be used in cron/CI). +""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import pandas as pd + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Check latest risk metrics entry.") + parser.add_argument("--csv", default="results/risk/metrics.csv", help="Metrics CSV path.") + parser.add_argument("--allow-rejects", type=int, default=0, help="Maximum rejects allowed before alert.") + parser.add_argument("--allow-status", choices=["pass", "any"], default="pass", help="Expected status for latest run.") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + csv_path = Path(args.csv) + if not csv_path.exists(): + raise SystemExit(f"metrics CSV not found: {csv_path}") + df = pd.read_csv(csv_path) + if df.empty: + raise SystemExit("metrics CSV is empty.") + latest = df.tail(1).iloc[0] + rejects = int(latest.get("rejects", 0)) + status = str(latest.get("status", "unknown")).lower() + + ok = True + messages = [] + if rejects > args.allow_rejects: + ok = False + messages.append(f"Rejects={rejects} exceed allow_rejects={args.allow_rejects}") + if args.allow_status == "pass" and status != "pass": + ok = False + messages.append(f"Status={status}") + + print( + f"[watch_risk_metrics] latest run={latest.get('run_id')} " + f"status={status} rejects={rejects} kills={latest.get('kills')}" + ) + if not ok: + raise SystemExit(" ; ".join(messages)) + + +if __name__ == "__main__": + main() diff --git a/Q Research/strategies/__init__.py b/Q Research/strategies/__init__.py new file mode 100644 index 0000000..6d12133 --- /dev/null +++ b/Q Research/strategies/__init__.py @@ -0,0 +1,37 @@ +# strategies/__init__.py +from typing import Dict, Callable, Any + +# 策略注册表:name -> class +_REGISTRY: Dict[str, Callable[..., Any]] = {} + +def register(name: str): + """用作装饰器:@register('sma_atr')""" + def deco(cls): + _REGISTRY[name] = cls + return cls + return deco + +def _lazy_import_all(): + """ + 懒加载:首次 load_strategy 时再导入具体策略文件, + 这样不会因为循环依赖或路径问题导致注册表是空的。 + """ + # 在这里逐个导入具体策略模块;导入发生时模块内的 @register 会把类放进 _REGISTRY + from . import sma_atr # noqa: F401 + from . import regime_sma # noqa: F401 + from . import band_mean_revert # noqa: F401 + from . import bollinger_mean_revert # noqa: F401 + from . import ma_crossover # noqa: F401 + from . import momentum # noqa: F401 + from . import xgb_signal # noqa: F401 + +def load_strategy(name: str, **kwargs): + # 第一次用时尝试懒加载,填充注册表 + if not _REGISTRY: + _lazy_import_all() + if name not in _REGISTRY: + # 再尝试一次(防止用户后来才添加文件) + _lazy_import_all() + if name not in _REGISTRY: + raise ValueError(f"Unknown strategy: {name}. Available: {list(_REGISTRY.keys())}") + return _REGISTRY[name](**kwargs) diff --git a/Q Research/strategies/__pycache__/__init__.cpython-312.pyc b/Q Research/strategies/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000..06e99fd Binary files /dev/null and b/Q Research/strategies/__pycache__/__init__.cpython-312.pyc differ diff --git a/Q Research/strategies/__pycache__/__init__.cpython-313.pyc b/Q Research/strategies/__pycache__/__init__.cpython-313.pyc new file mode 100644 index 0000000..c95b9a4 Binary files /dev/null and b/Q Research/strategies/__pycache__/__init__.cpython-313.pyc differ diff --git a/Q Research/strategies/__pycache__/band_mean_revert.cpython-312.pyc b/Q Research/strategies/__pycache__/band_mean_revert.cpython-312.pyc new file mode 100644 index 0000000..0567f0d Binary files /dev/null and b/Q Research/strategies/__pycache__/band_mean_revert.cpython-312.pyc differ diff --git a/Q Research/strategies/__pycache__/band_mean_revert.cpython-313.pyc b/Q Research/strategies/__pycache__/band_mean_revert.cpython-313.pyc new file mode 100644 index 0000000..23fd7ea Binary files /dev/null and b/Q Research/strategies/__pycache__/band_mean_revert.cpython-313.pyc differ diff --git a/Q Research/strategies/__pycache__/base.cpython-312.pyc b/Q Research/strategies/__pycache__/base.cpython-312.pyc new file mode 100644 index 0000000..69c0db1 Binary files /dev/null and b/Q Research/strategies/__pycache__/base.cpython-312.pyc differ diff --git a/Q Research/strategies/__pycache__/base.cpython-313.pyc b/Q Research/strategies/__pycache__/base.cpython-313.pyc new file mode 100644 index 0000000..6032be4 Binary files /dev/null and b/Q Research/strategies/__pycache__/base.cpython-313.pyc differ diff --git a/Q Research/strategies/__pycache__/bollinger_mean_revert.cpython-312.pyc b/Q Research/strategies/__pycache__/bollinger_mean_revert.cpython-312.pyc new file mode 100644 index 0000000..5d7bb96 Binary files /dev/null and b/Q Research/strategies/__pycache__/bollinger_mean_revert.cpython-312.pyc differ diff --git a/Q Research/strategies/__pycache__/bollinger_mean_revert.cpython-313.pyc b/Q Research/strategies/__pycache__/bollinger_mean_revert.cpython-313.pyc new file mode 100644 index 0000000..5408d0a Binary files /dev/null and b/Q Research/strategies/__pycache__/bollinger_mean_revert.cpython-313.pyc differ diff --git a/Q Research/strategies/__pycache__/ma_cross.cpython-313.pyc b/Q Research/strategies/__pycache__/ma_cross.cpython-313.pyc new file mode 100644 index 0000000..bbd8008 Binary files /dev/null and b/Q Research/strategies/__pycache__/ma_cross.cpython-313.pyc differ diff --git a/Q Research/strategies/__pycache__/ma_crossover.cpython-312.pyc b/Q Research/strategies/__pycache__/ma_crossover.cpython-312.pyc new file mode 100644 index 0000000..bd44477 Binary files /dev/null and b/Q Research/strategies/__pycache__/ma_crossover.cpython-312.pyc differ diff --git a/Q Research/strategies/__pycache__/ma_crossover.cpython-313.pyc b/Q Research/strategies/__pycache__/ma_crossover.cpython-313.pyc new file mode 100644 index 0000000..8424c08 Binary files /dev/null and b/Q Research/strategies/__pycache__/ma_crossover.cpython-313.pyc differ diff --git a/Q Research/strategies/__pycache__/mean_reversion.cpython-313.pyc b/Q Research/strategies/__pycache__/mean_reversion.cpython-313.pyc new file mode 100644 index 0000000..560e2d7 Binary files /dev/null and b/Q Research/strategies/__pycache__/mean_reversion.cpython-313.pyc differ diff --git a/Q Research/strategies/__pycache__/momentum.cpython-312.pyc b/Q Research/strategies/__pycache__/momentum.cpython-312.pyc new file mode 100644 index 0000000..a0d0c7b Binary files /dev/null and b/Q Research/strategies/__pycache__/momentum.cpython-312.pyc differ diff --git a/Q Research/strategies/__pycache__/momentum.cpython-313.pyc b/Q Research/strategies/__pycache__/momentum.cpython-313.pyc new file mode 100644 index 0000000..0fcbbbc Binary files /dev/null and b/Q Research/strategies/__pycache__/momentum.cpython-313.pyc differ diff --git a/Q Research/strategies/__pycache__/regime_sma.cpython-312.pyc b/Q Research/strategies/__pycache__/regime_sma.cpython-312.pyc new file mode 100644 index 0000000..28df1fb Binary files /dev/null and b/Q Research/strategies/__pycache__/regime_sma.cpython-312.pyc differ diff --git a/Q Research/strategies/__pycache__/regime_sma.cpython-313.pyc b/Q Research/strategies/__pycache__/regime_sma.cpython-313.pyc new file mode 100644 index 0000000..2436c2a Binary files /dev/null and b/Q Research/strategies/__pycache__/regime_sma.cpython-313.pyc differ diff --git a/Q Research/strategies/__pycache__/sma_atr.cpython-312.pyc b/Q Research/strategies/__pycache__/sma_atr.cpython-312.pyc new file mode 100644 index 0000000..3081275 Binary files /dev/null and b/Q Research/strategies/__pycache__/sma_atr.cpython-312.pyc differ diff --git a/Q Research/strategies/__pycache__/sma_atr.cpython-313.pyc b/Q Research/strategies/__pycache__/sma_atr.cpython-313.pyc new file mode 100644 index 0000000..26082da Binary files /dev/null and b/Q Research/strategies/__pycache__/sma_atr.cpython-313.pyc differ diff --git a/Q Research/strategies/__pycache__/xgb_signal.cpython-312.pyc b/Q Research/strategies/__pycache__/xgb_signal.cpython-312.pyc new file mode 100644 index 0000000..a49a3b9 Binary files /dev/null and b/Q Research/strategies/__pycache__/xgb_signal.cpython-312.pyc differ diff --git a/Q Research/strategies/__pycache__/xgb_signal.cpython-313.pyc b/Q Research/strategies/__pycache__/xgb_signal.cpython-313.pyc new file mode 100644 index 0000000..3ac244d Binary files /dev/null and b/Q Research/strategies/__pycache__/xgb_signal.cpython-313.pyc differ diff --git a/Q Research/strategies/band_mean_revert.py b/Q Research/strategies/band_mean_revert.py new file mode 100644 index 0000000..b0b8528 --- /dev/null +++ b/Q Research/strategies/band_mean_revert.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +from . import register +from .base import Strategy + + +@register("band_mean_revert") +class BandMeanRevert(Strategy): + """ + 简单区间/均值回归策略: + - 使用慢均线 +/- ATR*mult 作为区间带; + - 当价格跌破下带且 RSI 低于阈值时做多; + - 当价格突破上带且 RSI 高于阈值时做空(可选)。 + - 价格回到均值或 RSI 归中时离场。 + """ + + def __init__( + self, + band_atr_mult: float = 1.5, + rsi_long: float = 35.0, + rsi_short: float = 65.0, + exit_rsi_mid: float = 50.0, + allow_short: bool = True, + ) -> None: + super().__init__( + band_atr_mult=band_atr_mult, + rsi_long=rsi_long, + rsi_short=rsi_short, + exit_rsi_mid=exit_rsi_mid, + allow_short=allow_short, + ) + self.band_atr_mult = float(band_atr_mult) + self.rsi_long = float(rsi_long) + self.rsi_short = float(rsi_short) + self.exit_rsi_mid = float(exit_rsi_mid) + self.allow_short = bool(allow_short) + + def on_bar(self, state: dict) -> dict: + close = state.get("close") + sma_slow = state.get("sma_slow") + atr = state.get("curr_atr") + rsi = state.get("rsi") + position = state.get("position", 0) + + if close is None or sma_slow is None or atr is None or rsi is None: + return {"action": "HOLD"} + + upper = sma_slow + self.band_atr_mult * atr + lower = sma_slow - self.band_atr_mult * atr + + if position == 0: + if close <= lower and rsi <= self.rsi_long: + return {"action": "ENTER_LONG"} + if self.allow_short and close >= upper and rsi >= self.rsi_short: + return {"action": "ENTER_SHORT"} + elif position > 0: + if close >= sma_slow or rsi >= self.exit_rsi_mid: + return {"action": "EXIT_LONG"} + elif position < 0: + if close <= sma_slow or rsi <= self.exit_rsi_mid: + return {"action": "EXIT_SHORT"} + + return {"action": "HOLD"} diff --git a/Q Research/strategies/base.py b/Q Research/strategies/base.py new file mode 100644 index 0000000..d121483 --- /dev/null +++ b/Q Research/strategies/base.py @@ -0,0 +1,16 @@ +# strategies/base.py +import os +import sys +sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from typing import Dict, Any + +class Strategy: + """ + 只负责“发信号”,不做撮合和资金结算。 + on_bar 输入 state,输出 {"action": "ENTER_LONG"/"EXIT_LONG"/"HOLD"}。 + """ + def __init__(self, **params: Any) -> None: + self.params = params + + def on_bar(self, state: Dict[str, Any]) -> Dict[str, Any]: + return {"action": "HOLD"} \ No newline at end of file diff --git a/Q Research/strategies/bollinger_mean_revert.py b/Q Research/strategies/bollinger_mean_revert.py new file mode 100644 index 0000000..5785f5e --- /dev/null +++ b/Q Research/strategies/bollinger_mean_revert.py @@ -0,0 +1,78 @@ +from __future__ import annotations + +import numpy as np + +from . import register +from .base import Strategy + + +@register("bollinger_mean_revert") +class BollingerMeanRevert(Strategy): + """ + Simple Bollinger-band based mean reversion strategy skeleton. + - Enters long when price falls `enter_z` standard deviations below the mean. + - Enters short symmetrically (if allow_short). + - Exits when price reverts back within `exit_z` standard deviations. + This is meant to be combined with other sleeves in portfolio tests. + """ + + def __init__( + self, + window: int = 50, + num_std: float = 2.0, + enter_z: float = 1.0, + exit_z: float = 0.2, + allow_short: bool = True, + cooldown: int = 0, + ) -> None: + super().__init__( + window=window, + num_std=num_std, + enter_z=enter_z, + exit_z=exit_z, + allow_short=allow_short, + cooldown=cooldown, + ) + self.window = int(window) + self.num_std = float(num_std) + self.enter_z = float(enter_z) + self.exit_z = float(exit_z) + self.allow_short = bool(allow_short) + self.cooldown = int(cooldown or 0) + self._next_entry_bar = 0 + + def on_bar(self, state: dict) -> dict: + closes = state.get("close_history") + bar_idx = state.get("bar_idx", 0) + position = state.get("position", 0) + if closes is None or len(closes) < self.window: + return {"action": "HOLD"} + + window_data = np.array(closes[-self.window:], dtype=float) + mean = window_data.mean() + std = window_data.std(ddof=0) + if std == 0: + return {"action": "HOLD"} + + price = float(window_data[-1]) + z_score = (price - mean) / std + + # Enforce cooldown between fresh entries + if bar_idx < self._next_entry_bar and position == 0: + return {"action": "HOLD"} + + if position == 0: + if z_score <= -self.enter_z: + self._next_entry_bar = bar_idx + self.cooldown + return {"action": "ENTER_LONG"} + if self.allow_short and z_score >= self.enter_z: + self._next_entry_bar = bar_idx + self.cooldown + return {"action": "ENTER_SHORT"} + elif position > 0: + if z_score >= -self.exit_z: + return {"action": "EXIT_LONG"} + else: # position < 0 + if z_score <= self.exit_z: + return {"action": "EXIT_SHORT"} + + return {"action": "HOLD"} diff --git a/Q Research/strategies/ma_cross.py b/Q Research/strategies/ma_cross.py new file mode 100644 index 0000000..3d0da9b --- /dev/null +++ b/Q Research/strategies/ma_cross.py @@ -0,0 +1,33 @@ +# FX_BACKTEST/strategies/ma_cross.py +import os +import sys +sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from collections import deque +from core.events import TickEvent, SignalEvent + +class MACross: + def __init__(self, q, symbol: str, short: int = 20, long: int = 50, size: float = 10000.0): + assert short < long + self.q = q + self.symbol = symbol + self.short_n, self.long_n = short, long + self.short_win, self.long_win = deque(maxlen=short), deque(maxlen=long) + self.pos = 0.0 # 当前方向(>0 多 / <0 空 / =0 空仓) + self.size = size + + def on_event(self, ev): + if isinstance(ev, TickEvent) and ev.symbol == self.symbol: + mid = (ev.bid + ev.ask) / 2.0 + self.short_win.append(mid) + self.long_win.append(mid) + if len(self.long_win) < self.long_n: + return + sma_s = sum(self.short_win) / len(self.short_win) + sma_l = sum(self.long_win) / len(self.long_win) + # 交叉信号 + if self.pos <= 0 and sma_s > sma_l: # 金叉 -> 做多 + self.q.put(SignalEvent(ev.ts, self.symbol, "LONG", self.size)) + self.pos = 1 + elif self.pos >= 0 and sma_s < sma_l: # 死叉 -> 做空 + self.q.put(SignalEvent(ev.ts, self.symbol, "SHORT", self.size)) + self.pos = -1 \ No newline at end of file diff --git a/Q Research/strategies/ma_crossover.py b/Q Research/strategies/ma_crossover.py new file mode 100644 index 0000000..6e39c56 --- /dev/null +++ b/Q Research/strategies/ma_crossover.py @@ -0,0 +1,94 @@ +"""Simple moving-average crossover strategy registered for StrategyEngine combos.""" + +from __future__ import annotations + +from typing import Dict, Any + +from . import register +from .base import Strategy + + +@register("ma_crossover") +class MovingAverageCrossover(Strategy): + """ + Emits ENTER/EXIT signals when a fast SMA crosses a slow SMA. + + State inputs expected from StrategyEngine: + - sma_fast / sma_slow + - bar_idx (int) + - position_units (float) + - default_qty (float) + """ + + def __init__( + self, + size_mult: float = 1.0, + cooldown_bars: int = 0, + exit_buffer_pct: float = 0.0, + allow_short: bool = True, + ) -> None: + super().__init__( + size_mult=size_mult, + cooldown_bars=cooldown_bars, + exit_buffer_pct=exit_buffer_pct, + allow_short=allow_short, + ) + self.size_mult = float(size_mult) + self.cooldown_bars = int(max(0, cooldown_bars)) + self.exit_buffer_pct = float(max(0.0, exit_buffer_pct)) + self.allow_short = bool(allow_short) + + self._prev_fast: float | None = None + self._prev_slow: float | None = None + self._block_until: int = 0 + + def on_bar(self, state: Dict[str, Any]) -> Dict[str, Any]: + fast = state.get("sma_fast") + slow = state.get("sma_slow") + bar_idx = int(state.get("bar_idx", 0) or 0) + position = float(state.get("position_units", 0.0) or 0.0) + + if fast is None or slow is None: + return {"action": "HOLD"} + + prev_fast = self._prev_fast + prev_slow = self._prev_slow + self._prev_fast = fast + self._prev_slow = slow + + if prev_fast is None or prev_slow is None: + return {"action": "HOLD"} + + if bar_idx < self._block_until: + return {"action": "HOLD"} + + buffer = self.exit_buffer_pct + size = self._position_size(state) + + crossed_up = prev_fast <= prev_slow and fast > slow + crossed_down = prev_fast >= prev_slow and fast < slow + + if crossed_up: + self._block_until = bar_idx + self.cooldown_bars + return {"action": "ENTER_LONG", "size": size} + + if crossed_down and self.allow_short: + self._block_until = bar_idx + self.cooldown_bars + return {"action": "ENTER_SHORT", "size": size} + + slow_with_buffer = slow * (1.0 + buffer) + slow_lower = slow * (1.0 - buffer) + + if position > 0 and fast < slow_lower: + return {"action": "EXIT_LONG"} + + if position < 0 and (fast > slow_with_buffer or not self.allow_short): + return {"action": "EXIT_SHORT"} + + return {"action": "HOLD"} + + def _position_size(self, state: Dict[str, Any]) -> float | None: + default_qty = state.get("default_qty") + if default_qty is None: + return None + return float(default_qty) * self.size_mult diff --git a/Q Research/strategies/mean_reversion.py b/Q Research/strategies/mean_reversion.py new file mode 100644 index 0000000..d166b1f --- /dev/null +++ b/Q Research/strategies/mean_reversion.py @@ -0,0 +1,17 @@ +import pandas as pd + +def mean_reversion_strategy(df: pd.DataFrame, short_window=20, long_window=50): + """ + 简单的均值回归策略(示例): + - 价格高于长期均线则卖出 + - 价格低于长期均线则买入 + """ + df = df.copy() + df["short_ma"] = df["close"].rolling(window=short_window).mean() + df["long_ma"] = df["close"].rolling(window=long_window).mean() + + df["signal"] = 0 + df.loc[df["short_ma"] > df["long_ma"], "signal"] = 1 # 买入 + df.loc[df["short_ma"] < df["long_ma"], "signal"] = -1 # 卖出 + + return df \ No newline at end of file diff --git a/Q Research/strategies/momentum.py b/Q Research/strategies/momentum.py new file mode 100644 index 0000000..c543d5f --- /dev/null +++ b/Q Research/strategies/momentum.py @@ -0,0 +1,104 @@ +"""Basic momentum breakout strategy usable inside StrategyEngine combos.""" + +from __future__ import annotations + +from typing import Dict, Any, Sequence + +from . import register +from .base import Strategy + + +@register("momentum_breakout") +class MomentumBreakout(Strategy): + """ + Uses rate-of-change over a configurable lookback to enter in the dominant direction. + + Parameters + ---------- + lookback : int + Number of bars between comparisons. + enter_threshold : float + Minimum absolute return (%) to trigger an entry (e.g. 0.002 = 20 bps). + exit_threshold : float + Momentum magnitude below which existing positions are flattened. + size_mult : float + Multiplier applied to StrategyEngine default_qty when sizing trades. + allow_short : bool + Whether to take short trades when momentum turns negative. + cooldown_bars : int + Minimum bars between successive entries. + """ + + def __init__( + self, + lookback: int = 24, + enter_threshold: float = 0.0015, + exit_threshold: float = 0.0005, + size_mult: float = 1.0, + allow_short: bool = True, + cooldown_bars: int = 0, + ) -> None: + super().__init__( + lookback=lookback, + enter_threshold=enter_threshold, + exit_threshold=exit_threshold, + size_mult=size_mult, + allow_short=allow_short, + cooldown_bars=cooldown_bars, + ) + self.lookback = max(1, int(lookback)) + self.enter_threshold = float(enter_threshold) + self.exit_threshold = float(exit_threshold) + self.size_mult = float(size_mult) + self.allow_short = bool(allow_short) + self.cooldown_bars = int(max(0, cooldown_bars)) + + self._block_until: int = 0 + + def on_bar(self, state: Dict[str, Any]) -> Dict[str, Any]: + closes = state.get("close_history") + bar_idx = int(state.get("bar_idx", 0) or 0) + position = float(state.get("position_units", 0.0) or 0.0) + + if not self._has_enough_history(closes): + return {"action": "HOLD"} + + roc = self._rate_of_change(closes) + + if position > 0 and roc < self.exit_threshold: + return {"action": "EXIT_LONG"} + if position < 0 and roc > -self.exit_threshold: + return {"action": "EXIT_SHORT"} + + if bar_idx < self._block_until: + return {"action": "HOLD"} + + size = self._position_size(state) + if roc >= self.enter_threshold: + self._block_until = bar_idx + self.cooldown_bars + return {"action": "ENTER_LONG", "size": size} + if roc <= -self.enter_threshold and self.allow_short: + self._block_until = bar_idx + self.cooldown_bars + return {"action": "ENTER_SHORT", "size": size} + + return {"action": "HOLD"} + + def _has_enough_history(self, closes: Any) -> bool: + if closes is None: + return False + if isinstance(closes, Sequence): + return len(closes) > self.lookback + return False + + def _rate_of_change(self, closes: Sequence[float]) -> float: + recent = float(closes[-1]) + past = float(closes[-(self.lookback + 1)]) + if past == 0: + return 0.0 + return (recent - past) / past + + def _position_size(self, state: Dict[str, Any]) -> float | None: + default_qty = state.get("default_qty") + if default_qty is None: + return None + return float(default_qty) * self.size_mult diff --git a/Q Research/strategies/regime_sma.py b/Q Research/strategies/regime_sma.py new file mode 100644 index 0000000..42af889 --- /dev/null +++ b/Q Research/strategies/regime_sma.py @@ -0,0 +1,246 @@ +# strategies/regime_sma.py + +from __future__ import annotations + +from datetime import datetime + +from . import register +from .base import Strategy +from .sma_atr import SmaAtr + + +@register("regime_sma") +class RegimeSMAStrategy(Strategy): + """ + Wraps the SMA+ATR strategy with a simple regime filter. + + - 在趋势 regime(由 StrategyEngine 提供)下,沿用 SmaAtr 信号。 + - 在震荡 regime 下,可选择保持空仓或使用简单的 RSI 均值回归。 + """ + + def __init__( + self, + trend_params: dict | None = None, + range_mode: str = "flat", + range_rsi_high: float = 75.0, + range_rsi_low: float = 25.0, + trend_min_bars: int = 0, + atr_percentile_min: float | None = None, + size_tiers: list | None = None, + base_size_mult: float = 1.0, + htf_alignment: bool = False, + htf_rsi_range: tuple[float, float] | list | None = None, + risk_rules: list | None = None, + ) -> None: + super().__init__() + self.range_mode = (range_mode or "flat").lower() + self.range_rsi_high = float(range_rsi_high) + self.range_rsi_low = float(range_rsi_low) + self.range_exit_mid = (self.range_rsi_high + self.range_rsi_low) / 2.0 + self.trend_min_bars = int(trend_min_bars or 0) + self.atr_percentile_min = atr_percentile_min if atr_percentile_min is None else float(atr_percentile_min) + self.base_size_mult = float(base_size_mult or 1.0) + self.size_tiers = self._normalize_tiers(size_tiers) + self.htf_alignment = bool(htf_alignment) + if htf_rsi_range: + lo, hi = htf_rsi_range + self.htf_rsi_min = float(lo) + self.htf_rsi_max = float(hi) + else: + self.htf_rsi_min = 30.0 + self.htf_rsi_max = 70.0 + self.risk_rules = risk_rules or [] + + trend_params = trend_params or {} + self.trend_strategy = SmaAtr(**trend_params) + + def on_bar(self, state: dict): + cooldown = self._check_risk_rules(state) + if cooldown: + return {"action": "HOLD", "cooldown_bars": cooldown} + + regime_label = state.get("regime_label", "unknown") + trend_streak = int(state.get("regime_trend_bars", 0) or 0) + atr_percentile = state.get("atr_percentile") + base_signal = self.trend_strategy.on_bar(state) or {} + action = base_signal.get("action", "HOLD") + + if regime_label == "trend": + if self.trend_min_bars and trend_streak < self.trend_min_bars: + regime_label = "range" + elif self.atr_percentile_min is not None: + if atr_percentile is None or atr_percentile < self.atr_percentile_min: + regime_label = "range" + if regime_label == "trend" and action.startswith("ENTER"): + if not self._passes_htf_filter(action, state): + regime_label = "range" + else: + sized = dict(base_signal) + sized["size"] = self._position_size(action, state) + return sized + elif regime_label == "trend": + return base_signal + + # 在趋势 regime,直接沿用趋势策略的信号 + if regime_label == "trend": + return base_signal + + # 非趋势 regime:允许趋势策略发出的平仓指令生效,但屏蔽入场 + if action.startswith("EXIT"): + return base_signal + + # 根据 range_mode 决定行为 + range_action = self._range_signal(state) + if range_action: + return {"action": range_action} + + # 默认保持空仓 + if state.get("position"): + # 持仓状态下交给风控(risk exit)或趋势策略的平仓指令处理 + return {"action": "HOLD"} + return {"action": "HOLD"} + + def _range_signal(self, state: dict) -> str | None: + if self.range_mode != "mean_revert": + return None + rsi = state.get("rsi") + if rsi is None: + return None + + position = state.get("position", 0) + if position == 0: + if rsi >= self.range_rsi_high: + return "ENTER_SHORT" + if rsi <= self.range_rsi_low: + return "ENTER_LONG" + elif position > 0 and rsi >= self.range_exit_mid: + return "EXIT_LONG" + elif position < 0 and rsi <= self.range_exit_mid: + return "EXIT_SHORT" + return None + + def _passes_htf_filter(self, action: str, state: dict) -> bool: + if not self.htf_alignment: + return True + htf_ema = state.get("htf_ema") + if htf_ema is None: + return False + close = state.get("close") + if close is None: + return False + if action == "ENTER_LONG" and close < htf_ema: + return False + if action == "ENTER_SHORT" and close > htf_ema: + return False + htf_rsi = state.get("htf_rsi") + if htf_rsi is not None: + if htf_rsi < self.htf_rsi_min or htf_rsi > self.htf_rsi_max: + return False + return True + + def _normalize_tiers(self, tiers: list | None) -> list: + if not tiers: + return [{"size_mult": self.base_size_mult}] + normalized = [] + for tier in tiers: + if not isinstance(tier, dict): + continue + entry = tier.copy() + entry["size_mult"] = float(entry.get("size_mult", 1.0)) + normalized.append(entry) + if not normalized: + normalized.append({"size_mult": self.base_size_mult}) + normalized.sort(key=lambda t: t.get("size_mult", 0), reverse=True) + return normalized + + def _position_size(self, action: str, state: dict) -> float: + base_qty = float(state.get("default_qty") or 0.0) + if base_qty <= 0: + return 0.0 + atr_pct = state.get("atr_percentile") + trend_strength = state.get("trend_strength") + streak = int(state.get("regime_trend_bars", 0) or 0) + for tier in self.size_tiers: + if self._tier_matches(tier, atr_pct, trend_strength, streak, action): + return base_qty * tier.get("size_mult", 1.0) + return base_qty * self.base_size_mult + + def _tier_matches(self, tier: dict, atr_pct, trend_strength, streak: int, action: str) -> bool: + min_atr = tier.get("min_atr_pct") + max_atr = tier.get("max_atr_pct") + min_strength = tier.get("min_trend_strength") + min_streak = tier.get("min_trend_bars") + allow_short = tier.get("allow_short") + allow_long = tier.get("allow_long") + if min_atr is not None: + if atr_pct is None or atr_pct < float(min_atr): + return False + if max_atr is not None and atr_pct is not None: + if atr_pct > float(max_atr): + return False + if min_strength is not None: + if trend_strength is None or trend_strength < float(min_strength): + return False + if min_streak is not None and streak < int(min_streak): + return False + if action == "ENTER_LONG" and allow_long is False: + return False + if action == "ENTER_SHORT" and allow_short is False: + return False + return True + + def _check_risk_rules(self, state: dict) -> int: + if not self.risk_rules: + return 0 + atr_pct = state.get("atr_percentile") + ts = self._to_datetime(state.get("ts")) + for rule in self.risk_rules: + rtype = rule.get("type") + if rtype == "atr_percentile": + min_v = rule.get("min") + max_v = rule.get("max") + triggered = False + if min_v is not None: + if atr_pct is None or atr_pct < float(min_v): + triggered = True + if max_v is not None and atr_pct is not None and atr_pct > float(max_v): + triggered = True + if triggered: + return int(rule.get("cooldown_bars", 0) or 0) + elif rtype == "calendar": + if ts is None: + continue + dates = rule.get("dates") or [] + date_str = ts.strftime("%Y-%m-%d") + if date_str in dates: + return int(rule.get("cooldown_bars", 0) or 0) + windows = rule.get("windows") or [] + for window in windows: + start = self._to_datetime(window.get("start")) + end = self._to_datetime(window.get("end")) + if start and end and start <= ts <= end: + return int(rule.get("cooldown_bars", 0) or 0) + elif rtype == "time_window": + start = self._to_datetime(rule.get("start")) + end = self._to_datetime(rule.get("end")) + if ts and start and end and start <= ts <= end: + return int(rule.get("cooldown_bars", 0) or 0) + return 0 + + def _to_datetime(self, value): + if value is None: + return None + if hasattr(value, "to_pydatetime"): + try: + return value.to_pydatetime() + except Exception: + pass + if isinstance(value, datetime): + return value + text = str(value) + if not text: + return None + try: + return datetime.fromisoformat(text.replace("Z", "+00:00")) + except Exception: + return None diff --git a/Q Research/strategies/sma_atr.py b/Q Research/strategies/sma_atr.py new file mode 100644 index 0000000..bb51a96 --- /dev/null +++ b/Q Research/strategies/sma_atr.py @@ -0,0 +1,110 @@ +# strategies/sma_atr.py +from .base import Strategy +from . import register + +@register("sma_atr") +class SmaAtr(Strategy): + """ + 参数(与 runner 对齐): + - long_only_above_slow: bool + - allow_short: bool + - short_only_below_slow: bool + - slope_lookback: int + - cooldown: int + - fast_win: int + - slow_win: int + - atr_sl / atr_tp / atr_window(仅用于记录,不在策略层计算) + """ + def on_bar(self, state): + c = state["close"] + position = state["position"] + rsi = state.get("rsi") + sma_fast = state.get("sma_fast") + sma_slow = state.get("sma_slow") + bar_idx = state["bar_idx"] + next_entry_bar_idx_long = state.get("next_entry_bar_idx_long", 0) + next_entry_bar_idx_short = state.get("next_entry_bar_idx_short", 0) + sma_fast_hist = state["sma_fast_hist"] # deque,最近值在右侧 + + fast_win = self.params.get("fast_win") + slow_win = self.params.get("slow_win") + long_only_above_slow = self.params.get("long_only_above_slow", False) + slope_lookback = self.params.get("slope_lookback", 0) + cooldown = self.params.get("cooldown", 0) + allow_short = self.params.get("allow_short", True) + short_only_below_slow = self.params.get("short_only_below_slow", False) + rsi_long_thresh = self.params.get("rsi_long_thresh") + rsi_short_thresh = self.params.get("rsi_short_thresh") + + # 均线就绪才判断 + if sma_fast is None or sma_slow is None: + return {"action": "HOLD"} + + go_long = (position == 0) and (sma_fast > sma_slow) + exit_long = (position == 1) and (sma_fast < sma_slow) + + # 仅做多需在慢均线上方 + if long_only_above_slow and go_long: + if not (c > sma_slow): + go_long = False + + # fast 斜率确认 + if slope_lookback and go_long: + if len(sma_fast_hist) > slope_lookback: + # 最近一个值与 L 根前比较 + if not (sma_fast_hist[-1] > sma_fast_hist[-1 - slope_lookback]): + go_long = False + else: + go_long = False + + # 冷却 + if cooldown and go_long: + if bar_idx < next_entry_bar_idx_long: + go_long = False + + # RSI 过滤(如果配置了阈值且 RSI 可用) + if go_long and rsi_long_thresh is not None: + if rsi is None: + go_long = False + else: + if not (rsi > float(rsi_long_thresh)): + go_long = False + + go_short = False + exit_short = False + if allow_short: + go_short = (position == 0) and (sma_fast < sma_slow) + exit_short = (position == -1) and (sma_fast > sma_slow) + + if short_only_below_slow and go_short: + if not (c < sma_slow): + go_short = False + + if slope_lookback and go_short: + if len(sma_fast_hist) > slope_lookback: + if not (sma_fast_hist[-1] < sma_fast_hist[-1 - slope_lookback]): + go_short = False + else: + go_short = False + + if cooldown and go_short: + if bar_idx < next_entry_bar_idx_short: + go_short = False + + # RSI 过滤(空头侧) + if go_short and rsi_short_thresh is not None: + if rsi is None: + go_short = False + else: + if not (rsi < float(rsi_short_thresh)): + go_short = False + + if exit_long: + return {"action": "EXIT_LONG"} + if exit_short: + return {"action": "EXIT_SHORT"} + if go_long: + return {"action": "ENTER_LONG"} + if go_short: + return {"action": "ENTER_SHORT"} + return {"action": "HOLD"} diff --git a/Q Research/strategies/xgb_signal.py b/Q Research/strategies/xgb_signal.py new file mode 100644 index 0000000..eaf3b84 --- /dev/null +++ b/Q Research/strategies/xgb_signal.py @@ -0,0 +1,274 @@ +"""XGBoost-based signal strategy (long-only v1). + +Loads a trained Booster + feature list + thresholds and emits ENTER_LONG/EXIT_LONG +decisions based on predicted probability compared to configured thresholds. + +Registration name: "xgb_signal" +""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any, Dict, List, Optional + +import numpy as np +from loguru import logger + +from . import register +from .base import Strategy + + +def _safe_get(d: Dict[str, Any], key: str, default=None): + v = d.get(key, default) + return v if v is not None else default + + +def _rsi_from_series(arr: np.ndarray, period: int = 14) -> float | None: + if arr.size < period + 1: + return None + diffs = np.diff(arr) + gains = diffs[diffs > 0] + losses = -diffs[diffs < 0] + avg_gain = gains.mean() if gains.size > 0 else 0.0 + avg_loss = losses.mean() if losses.size > 0 else 0.0 + if avg_loss == 0.0 and avg_gain == 0.0: + return 50.0 + if avg_loss == 0.0: + return 100.0 + rs = avg_gain / avg_loss + return 100.0 - (100.0 / (1.0 + rs)) + + +@register("xgb_signal") +class XGBSignal(Strategy): + def __init__( + self, + model_dir: Optional[str] = None, + latest_ptr: str = "QuantResearch/artifacts/models/usdjpy_h1_xgb_latest.json", + prob_long: Optional[float] = None, + prob_exit: Optional[float] = None, + size_mult: float = 1.0, + cooldown_bars: int = 0, + min_atr_pct: Optional[float] = None, + low_atr_pct: Optional[float] = None, + prob_long_low: Optional[float] = None, + cooldown_low: Optional[int] = None, + min_vol_24: Optional[float] = None, + atr_relax_pct: Optional[float] = None, + prob_long_relaxed: Optional[float] = None, + cooldown_relaxed: Optional[int] = None, + debug_log_hits: bool = False, + ) -> None: + super().__init__() + # Resolve model directory + if not model_dir: + ptr = Path(latest_ptr) + if not ptr.exists(): + raise RuntimeError(f"latest.json not found: {ptr}") + latest = json.loads(ptr.read_text(encoding="utf-8")) + model_dir = latest.get("model_dir") + if not model_dir: + raise RuntimeError("latest.json missing 'model_dir'") + self.model_dir = Path(model_dir) + # Load artifacts + self.feature_list: List[str] = json.loads((self.model_dir / "feature_list.json").read_text(encoding="utf-8")) + thr = json.loads((self.model_dir / "thresholds.json").read_text(encoding="utf-8")) + self.p_long = float(prob_long) if prob_long is not None else float(thr.get("p_long", 0.6)) + self.p_exit = float(prob_exit) if prob_exit is not None else float(thr.get("p_exit", 0.5)) + try: + import xgboost as xgb + except Exception as exc: + raise RuntimeError("xgboost is required at runtime for xgb_signal.") from exc + self._xgb = xgb + self._booster = xgb.Booster() + self._booster.load_model(str(self.model_dir / "model.json")) + + self.size_mult = float(size_mult) + self.cooldown_bars = int(max(0, cooldown_bars)) + self.cooldown_relaxed = int(max(0, cooldown_relaxed)) if cooldown_relaxed is not None else None + self.min_atr_pct = float(min_atr_pct) if min_atr_pct is not None else None + self.low_atr_pct = float(low_atr_pct) if low_atr_pct is not None else None + self.prob_long_low = float(prob_long_low) if prob_long_low is not None else None + self.cooldown_low = int(max(0, cooldown_low)) if cooldown_low is not None else None + self.min_vol_24 = float(min_vol_24) if min_vol_24 is not None else None + self.atr_relax_pct = float(atr_relax_pct) if atr_relax_pct is not None else None + self.p_long_relaxed = float(prob_long_relaxed) if prob_long_relaxed is not None else None + self.debug_log_hits = bool(debug_log_hits) + self._block_until: int = 0 + self._debug_max = 0.0 + self._debug_none = 0 + + def _note_feature_miss(self, reason: str) -> None: + self._debug_none += 1 + if self.debug_log_hits and self._debug_none <= 10: + logger.warning(f"[xgb_signal] feature unavailable ({reason})") + + def _features_from_state(self, state: Dict[str, Any]) -> tuple[Optional[np.ndarray], Optional[float]]: + # Close history for returns/volatility + ch = state.get("close_history") + if ch is None: + self._note_feature_miss("close_history missing") + return None, None + closes = np.asarray(ch, dtype=float) + if closes.size < 80: # need at least slow window context + self._note_feature_miss("insufficient history") + return None, None + close = float(state.get("close", closes[-1])) + + # Returns & rolling vol + def pct_change(arr: np.ndarray, k: int) -> float | None: + if arr.size <= k: + self._note_feature_miss(f"ret_{k} insufficient") + return None, None + a, b = arr[-k - 1], arr[-1] + return (b - a) / a if a else None + + ret_1 = pct_change(closes, 1) + ret_3 = pct_change(closes, 3) + ret_6 = pct_change(closes, 6) + vol_24 = None + if closes.size >= 25: + rets = np.diff(closes[-25:]) / closes[-25:-1] + vol_24 = float(np.std(rets)) if rets.size else None + + # SMA diff + sma_fast = _safe_get(state, "sma_fast") + sma_slow = _safe_get(state, "sma_slow") + if sma_fast is None or sma_slow is None: + sma_fast = float(np.mean(closes[-20:])) if closes.size >= 20 else None + sma_slow = float(np.mean(closes[-80:])) if closes.size >= 80 else None + if sma_fast is None or sma_slow is None: + self._note_feature_miss("sma missing") + return None, None + sma_diff = (float(sma_fast) - float(sma_slow)) / close if close else 0.0 + + # RSI (prefer engine state, else compute) + rsi_val = state.get("rsi") + if rsi_val is None: + r = _rsi_from_series(closes, 14) + rsi_val = r if r is not None else 50.0 + + # ATR normalized + curr_atr = state.get("curr_atr") + atr_norm = float(curr_atr) / close if (curr_atr is not None and close) else 0.0 + + # Time features from ts + ts = state.get("ts") + if ts is None: + self._note_feature_miss("timestamp missing") + return None, None + try: + import pandas as pd + ts_pd = pd.Timestamp(ts) + hour = float(ts_pd.hour) + dow = float(ts_pd.dayofweek) + except Exception: + self._note_feature_miss("timestamp parse") + return None, None + hour_sin, hour_cos = np.sin(2 * np.pi * hour / 24.0), np.cos(2 * np.pi * hour / 24.0) + dow_sin, dow_cos = np.sin(2 * np.pi * dow / 7.0), np.cos(2 * np.pi * dow / 7.0) + + feat_map = { + "ret_1": ret_1, + "ret_3": ret_3, + "ret_6": ret_6, + "vol_24": vol_24, + "sma_diff": sma_diff, + "rsi": float(rsi_val), + "atr_norm": atr_norm, + "hour_sin": float(hour_sin), + "hour_cos": float(hour_cos), + "dow_sin": float(dow_sin), + "dow_cos": float(dow_cos), + } + vec = [] + for name in self.feature_list: + val = feat_map.get(name) + if val is None or (isinstance(val, float) and (np.isnan(val) or np.isinf(val))): + self._note_feature_miss(f"feature {name} invalid") + return None, None + vec.append(float(val)) + return np.asarray(vec, dtype=float), float(vol_24) if vol_24 is not None else None + + def on_bar(self, state: Dict[str, Any]) -> Dict[str, Any]: + bar_idx = int(state.get("bar_idx", 0) or 0) + position_units = float(state.get("position_units", 0.0) or 0.0) + default_qty = float(state.get("default_qty", 0.0) or 0.0) + close = float(state.get("close", 0.0) or 0.0) + + atr_pct = state.get("atr_percentile") + if self.min_atr_pct is not None: + if atr_pct is None or float(atr_pct) < self.min_atr_pct: + return {"action": "HOLD"} + + # Determine per-bar entry threshold / cooldown after ATR gating + effective_prob_long = self.p_long + effective_cooldown = self.cooldown_bars + if ( + self.low_atr_pct is not None + and atr_pct is not None + and float(atr_pct) < self.low_atr_pct + ): + if self.prob_long_low is not None: + effective_prob_long = self.prob_long_low + if self.cooldown_low is not None: + effective_cooldown = self.cooldown_low + if ( + self.atr_relax_pct is not None + and atr_pct is not None + and float(atr_pct) >= self.atr_relax_pct + ): + if self.p_long_relaxed is not None: + effective_prob_long = self.p_long_relaxed + if self.cooldown_relaxed is not None: + effective_cooldown = self.cooldown_relaxed + + # Cooldown gate (updated per bar) + if bar_idx < self._block_until: + return {"action": "HOLD"} + + feats, vol_24 = self._features_from_state(state) + if feats is None: + return {"action": "HOLD"} + if self.min_vol_24 is not None: + if vol_24 is None or vol_24 < self.min_vol_24: + return {"action": "HOLD"} + + dmat = self._xgb.DMatrix(feats.reshape(1, -1), feature_names=self.feature_list) + p_up = float(self._booster.predict(dmat)[0]) + if p_up > self._debug_max: + self._debug_max = p_up + if self.debug_log_hits: + logger.info( + "[xgb_signal] new max prob %.4f (ts=%s close=%.5f position=%s)", + p_up, + state.get("ts"), + close, + position_units, + ) + + # Long-only logic + if position_units == 0.0: + if p_up >= effective_prob_long and default_qty > 0.0: + self._block_until = bar_idx + effective_cooldown + if self.debug_log_hits: + logger.info( + "[xgb_signal] ENTER signal p=%.4f (thr=%.4f) ts=%s", + p_up, + effective_prob_long, + state.get("ts"), + ) + return {"action": "ENTER_LONG", "size": default_qty * self.size_mult} + return {"action": "HOLD"} + else: + if p_up < self.p_exit: + if self.debug_log_hits: + logger.info( + "[xgb_signal] EXIT signal p=%.4f (thr=%.4f) ts=%s", + p_up, + self.p_exit, + state.get("ts"), + ) + return {"action": "EXIT_LONG"} + return {"action": "HOLD"}