mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 19:04:15 +08:00
PR2b 后端 (叠在 PR2a 优化器上): app/backtest/walkforward.py: - generate_folds: 滚动窗口切分 (训练固定长度 + 紧邻测试窗口 + step 前移), 测试超出 end 即停; 放不下一折抛错; 非正窗口拒绝。 - aggregate_oos: 从各折 OOS 结果聚合 — 复利净值曲线 / IS-vs-OOS 退化 (样本内目标均值 - 样本外目标均值, 正值=过拟合信号) / 一致性 (OOS 目标为正的折占比)。 - WalkForwardService.run: 每折在训练区间调 PR2a 优化器选最优参数, 再在测试区间 用该参数做 OOS 回测。核心产出是纯 OOS 拼接 + 每折 IS/OOS 对比 — 单次样本内 回测的过拟合一眼看穿。支持进度回调 (fold i/N) 与 cancel。 API app/api/backtest.py: - GET /walkforward/stream (SSE, 复用 _BacktestJob + job_key 回吐) + POST /walkforward/cancel (按回吐 key 查表)。 测试 13 例: fold 切分 (滚动/边界/不足/非正) + OOS 聚合 (复利/退化/一致性/空) + 编排 (训练优化→测试OOS/退化上报/取消) + API (job_key 区分窗口/cancel按key)。 全量 169 测试通过。
210 lines
7.4 KiB
Python
210 lines
7.4 KiB
Python
"""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)
|