mirror of
https://github.com/xavierchuan/FX-ML-Trading-Engine.git
synced 2026-07-27 18:17:44 +00:00
Add files via upload
This commit is contained in:
@@ -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/<run_id>/summary.json` 记录本次运行的 KPI + 数据签名,提交代码时请一并引用该 run_id 便于审核。
|
||||
- 提交前可执行 `python scripts/validate_results.py results/<run_id>`,快速检查 KPI 字段、数据报告引用是否完整。
|
||||
|
||||
## 风控与 Diagnostics 流程
|
||||
|
||||
1. **风险仿真守门**
|
||||
```bash
|
||||
cd QuantResearch
|
||||
RUN=<run_id> ./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_<timestamp>/`,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`。
|
||||
@@ -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
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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()
|
||||
@@ -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()
|
||||
@@ -0,0 +1,146 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Backfill results/risk/metrics.csv using historical execution runs.
|
||||
|
||||
For each run under results/execution/<run_id>/, 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()
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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.")
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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}")
|
||||
@@ -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()
|
||||
@@ -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="<green>{time:YYYY-MM-DD HH:mm:ss}</green> | <level>{level}</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()
|
||||
@@ -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}")
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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/<symbol>_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()
|
||||
@@ -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()
|
||||
@@ -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")
|
||||
@@ -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}")
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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"
|
||||
@@ -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/<run_id> 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()
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
@@ -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/<ts>/
|
||||
- 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/<ts>"}
|
||||
"""
|
||||
|
||||
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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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/<run_id> directory.")
|
||||
parser.add_argument("path", help="Path to results/<run_id> 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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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"}
|
||||
@@ -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"}
|
||||
@@ -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"}
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"}
|
||||
@@ -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"}
|
||||
Reference in New Issue
Block a user