mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 15:34:16 +08:00
feat(walkforward): 真·walk-forward 优化后端 — 滚动窗口 IS 优化 + OOS 验证
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 测试通过。
This commit is contained in:
@@ -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": "任务不存在或已完成"}
|
||||
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user