Files
easy_tdx_max/tests/unit/test_backtest_fitness_benchmark.py
T
GitHub b2509a8d0e release: v1.25.0 — Walk-Forward/适配性/一条龙评估防过拟合链 + 评分评级后端化 + 寻优加速
升级计划 P1:补上两个下游项目都在自研的样本外验证空白。

- Walk-Forward 引擎(walkforward.py):7 窗样本外、每窗独立开仓(backtest-system v1.2.1 踩坑语义)、
  上下文预热不污染;CLI --wf、REST /backtest/wf/run/async
- 适配性评估(fitness.py):train/valid/test 三段 + 8 项可解释检查 + 高适配标记;
  evaluate_prefix 滚动过滤原语(无未来泄漏)
- 一条龙评估(benchmark.py evaluate_strategy):回测+WF+适配性+评分+评级+买入持有基准对比;
  CLI --evaluate、REST /backtest/evaluate/run/async
- 综合评分(scoring.py,收益50/夏普15/回撤10/Sortino5/WF20)+ 评级后端化(grading.py,
  前端 TS 忠实移植,REST 响应新增 grade/score 字段)
- 多 seed 验证 + 四项晋级门槛(validation.py);REST /backtest/multiseed/run/async
- 寻优两段式加速:IndicatorCache(36 点网格命中率 41.7%)+ workers 进程并行(实测约 2x);
  诚实注:指标缓存墙钟 ~1.01x,瓶颈在逐 bar 循环,后续向量化
- strategy.I() 指标缓存钩子 + 数据代理零拷贝(astype copy=False);
  types.to_json_native 统一 numpy 清洗
2026-09-01 22:17:08 +08:00

197 lines
6.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""适配性评估(fitness+ 一条龙评估(benchmark)测试。"""
from __future__ import annotations
import json
import numpy as np
import pandas as pd
import pytest
from easy_tdx.backtest.benchmark import evaluate_strategy, run_buy_hold_benchmark
from easy_tdx.backtest.fitness import FitnessEngine, rolling_fitness_scores
from easy_tdx.backtest.strategy import Strategy
class _BuyFirstBar(Strategy):
def init(self) -> None:
self._bought = False
def next(self) -> None:
if not self._bought:
self.buy()
self._bought = True
class _CycleTrader(Strategy):
"""每 10 根切换一次持仓(买卖交替),保证各段有完整回合(total_trades>0)。"""
def init(self) -> None:
self._count = 0
self._holding = False
def next(self) -> None:
self._count += 1
if self._count % 10 == 0:
if self._holding:
self.sell()
self._holding = False
else:
self.buy()
self._holding = True
def _df(n: int = 500, drift: float = 0.004) -> pd.DataFrame:
rng = np.random.default_rng(3)
dates = pd.date_range("2018-01-01", periods=n, freq="B")
close = 10.0 * np.cumprod(1.0 + drift + rng.normal(0, 0.006, n))
return pd.DataFrame(
{
"datetime": dates,
"open": close * 0.999,
"high": close * 1.01,
"low": close * 0.99,
"close": close,
"vol": 1000.0,
}
)
# ── FitnessEngine ─────────────────────────────────────────────────────────────
def test_fitness_three_segments_and_checks():
rep = FitnessEngine(_CycleTrader).evaluate(_df(600))
assert [s.name for s in rep.segments] == ["train", "valid", "test"]
assert len(rep.checks) == 8
names = {c.name for c in rep.checks}
assert names == {
"train_profitable",
"valid_profitable",
"test_profitable",
"sign_consistent",
"drawdown_bounded",
"train_enough_trades",
"test_active",
"oos_sharpe_positive",
}
# 上涨行情 + 买入持有 → 大部分检查通过 → 高适配
assert rep.pass_ratio >= 0.75
assert rep.high_fitness
def test_fitness_checks_carry_values():
rep = FitnessEngine(_CycleTrader).evaluate(_df(600))
for c in rep.checks:
assert c.detail # 每条检查附实际值(可解释性)
assert isinstance(c.passed, bool)
def test_fitness_losing_market_fails():
rep = FitnessEngine(_BuyFirstBar).evaluate(_df(600, drift=-0.002))
assert rep.pass_ratio < 0.75
assert not rep.high_fitness
# 三段全亏 → sign_consistent 通过(同号),但盈利检查全挂
by_name = {c.name: c.passed for c in rep.checks}
assert by_name["train_profitable"] is False
assert by_name["valid_profitable"] is False
assert by_name["test_profitable"] is False
def test_fitness_insufficient_data_returns_empty():
rep = FitnessEngine(_BuyFirstBar).evaluate(_df(60)) # valid 段 = 12 根 < 20 → 空报告
assert rep.segments == []
assert rep.checks == []
assert rep.high_fitness is False
def test_fitness_invalid_split_raises():
with pytest.raises(ValueError, match="split"):
FitnessEngine(_BuyFirstBar, split=(0.5, 0.2, 0.2))
def test_fitness_prefix_no_lookahead():
"""evaluate_prefix 只用前缀:末段测试段终点必须早于 end_index。"""
df = _df(600)
rep = FitnessEngine(_BuyFirstBar).evaluate_prefix(df, 400)
assert [s.name for s in rep.segments] == ["train", "valid", "test"]
# 前缀评估的测试段末日期 < 第 400 根的日期
dt_col = "datetime"
cutoff = pd.Timestamp(df[dt_col].iloc[399]).strftime("%Y-%m-%d")
assert rep.segments[-1].end <= cutoff
def test_rolling_fitness_scores_series():
df = _df(700)
scores = rolling_fitness_scores(df, _BuyFirstBar, step=100, min_prefix=300)
assert len(scores) >= 3
assert all(s["index"] < 700 for s in scores)
assert all(0.0 <= s["pass_ratio"] <= 1.0 for s in scores)
# 时间升序
idxs = [s["index"] for s in scores]
assert idxs == sorted(idxs)
def test_fitness_report_serializable():
rep = FitnessEngine(_BuyFirstBar).evaluate(_df(500))
d = rep.to_dict()
json.dumps(d, default=str)
assert d["total_checks"] == 8
assert "high_fitness" in d
# ── benchmark(一条龙评估)────────────────────────────────────────────────────
def test_buy_hold_benchmark_matches_trend():
df = _df(300, drift=0.002)
bh = run_buy_hold_benchmark(df)
total = df["close"].iloc[-1] / df["close"].iloc[0] - 1
assert bh["total_return"] == pytest.approx(total, rel=0.05) # 扣少量费用
def test_evaluate_strategy_full_report_structure():
report = evaluate_strategy(_BuyFirstBar, _df(500))
for key in (
"performance",
"score",
"grade",
"walkforward",
"fitness",
"benchmark",
"config",
):
assert key in report
# 绩效 19 项
assert "total_return" in report["performance"]
assert "sharpe" in report["performance"]
# 评分/评级结构
assert 0 <= report["score"]["total"] <= 100
assert report["grade"]["grade"] in ("S", "A", "B", "C", "D")
# WF
assert report["walkforward"]["n_windows"] == 7
# 适配性
assert report["fitness"]["total_checks"] == 8
# 基准
assert "buy_hold" in report["benchmark"]
assert "excess_return" in report["benchmark"]
def test_evaluate_strategy_excess_return_sign():
"""上涨行情 + 买入持有策略 ≈ 基准本身,excess_return 接近 0(扣费差异)。"""
report = evaluate_strategy(_BuyFirstBar, _df(400))
excess = report["benchmark"]["excess_return"]
assert abs(excess) < 0.05
def test_evaluate_strategy_serializable():
report = evaluate_strategy(_BuyFirstBar, _df(300), n_windows=3)
text = json.dumps(report, default=str)
assert "excess_return" in text
def test_evaluate_strategy_auto_fees_for_etf():
report = evaluate_strategy(_BuyFirstBar, _df(300), symbol="SH:510300", auto_fees=True)
assert report["config"]["symbol"] == "SH:510300"
assert report["config"]["auto_fees"] is True