diff --git a/backend/app/api/backtest.py b/backend/app/api/backtest.py index 7d89704..7ed9a7c 100644 --- a/backend/app/api/backtest.py +++ b/backend/app/api/backtest.py @@ -719,3 +719,158 @@ async def optimize_cancel(request: Request): return {"ok": True} return {"ok": False, "message": "任务不存在或已完成"} + +# ══════════════════════════════════════════════════════════════ +# Walk-forward 优化 — 每折训练区间优化 + 测试区间 OOS 验证 (复用优化器 + job_key 回吐) +# ══════════════════════════════════════════════════════════════ + +def _make_wf_job_key(strategy_id, symbols, start, end, param_grid, objective, direction, windows, bt_sig) -> str: + raw = f"WF|{strategy_id}|{symbols}|{start}|{end}|{param_grid}|{objective}|{direction}|{windows}|{bt_sig}" + return hashlib.md5(raw.encode()).hexdigest()[:12] + + +@router.get("/walkforward/stream") +async def walkforward_stream( + request: Request, + strategy_id: str, + param_grid: str, + objective: str = "sortino", + direction: str | None = None, + train_days: int = 252, + test_days: int = 63, + step_days: int = 63, + max_workers: int = 4, + symbols: str | None = None, + start: str | None = None, + end: str | None = None, + matching: str = "open_t+1", + fees_pct: float = 0.0002, + commission_pct: float | None = None, + stamp_tax_pct: float | None = None, + slippage_bps: float = 5.0, + max_positions: int = 10, + max_exposure_pct: float = 1.0, + initial_capital: float = 1_000_000.0, + position_sizing: str = "equal", + mode: str = "position", + holding_days: int = 5, +): + """SSE 流式 walk-forward: 每折训练区间网格优化 -> 测试区间 OOS 回测。 + + 事件: job {key} / progress {type:walkforward_progress,done,total,fold} / done {result} / error {message} + """ + from app.backtest.optimizer import StrategyOptimizer + from app.backtest.strategy import StrategyBacktestService + from app.backtest.walkforward import WalkForwardConfig, WalkForwardService + + direction = direction or None + engine = _get_engine(request) + strategy_engine = request.app.state.strategy_engine + svc = StrategyBacktestService(engine, strategy_engine) + optimizer = StrategyOptimizer(svc, strategy_engine) + + end_date = date.fromisoformat(end) if end else date.today() + if start: + start_date = date.fromisoformat(start) + else: + earliest = request.app.state.repo.earliest_daily_date() + start_date = earliest or (end_date - timedelta(days=STRATEGY_DEFAULT_DAYS)) + + bt_kwargs = _opt_backtest_kwargs( + matching, fees_pct, commission_pct, stamp_tax_pct, slippage_bps, + max_positions, max_exposure_pct, initial_capital, position_sizing, mode, holding_days, + ) + bt_sig = "|".join(f"{k}={bt_kwargs[k]}" for k in _OPT_BT_FIELDS) + windows = f"{train_days}/{test_days}/{step_days}" + job_key = _make_wf_job_key(strategy_id, symbols, start, end, param_grid, objective, direction, windows, bt_sig) + + _cleanup_stale_jobs() + with _jobs_lock: + job = _running_jobs.get(job_key) + if job is None: + job = _BacktestJob(job_key) + _running_jobs[job_key] = job + is_new = True + else: + is_new = False + + async def event_generator(): + yield f"event: job\ndata: {json.dumps({'key': job_key}, ensure_ascii=False)}\n\n" + + if is_new and not job.done: + try: + grid = json.loads(param_grid) + except (json.JSONDecodeError, TypeError): + grid = None + if not isinstance(grid, dict) or not grid: + job.error = "param_grid 必须是非空的参数网格对象" + job.done = True + job.finish_ts = time.time() + grid = None + + if grid is not None: + wf_cfg = WalkForwardConfig( + strategy_id=strategy_id, + symbols=[s.strip() for s in symbols.split(",") if s.strip()] if symbols else None, + start=start_date, + end=end_date, + param_grid=grid, + objective=objective, + direction=direction, + train_days=int(train_days), + test_days=int(test_days), + step_days=int(step_days), + max_workers=int(max_workers), + backtest_kwargs=bt_kwargs, + ) + + def _run_wf(): + try: + wf = WalkForwardService(optimizer, svc, strategy_engine) + job.result = wf.run(wf_cfg, lambda d: job.progress.append(d), job.cancel_event) + job.done = True + job.finish_ts = time.time() + except Exception as e: + job.error = str(e) + job.done = True + job.finish_ts = time.time() + + threading.Thread(target=_run_wf, daemon=True).start() + + cursor = 0 + tick = 0 + try: + while True: + if job.done: + if job.error: + yield f"event: error\ndata: {json.dumps({'message': job.error}, ensure_ascii=False)}\n\n" + elif job.cancel_event.is_set(): + yield f"event: error\ndata: {json.dumps({'message': 'walk-forward 已取消'}, ensure_ascii=False)}\n\n" + elif job.result is not None: + yield f"event: done\ndata: {json.dumps(job.result, ensure_ascii=False, default=str)}\n\n" + return + tick += 1 + if tick % 4 == 0 and await request.is_disconnected(): + break + while cursor < len(job.progress): + msg = job.progress[cursor] + cursor += 1 + yield f"event: progress\ndata: {json.dumps(msg, ensure_ascii=False, default=str)}\n\n" + await asyncio.sleep(0.5) + except asyncio.CancelledError: + raise + + return StreamingResponse(event_generator(), media_type="text/event-stream") + + +@router.post("/walkforward/cancel") +async def walkforward_cancel(request: Request): + """取消 walk-forward 任务 — 传 stream 首事件回吐的 job_key。""" + body = await request.json() + job_key = body.get("job_key", "") + job = _running_jobs.get(job_key) + if job and not job.done: + job.cancel_event.set() + return {"ok": True} + return {"ok": False, "message": "任务不存在或已完成"} + diff --git a/backend/app/backtest/walkforward.py b/backend/app/backtest/walkforward.py new file mode 100644 index 0000000..969521a --- /dev/null +++ b/backend/app/backtest/walkforward.py @@ -0,0 +1,216 @@ +"""Walk-forward 优化 — 滚动窗口的样本内优化 + 样本外验证。 + +每折在训练区间用参数网格优化选出最优参数, 再在紧邻的测试区间用该参数做样本外(OOS) +回测。滚动前移。核心产出是 OOS 拼接净值 + 每折 IS-vs-OOS 退化 —— 样本内漂亮、样本外 +崩溃即过拟合信号, 单次样本内回测看不到。 + +依赖 PR2a 的 StrategyOptimizer 做每折训练区间的网格优化。 +""" +from __future__ import annotations + +import logging +import time +from dataclasses import dataclass, field +from datetime import date, timedelta + +logger = logging.getLogger(__name__) + + +@dataclass +class Fold: + index: int + train_start: date + train_end: date + test_start: date + test_end: date + + +def generate_folds( + start: date, + end: date, + train_days: int, + test_days: int, + step_days: int, +) -> list[Fold]: + """滚动窗口 fold 切分: 训练窗口固定长度, 测试窗口紧接其后, 按 step 前移。 + + 测试区间超出 end 即停止。数据区间放不下一折则抛错。 + """ + if train_days <= 0 or test_days <= 0 or step_days <= 0: + raise ValueError("train_days / test_days / step_days 必须为正") + + folds: list[Fold] = [] + i = 0 + train_start = start + while True: + train_end = train_start + timedelta(days=train_days) + test_start = train_end + test_end = test_start + timedelta(days=test_days) + if test_end > end: + break + folds.append(Fold(i, train_start, train_end, test_start, test_end)) + i += 1 + train_start = train_start + timedelta(days=step_days) + + if not folds: + raise ValueError( + f"数据区间不足以切出至少一折 (需 train+test={train_days + test_days}天, " + f"实有 {(end - start).days}天)" + ) + return folds + + +def aggregate_oos(fold_records: list[dict], objective: str) -> dict: + """从各折 OOS 结果聚合: 复利净值曲线 / IS-OOS 退化 / 一致性。 + + fold_records: [{index, test_end, best_params, is_score, oos_stats}] + - compounded_oos_return: 各折 OOS 总收益复利 + - avg_is_objective / avg_oos_objective / degradation: IS 目标均值 - OOS 目标均值, + 正值 = 样本外退化 = 过拟合信号 + - consistency: OOS 目标 > 0 的折占比 + """ + n = len(fold_records) + if n == 0: + return { + "n_folds": 0, + "compounded_oos_return": 0.0, + "avg_is_objective": None, + "avg_oos_objective": None, + "degradation": None, + "consistency": 0.0, + "oos_equity_curve": [], + } + + equity = 1.0 + curve: list[dict] = [] + for f in fold_records: + r = float(f["oos_stats"].get("total_return", 0.0) or 0.0) + equity *= (1 + r) + curve.append({"fold": f["index"], "date": str(f["test_end"]), "value": round(equity, 4)}) + + is_vals = [f["is_score"] for f in fold_records if f["is_score"] is not None] + oos_vals = [f["oos_stats"].get(objective) for f in fold_records] + oos_vals = [v for v in oos_vals if v is not None] + + avg_is = round(float(sum(is_vals) / len(is_vals)), 4) if is_vals else None + avg_oos = round(float(sum(oos_vals) / len(oos_vals)), 4) if oos_vals else None + degradation = round(avg_is - avg_oos, 4) if (avg_is is not None and avg_oos is not None) else None + n_positive = sum(1 for v in oos_vals if v > 0) + consistency = round(n_positive / len(oos_vals), 4) if oos_vals else 0.0 + + return { + "n_folds": n, + "compounded_oos_return": round(equity - 1.0, 4), + "avg_is_objective": avg_is, + "avg_oos_objective": avg_oos, + "degradation": degradation, + "consistency": consistency, + "oos_equity_curve": curve, + } + + +@dataclass +class WalkForwardConfig: + strategy_id: str + symbols: list[str] | None + start: date + end: date + param_grid: dict + objective: str = "sortino" + train_days: int = 252 + test_days: int = 63 + step_days: int = 63 + direction: str | None = None + max_workers: int = 4 + base_params: dict = field(default_factory=dict) + overrides: dict | None = None + backtest_kwargs: dict = field(default_factory=dict) + + +class WalkForwardService: + """滚动窗口 walk-forward: 每折训练区间优化 -> 测试区间 OOS 验证 -> 聚合。""" + + def __init__(self, optimizer, service, strategy_engine) -> None: + self.optimizer = optimizer + self.service = service + self.strategy_engine = strategy_engine + + def run( + self, + cfg: WalkForwardConfig, + progress_cb=None, + cancel_event=None, + ) -> dict: + from app.backtest.optimizer import OptimizeConfig + from app.backtest.strategy import StrategyBacktestConfig + + t0 = time.perf_counter() + folds = generate_folds(cfg.start, cfg.end, cfg.train_days, cfg.test_days, cfg.step_days) + n_total = len(folds) + + fold_records: list[dict] = [] + for f in folds: + if cancel_event is not None and cancel_event.is_set(): + break + + # 训练区间: 网格优化选最优参数 + opt_cfg = OptimizeConfig( + strategy_id=cfg.strategy_id, + symbols=cfg.symbols, + start=f.train_start, + end=f.train_end, + param_grid=cfg.param_grid, + objective=cfg.objective, + direction=cfg.direction, + max_workers=cfg.max_workers, + base_params=cfg.base_params, + overrides=cfg.overrides, + backtest_kwargs=cfg.backtest_kwargs, + ) + opt_res = self.optimizer.optimize(opt_cfg, cancel_event=cancel_event) + best_params = opt_res.get("best_params") + is_score = opt_res.get("best_score") + + # 测试区间: 用最优参数做样本外回测 + merged = {**cfg.base_params, **(best_params or {})} + oos_cfg = StrategyBacktestConfig( + strategy_id=cfg.strategy_id, + symbols=cfg.symbols, + start=f.test_start, + end=f.test_end, + params=merged, + overrides=cfg.overrides, + **cfg.backtest_kwargs, + ) + oos_res = self.service.run(oos_cfg, cancel_event=cancel_event) + oos_stats = {} if oos_res.error else oos_res.stats + + fold_records.append({ + "index": f.index, + "train_start": str(f.train_start), + "train_end": str(f.train_end), + "test_start": str(f.test_start), + "test_end": str(f.test_end), + "best_params": best_params, + "is_score": is_score, + "oos_objective": oos_stats.get(cfg.objective), + "oos_stats": oos_stats, + }) + + if progress_cb is not None: + progress_cb({ + "type": "walkforward_progress", + "done": len(fold_records), + "total": n_total, + "fold": f.index, + }) + + summary = aggregate_oos(fold_records, cfg.objective) + return { + "objective": cfg.objective, + "n_folds": len(fold_records), + "n_planned_folds": n_total, + "folds": fold_records, + "summary": summary, + "elapsed_ms": round((time.perf_counter() - t0) * 1000, 1), + } diff --git a/backend/tests/backtest/test_walkforward.py b/backend/tests/backtest/test_walkforward.py new file mode 100644 index 0000000..9f2406f --- /dev/null +++ b/backend/tests/backtest/test_walkforward.py @@ -0,0 +1,209 @@ +"""Walk-forward 核心测试 — 滚动窗口 fold 生成 + OOS 聚合 + 编排。 + +被测: +- generate_folds: 滚动训练/测试窗口切分 +- aggregate_oos: 从各折 OOS 结果聚合 (复利净值/IS-OOS 退化/一致性) +- WalkForwardService.run: 每折 训练区间优化 -> 测试区间 OOS 验证 +""" +from __future__ import annotations + +from dataclasses import dataclass +from datetime import date + +import pytest + +from app.backtest.walkforward import ( + WalkForwardConfig, + WalkForwardService, + aggregate_oos, + generate_folds, +) + +# --------------------------------------------------------------- +# fold 生成 +# --------------------------------------------------------------- + +def test_folds_rolling_windows(): + # 1 年数据, 训练 90d / 测试 30d / 步进 30d + folds = generate_folds(date(2024, 1, 1), date(2024, 12, 31), train_days=90, test_days=30, step_days=30) + assert len(folds) > 0 + f0 = folds[0] + assert f0.train_start == date(2024, 1, 1) + assert f0.train_end == date(2024, 3, 31) # +90d (2024 闰年) + assert f0.test_start == date(2024, 3, 31) # 紧接训练 + assert f0.test_end == date(2024, 4, 30) # +30d + # 滚动: 下一折训练起点 +step + assert folds[1].train_start == date(2024, 1, 31) # +30d + + +def test_folds_no_test_beyond_end(): + folds = generate_folds(date(2024, 1, 1), date(2024, 12, 31), train_days=90, test_days=30, step_days=30) + for f in folds: + assert f.test_end <= date(2024, 12, 31) + + +def test_folds_insufficient_span_raises(): + # 训练90+测试30=120d, 但只有 100d 数据 -> 0 折 + with pytest.raises(ValueError, match=r"数据区间不足|至少"): + generate_folds(date(2024, 1, 1), date(2024, 4, 10), train_days=90, test_days=30, step_days=30) + + +def test_folds_reject_nonpositive_windows(): + with pytest.raises(ValueError, match=r"必须为正"): + generate_folds(date(2024, 1, 1), date(2024, 12, 31), train_days=0, test_days=30, step_days=30) + + +# --------------------------------------------------------------- +# OOS 聚合 +# --------------------------------------------------------------- + +def _rec(index, is_score, total_return, obj): + return { + "index": index, + "test_end": date(2024, 1, 1), + "best_params": {"p": index}, + "is_score": is_score, + "oos_stats": {"total_return": total_return, "sortino": obj}, + } + + +def test_aggregate_compounds_oos_returns(): + recs = [_rec(0, 2.0, 0.10, 1.5), _rec(1, 2.0, -0.05, 0.8), _rec(2, 2.0, 0.08, 1.2)] + agg = aggregate_oos(recs, objective="sortino") + # 复利: 1.1 * 0.95 * 1.08 - 1 + assert abs(agg["compounded_oos_return"] - (1.10 * 0.95 * 1.08 - 1)) < 1e-9 + assert len(agg["oos_equity_curve"]) == 3 + + +def test_aggregate_is_oos_degradation(): + # IS 目标平均远高于 OOS -> 退化为正 (过拟合信号) + recs = [_rec(0, 3.0, 0.05, 0.5), _rec(1, 3.0, 0.02, 0.3)] + agg = aggregate_oos(recs, objective="sortino") + assert agg["avg_is_objective"] == 3.0 + assert abs(agg["avg_oos_objective"] - 0.4) < 1e-9 + assert agg["degradation"] > 0 # IS 3.0 - OOS 0.4 = 2.6 + + +def test_aggregate_consistency_fraction_positive(): + # 3 折 OOS sortino: 1.5>0, -0.2<=0, 0.8>0 -> 2/3 正 + recs = [_rec(0, 1, 0.1, 1.5), _rec(1, 1, -0.1, -0.2), _rec(2, 1, 0.1, 0.8)] + agg = aggregate_oos(recs, objective="sortino") + assert agg["consistency"] == round(2 / 3, 4) # 0.6667 + + +def test_aggregate_empty_folds(): + agg = aggregate_oos([], objective="sortino") + assert agg["n_folds"] == 0 + assert agg["compounded_oos_return"] == 0.0 + + +# --------------------------------------------------------------- +# 编排 (假 optimizer / service) +# --------------------------------------------------------------- + +@dataclass +class _FakeResult: + stats: dict + error: str | None = None + + +class _FakeOptimizer: + """optimize 返回受控 best_params/best_score, 记录被优化的训练区间。""" + def __init__(self): + self.train_ranges = [] + + def optimize(self, cfg, progress_cb=None, cancel_event=None): + self.train_ranges.append((cfg.start, cfg.end)) + # best_params 随训练起点变化, best_score 固定 + return {"best_params": {"p": cfg.start.month}, "best_score": 2.0, "results": [], "n_completed": 1} + + +class _FakeService: + """run 返回受控 OOS stats, 记录测试区间 + 收到的 params。""" + def __init__(self): + self.calls = [] + + def run(self, config, progress_cb=None, cancel_event=None): + self.calls.append({"start": config.start, "end": config.end, "params": dict(config.params or {})}) + return _FakeResult(stats={"total_return": 0.05, "sortino": 1.0}) + + +def _wf_cfg(**kw): + base = dict( + strategy_id="s", symbols=None, start=date(2024, 1, 1), end=date(2024, 12, 31), + param_grid={"p": [1, 2]}, objective="sortino", + train_days=90, test_days=30, step_days=30, + ) + base.update(kw) + return WalkForwardConfig(**base) + + +def test_walkforward_optimizes_train_applies_oos(): + opt, svc = _FakeOptimizer(), _FakeService() + wf = WalkForwardService(opt, svc, strategy_engine=None) + out = wf.run(_wf_cfg()) + + assert out["n_folds"] > 0 + # 每折: optimizer 在训练区间跑, service 在测试区间用最优参数跑 + assert len(opt.train_ranges) == out["n_folds"] + assert len(svc.calls) == out["n_folds"] + # OOS 回测用的是该折优化出的 best_params (来自训练起点月份) + first_fold = out["folds"][0] + assert svc.calls[0]["params"] == first_fold["best_params"] + # 训练区间与测试区间不重叠 (测试在训练之后) + assert svc.calls[0]["start"] >= opt.train_ranges[0][1] + + +def test_walkforward_reports_degradation(): + opt, svc = _FakeOptimizer(), _FakeService() + wf = WalkForwardService(opt, svc, strategy_engine=None) + out = wf.run(_wf_cfg()) + # IS best_score=2.0, OOS sortino=1.0 -> 退化 1.0 + assert out["summary"]["avg_is_objective"] == 2.0 + assert out["summary"]["avg_oos_objective"] == 1.0 + assert abs(out["summary"]["degradation"] - 1.0) < 1e-9 + + +def test_walkforward_cancel_stops(): + import threading + ev = threading.Event() + ev.set() + opt, svc = _FakeOptimizer(), _FakeService() + wf = WalkForwardService(opt, svc, strategy_engine=None) + out = wf.run(_wf_cfg(), cancel_event=ev) + # 取消 -> 不跑任何折 + assert svc.calls == [] + assert out["n_folds"] == 0 + + +# --------------------------------------------------------------- +# API: job_key 回吐 + cancel 按 key 查表 +# --------------------------------------------------------------- + +def test_wf_job_key_distinguishes_windows(): + from app.api.backtest import _make_wf_job_key + base = _make_wf_job_key("s", None, None, None, '{"p":[1]}', "sortino", None, "252/63/63", "sig") + assert base != _make_wf_job_key("s", None, None, None, '{"p":[1]}', "sortino", None, "120/30/30", "sig") + + +def test_wf_cancel_by_echoed_key(): + import asyncio + + from app.api.backtest import _BacktestJob, _running_jobs, walkforward_cancel + + class _Req: + def __init__(self, body): + self._body = body + async def json(self): + return self._body + + key = "wfkey_test_1" + _running_jobs[key] = _BacktestJob(key) + try: + res = asyncio.run(walkforward_cancel(_Req({"job_key": key}))) + assert res["ok"] is True + assert _running_jobs[key].cancel_event.is_set() + res2 = asyncio.run(walkforward_cancel(_Req({"job_key": "nope"}))) + assert res2["ok"] is False + finally: + _running_jobs.pop(key, None)