Add files via upload

This commit is contained in:
xiaochuan
2025-11-14 23:16:51 +00:00
committed by GitHub
parent 8c9371683e
commit 53c7aa9182
85 changed files with 6789 additions and 0 deletions
+41
View File
@@ -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 推送(Phase4**
```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`。
+12
View File
@@ -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
@@ -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()
+123
View File
@@ -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()
+146
View File
@@ -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()
+828
View File
@@ -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()
+110
View File
@@ -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()
+42
View File
@@ -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()
+17
View File
@@ -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()
+82
View File
@@ -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()
+50
View File
@@ -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}")
+65
View File
@@ -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()
+200
View File
@@ -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}")
+218
View File
@@ -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()
+25
View File
@@ -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()
+36
View File
@@ -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")
+138
View File
@@ -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}")
+63
View File
@@ -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()
+107
View File
@@ -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()
+268
View File
@@ -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()
+35
View File
@@ -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"
+181
View File
@@ -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()
+55
View File
@@ -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
+359
View File
@@ -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()
+99
View File
@@ -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
+277
View File
@@ -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()
+234
View File
@@ -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()
+206
View File
@@ -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()
+101
View File
@@ -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()
+79
View File
@@ -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()
+80
View File
@@ -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()
+56
View File
@@ -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()
+37
View File
@@ -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)
+63
View File
@@ -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"}
+16
View File
@@ -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"}
+33
View File
@@ -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
+94
View File
@@ -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
+17
View File
@@ -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
+104
View File
@@ -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
+246
View File
@@ -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
+110
View File
@@ -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"}
+274
View File
@@ -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"}