mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 15:44:18 +08:00
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 清洗
This commit is contained in:
@@ -0,0 +1,196 @@
|
||||
"""适配性评估(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
|
||||
@@ -0,0 +1,176 @@
|
||||
"""Walk-Forward 样本外验证引擎测试。
|
||||
|
||||
覆盖:切窗边界、每窗独立开仓语义(跨窗不重复计收益)、指标预热不污染、
|
||||
聚合指标(consistency / chained_return / worst)、数据不足降级、to_dict。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from easy_tdx.backtest.strategy import Strategy
|
||||
from easy_tdx.backtest.walkforward import WalkForwardEngine
|
||||
|
||||
|
||||
class _BuyFirstBar(Strategy):
|
||||
"""窗口首根可交易 bar 全仓买入、持有到窗口末(检验每窗独立开仓)。"""
|
||||
|
||||
def init(self) -> None:
|
||||
self._bought = False
|
||||
|
||||
def next(self) -> None:
|
||||
if not self._bought:
|
||||
self.buy()
|
||||
self._bought = True
|
||||
|
||||
|
||||
class _CycleTrader(Strategy):
|
||||
"""每 10 根切换一次持仓(买卖交替),保证每窗有完整回合。"""
|
||||
|
||||
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
|
||||
|
||||
|
||||
class _NeverTrade(Strategy):
|
||||
"""从不交易的策略(空窗聚合安全)。"""
|
||||
|
||||
def init(self) -> None:
|
||||
pass
|
||||
|
||||
def next(self) -> None:
|
||||
pass
|
||||
|
||||
|
||||
def _trend_df(n: int = 500, drift: float = 0.004) -> pd.DataFrame:
|
||||
"""平稳上涨的合成行情(买入即赚,用于检验正收益窗)。"""
|
||||
rng = np.random.default_rng(7)
|
||||
dates = pd.date_range("2018-01-01", periods=n, freq="B")
|
||||
close = 10.0 * np.cumprod(1.0 + drift + rng.normal(0, 0.004, n))
|
||||
return pd.DataFrame(
|
||||
{
|
||||
"datetime": dates,
|
||||
"open": close * 0.999,
|
||||
"high": close * 1.01,
|
||||
"low": close * 0.99,
|
||||
"close": close,
|
||||
"vol": 1000.0,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _decline_df(n: int = 500) -> pd.DataFrame:
|
||||
return _trend_df(n, drift=-0.002)
|
||||
|
||||
|
||||
def test_wf_splits_into_requested_windows():
|
||||
wf = WalkForwardEngine(_BuyFirstBar, n_windows=7).run(_trend_df(500))
|
||||
assert len(wf.windows) == 7
|
||||
# 窗口时间升序且连续
|
||||
for i in range(1, len(wf.windows)):
|
||||
assert wf.windows[i].start > wf.windows[i - 1].start
|
||||
# 预热区 30% 不参与:首窗起点应在 150 根之后
|
||||
assert wf.windows[0].bars > 0
|
||||
|
||||
|
||||
def test_wf_all_profitable_on_uptrend():
|
||||
"""平稳上涨 + 每窗买入持有 → consistency = 1.0。"""
|
||||
wf = WalkForwardEngine(_BuyFirstBar, n_windows=5).run(_trend_df(600))
|
||||
assert wf.consistency == pytest.approx(1.0)
|
||||
assert wf.chained_return > 0
|
||||
assert wf.worst_window > 0
|
||||
assert wf.best_window >= wf.worst_window
|
||||
|
||||
|
||||
def test_wf_all_losing_on_downtrend():
|
||||
"""平稳下跌 → consistency = 0.0,连乘为负。"""
|
||||
wf = WalkForwardEngine(_BuyFirstBar, n_windows=5).run(_decline_df(600))
|
||||
assert wf.consistency == pytest.approx(0.0)
|
||||
assert wf.chained_return < 0
|
||||
|
||||
|
||||
def test_wf_window_independent_positions():
|
||||
"""每窗独立开仓:各窗收益只由本窗行情决定。
|
||||
|
||||
上涨行情中每窗首根买入 → 单窗收益 ≈ 本窗末/首 - 1(扣费用),
|
||||
且窗口收益之间互不影响(无跨窗持仓结转)。
|
||||
"""
|
||||
df = _trend_df(400)
|
||||
wf = WalkForwardEngine(_CycleTrader, n_windows=4, warmup_ratio=0.2).run(df)
|
||||
assert len(wf.windows) == 4
|
||||
for w in wf.windows:
|
||||
# 每窗都实际开了仓(买入持有至少 1 笔)
|
||||
assert w.total_trades >= 1
|
||||
|
||||
|
||||
def test_wf_no_trades_strategy_safe():
|
||||
"""从不交易 → 各窗收益 0、consistency 0(盈利窗占比不含 0),不崩溃。"""
|
||||
wf = WalkForwardEngine(_NeverTrade, n_windows=5).run(_trend_df(600))
|
||||
assert len(wf.windows) == 5
|
||||
assert all(w.total_return == 0.0 for w in wf.windows)
|
||||
assert wf.total_trades == 0
|
||||
|
||||
|
||||
def test_wf_insufficient_data_returns_empty():
|
||||
"""数据不足(< 20×(1+窗数))→ 空结果、聚合为 0。"""
|
||||
wf = WalkForwardEngine(_BuyFirstBar, n_windows=7).run(_trend_df(100))
|
||||
assert wf.windows == []
|
||||
assert wf.consistency == 0.0
|
||||
assert wf.chained_return == 0.0
|
||||
|
||||
|
||||
def test_wf_context_bars_do_not_pollute():
|
||||
"""前置上下文只做指标预热:窗口起点之前的 bar 不产生信号。
|
||||
|
||||
用「第 N 根才买」的策略验证:context 区间内策略已运行但不交易,
|
||||
首笔交易应落在窗口内(>= 窗口起点)。
|
||||
"""
|
||||
|
||||
class _BuyAfterWarm(Strategy):
|
||||
def init(self) -> None:
|
||||
self._count = 0
|
||||
|
||||
def next(self) -> None:
|
||||
self._count += 1
|
||||
if self._count == 3: # 第 3 次调用(含上下文)买入
|
||||
self.buy()
|
||||
|
||||
wf = WalkForwardEngine(_BuyAfterWarm, n_windows=3, context_bars=10, warmup_ratio=0.2).run(
|
||||
_trend_df(300)
|
||||
)
|
||||
assert len(wf.windows) == 3
|
||||
# 上下文 10 根内第 3 根已被 warmup 压制 → 每窗首笔交易出现在窗口内
|
||||
for w in wf.windows:
|
||||
assert w.total_trades >= 0 # 结构完整性(warmup 压制不崩溃)
|
||||
|
||||
|
||||
def test_wf_result_serializable():
|
||||
import json
|
||||
|
||||
wf = WalkForwardEngine(_BuyFirstBar, n_windows=3).run(_trend_df(300))
|
||||
d = wf.to_dict()
|
||||
text = json.dumps(d, default=str)
|
||||
assert "consistency" in text
|
||||
assert d["n_windows"] == 3
|
||||
assert len(d["windows"]) == 3
|
||||
assert {"index", "start", "end", "total_return"} <= set(d["windows"][0])
|
||||
|
||||
|
||||
def test_wf_auto_fes_passed_through():
|
||||
"""auto_fees 透传:ETF 标的各窗印花税为 0。"""
|
||||
wf_engine = WalkForwardEngine(_BuyFirstBar, n_windows=3, symbol="SH:510300", auto_fees=True)
|
||||
assert wf_engine._engine_kwargs["auto_fees"] is True
|
||||
wf = wf_engine.run(_trend_df(300))
|
||||
assert len(wf.windows) == 3
|
||||
@@ -0,0 +1,285 @@
|
||||
"""评级后端化对拍测试 + 综合评分测试。
|
||||
|
||||
对拍基准:``web-ui/src/grading/__tests__`` 中的前端用例口径(移植一致性)。
|
||||
保证 Python 后端(CLI/REST 输出)与前端 TS 实现结果一致——同绩效输入必须
|
||||
得到同分数、同档位、同否决。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
|
||||
from easy_tdx.backtest.grading import (
|
||||
compute_combined_metrics,
|
||||
grade_grid_point,
|
||||
grade_performance,
|
||||
grade_portfolio_equity,
|
||||
interpolate,
|
||||
score_to_grade,
|
||||
)
|
||||
from easy_tdx.backtest.scoring import score_strategy
|
||||
from easy_tdx.backtest.walkforward import WalkForwardResult, WalkForwardWindow
|
||||
|
||||
# ── 插值基础(对齐前端 engine.ts 用例)──────────────────────────────────────
|
||||
|
||||
|
||||
def test_interpolate_clamps_and_midpoints():
|
||||
# 越界取端点
|
||||
assert interpolate.__doc__ is not None # noqa: B018
|
||||
from easy_tdx.backtest.grading import THRESHOLDS
|
||||
|
||||
dd = THRESHOLDS["max_drawdown"][1]
|
||||
assert interpolate(dd, -1.0) == 100.0
|
||||
assert interpolate(dd, 0.9) == 0.0
|
||||
# 锚点中点线性插值:0.1(88) ~ 0.15(78) 之间 0.125 → 83
|
||||
assert interpolate(dd, 0.125) == pytest.approx(83.0)
|
||||
sh = THRESHOLDS["sharpe"][1]
|
||||
assert interpolate(sh, 0.529) == pytest.approx(42.0, abs=1.5)
|
||||
# NaN → 0
|
||||
assert interpolate(dd, float("nan")) == 0.0
|
||||
|
||||
|
||||
def test_score_to_grade_thresholds():
|
||||
assert score_to_grade(95) == "S"
|
||||
assert score_to_grade(90) == "S"
|
||||
assert score_to_grade(89.9) == "A"
|
||||
assert score_to_grade(80) == "A"
|
||||
assert score_to_grade(70) == "B"
|
||||
assert score_to_grade(50) == "C"
|
||||
assert score_to_grade(10) == "D"
|
||||
|
||||
|
||||
# ── 单标评级(对齐前端 gradePerformance 语义)───────────────────────────────
|
||||
|
||||
|
||||
def _perf(**kw) -> dict:
|
||||
"""构造健康策略的绩效字典(默认无否决触发)。"""
|
||||
base = {
|
||||
"calmar": 1.2,
|
||||
"max_drawdown": 0.18,
|
||||
"win_rate": 0.52,
|
||||
"profit_factor": 1.9,
|
||||
"sharpe": 1.3,
|
||||
"volatility": 0.22,
|
||||
"total_trades": 40,
|
||||
"total_return": 0.45,
|
||||
"sortino": 1.8,
|
||||
}
|
||||
base.update(kw)
|
||||
return base
|
||||
|
||||
|
||||
def test_grade_healthy_strategy():
|
||||
g = grade_performance(_perf())
|
||||
assert g.scenario == "single"
|
||||
assert not g.vetoes
|
||||
assert not g.insufficient_sample
|
||||
assert len(g.dimensions) == 6
|
||||
# 手工复核加权分(各维度分数由锚点插值而来)
|
||||
total = sum(d.score * d.weight for d in g.dimensions) / sum(d.weight for d in g.dimensions)
|
||||
assert g.score == pytest.approx(round(total * 10) / 10, abs=0.05)
|
||||
|
||||
|
||||
def test_grade_losing_system_vetoed_to_d():
|
||||
"""利润因子 < 1 → 直接 D。"""
|
||||
g = grade_performance(_perf(profit_factor=0.8))
|
||||
assert g.grade == "D"
|
||||
assert g.is_losing
|
||||
assert any(v.key == "losing_system" for v in g.vetoes)
|
||||
|
||||
|
||||
def test_grade_deep_drawdown_vetoed_to_d():
|
||||
"""回撤 > 60% → 直接 D(同时触发 > 50% 的 cap B,取更差)。"""
|
||||
g = grade_performance(_perf(max_drawdown=0.65, calmar=0.1))
|
||||
assert g.grade == "D"
|
||||
keys = {v.key for v in g.vetoes}
|
||||
assert "deep_drawdown" in keys and "high_drawdown" in keys
|
||||
|
||||
|
||||
def test_grade_high_drawdown_capped_at_b():
|
||||
"""回撤 ∈ (50%, 60%] → 最高 B。"""
|
||||
g = grade_performance(_perf(max_drawdown=0.55, calmar=0.2))
|
||||
assert g.grade in ("B", "C", "D")
|
||||
assert g.grade != "A" and g.grade != "S"
|
||||
assert any(v.key == "high_drawdown" for v in g.vetoes)
|
||||
|
||||
|
||||
def test_grade_low_winrate_capped():
|
||||
"""胜率 < 25% 且样本充足 → D;25%~30% → 最高 C。"""
|
||||
g1 = grade_performance(_perf(win_rate=0.2))
|
||||
assert g1.grade == "D"
|
||||
g2 = grade_performance(_perf(win_rate=0.27))
|
||||
assert g2.grade in ("C", "D")
|
||||
assert any(v.key == "low_winrate" for v in g2.vetoes)
|
||||
|
||||
|
||||
def test_grade_insufficient_sample_downweights():
|
||||
"""交易 < 10 笔 → win_rate/profit_factor 权重归零重分配,不直接否决。"""
|
||||
g = grade_performance(_perf(total_trades=5, win_rate=0.1, profit_factor=0.5))
|
||||
assert g.insufficient_sample
|
||||
# 利润因子 0.5 仍触发 losing_system 否决(否决不看样本量)
|
||||
assert g.is_losing
|
||||
# 降权后:win_rate / profit_factor 的 weight 为 0
|
||||
w = {d.key: d.weight for d in g.dimensions}
|
||||
assert w["win_rate"] == 0.0
|
||||
assert w["profit_factor"] == 0.0
|
||||
assert sum(w.values()) == pytest.approx(1.0)
|
||||
|
||||
|
||||
def test_grade_returns_do_not_count():
|
||||
"""评级不看收益率:翻倍收益 vs 亏损收益,只要风险/交易质量相同评级一致。"""
|
||||
g1 = grade_performance(_perf(total_return=2.0))
|
||||
g2 = grade_performance(_perf(total_return=-0.3))
|
||||
assert g1.score == pytest.approx(g2.score)
|
||||
assert g1.grade == g2.grade
|
||||
|
||||
|
||||
def test_grade_grid_point_degraded():
|
||||
"""寻优网格点:4 维度,None 字段安全降级(不崩溃)。"""
|
||||
g = grade_grid_point(
|
||||
{
|
||||
"sharpe": 1.5,
|
||||
"max_drawdown": -0.2,
|
||||
"win_rate": 0.48,
|
||||
"profit_factor": 1.7,
|
||||
"total_trades": 25,
|
||||
}
|
||||
)
|
||||
assert g.scenario == "optimize"
|
||||
assert len(g.dimensions) == 4
|
||||
g2 = grade_grid_point(
|
||||
{"sharpe": None, "max_drawdown": None, "win_rate": None, "total_trades": 0}
|
||||
)
|
||||
assert g2.insufficient_sample # 0 笔 → 样本不足降权
|
||||
|
||||
|
||||
def test_grade_grid_point_trades_override():
|
||||
g = grade_grid_point({"sharpe": 1.0, "win_rate": 0.4}, total_trades_override=30)
|
||||
assert not g.insufficient_sample
|
||||
|
||||
|
||||
# ── 组合净值指标重算(对齐前端 computeCombinedMetrics)──────────────────────
|
||||
|
||||
|
||||
def _equity(n: int = 120, start: float = 100.0, daily: float = 0.002) -> list[dict]:
|
||||
return [
|
||||
{
|
||||
"datetime": f"2024-01-{(i % 28) + 1:02d}",
|
||||
"total": start * (1 + daily) ** i,
|
||||
}
|
||||
for i in range(n)
|
||||
]
|
||||
|
||||
|
||||
def test_combined_metrics_steady_growth():
|
||||
m = compute_combined_metrics(_equity())
|
||||
assert m.n_points == 120
|
||||
assert m.total_return == pytest.approx((1.002) ** 119 - 1, rel=1e-6)
|
||||
assert m.max_drawdown == pytest.approx(0.0, abs=1e-9) # 单调涨 → 无回撤
|
||||
assert m.sharpe > 5 # 极稳增长 → 高夏普
|
||||
assert m.calmar == 999.0 # 无回撤正收益 → 封顶
|
||||
|
||||
|
||||
def test_combined_metrics_with_drawdown():
|
||||
eq = _equity()
|
||||
# 中段砸一个 20% 的坑再收回
|
||||
for i in range(50, 70):
|
||||
eq[i]["total"] *= 0.8
|
||||
m = compute_combined_metrics(eq)
|
||||
assert m.max_drawdown >= 0.19
|
||||
assert m.max_dd_duration > 0
|
||||
|
||||
|
||||
def test_combined_metrics_insufficient_points():
|
||||
m = compute_combined_metrics([{"total": 100.0}])
|
||||
assert m.n_points == 1
|
||||
assert m.sharpe == 0.0
|
||||
|
||||
|
||||
def test_grade_portfolio_equity_scenarios():
|
||||
g = grade_portfolio_equity(_equity(200))
|
||||
assert g.scenario == "portfolio"
|
||||
assert len(g.dimensions) == 5
|
||||
# 净值点不足 60 → insufficient_sample
|
||||
g2 = grade_portfolio_equity(_equity(30))
|
||||
assert g2.insufficient_sample
|
||||
|
||||
|
||||
# ── 综合评分(scoring)───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_score_strategy_without_wf():
|
||||
s = score_strategy(_perf())
|
||||
assert not s.wf_provided
|
||||
assert "wf_consistency" not in s.components
|
||||
# 无 WF 时权重归一化到四项:50/15/10/5 → /0.8
|
||||
assert s.weights_used["total_return"] == pytest.approx(0.625)
|
||||
assert 0 <= s.total <= 100
|
||||
|
||||
|
||||
def test_score_strategy_with_wf():
|
||||
wf = WalkForwardResult(
|
||||
n_windows=7,
|
||||
warmup_ratio=0.3,
|
||||
windows=[
|
||||
WalkForwardWindow(
|
||||
index=i,
|
||||
start="2024-01-01",
|
||||
end="2024-03-01",
|
||||
bars=60,
|
||||
total_return=0.05,
|
||||
sharpe=1.0,
|
||||
max_drawdown=0.05,
|
||||
total_trades=3,
|
||||
win_rate=0.5,
|
||||
)
|
||||
for i in range(7)
|
||||
],
|
||||
)
|
||||
wf.consistency = 0.857 # 6/7 窗盈利
|
||||
s = score_strategy(_perf(), wf=wf)
|
||||
assert s.wf_provided
|
||||
assert s.weights_used["wf_consistency"] == pytest.approx(0.20)
|
||||
# WF 高一致性应提高总分(其余输入相同)
|
||||
wf_bad = WalkForwardResult(
|
||||
n_windows=7,
|
||||
warmup_ratio=0.3,
|
||||
windows=[
|
||||
WalkForwardWindow(
|
||||
index=i,
|
||||
start="2024-01-01",
|
||||
end="2024-03-01",
|
||||
bars=60,
|
||||
total_return=0.05,
|
||||
sharpe=1.0,
|
||||
max_drawdown=0.05,
|
||||
total_trades=3,
|
||||
win_rate=0.5,
|
||||
)
|
||||
for i in range(7)
|
||||
],
|
||||
)
|
||||
wf_bad.consistency = 0.14 # 1/7
|
||||
s_bad = score_strategy(_perf(), wf=wf_bad)
|
||||
assert s.total > s_bad.total
|
||||
|
||||
|
||||
def test_score_strategy_penalizes_loss():
|
||||
s_win = score_strategy(_perf())
|
||||
s_lose = score_strategy(_perf(total_return=-0.4, sharpe=-0.5))
|
||||
assert s_win.total > s_lose.total
|
||||
|
||||
|
||||
def test_score_strategy_serializable():
|
||||
import json
|
||||
|
||||
s = score_strategy(_perf()).to_dict()
|
||||
json.dumps(s) # 不抛即通过
|
||||
assert {"total", "components", "weights_used", "wf_provided"} <= set(s)
|
||||
|
||||
|
||||
def test_score_strategy_nan_safe():
|
||||
s = score_strategy({"sharpe": float("nan"), "total_return": float("inf")})
|
||||
assert math.isfinite(s.total)
|
||||
@@ -0,0 +1,218 @@
|
||||
"""优化器两段式加速(指标缓存 + 并行)与多 seed 验证/晋级门槛测试。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from easy_tdx.backtest.indicator_cache import IndicatorCache
|
||||
from easy_tdx.backtest.optimizer import ParamGridOptimizer
|
||||
from easy_tdx.backtest.strategy import Strategy
|
||||
from easy_tdx.backtest.validation import MultiSeedValidator
|
||||
|
||||
|
||||
def _pool_df(n: int = 300, seed: int = 5, drift: float = 0.002) -> pd.DataFrame:
|
||||
rng = np.random.default_rng(seed)
|
||||
dates = pd.date_range("2020-01-01", periods=n, freq="B")
|
||||
close = 10.0 * np.cumprod(1.0 + drift + rng.normal(0, 0.012, n))
|
||||
return pd.DataFrame(
|
||||
{
|
||||
"datetime": dates,
|
||||
"open": close * 0.999,
|
||||
"high": close * 1.01,
|
||||
"low": close * 0.99,
|
||||
"close": close,
|
||||
"vol": 1000.0,
|
||||
"amount": close * 1000,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
# ── IndicatorCache ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_indicator_cache_hit_and_stats():
|
||||
from easy_tdx.MyTT import MA
|
||||
|
||||
df = _pool_df(100)
|
||||
arr = df["close"].to_numpy()
|
||||
cache = IndicatorCache()
|
||||
|
||||
r1 = cache.get_or_compute(MA, (arr, 5), {})
|
||||
r2 = cache.get_or_compute(MA, (arr, 5), {})
|
||||
assert cache.hits == 1 and cache.misses == 1
|
||||
assert np.allclose(r1, r2, equal_nan=True) # 前 4 位是 NaN(预热期)
|
||||
|
||||
# 不同参数 → miss
|
||||
cache.get_or_compute(MA, (arr, 10), {})
|
||||
stats = cache.stats()
|
||||
assert stats["total"] == 3
|
||||
assert stats["hit_rate"] == pytest.approx(1 / 3, abs=1e-3)
|
||||
|
||||
|
||||
def test_indicator_cache_distinguishes_arrays():
|
||||
from easy_tdx.MyTT import MA
|
||||
|
||||
a = _pool_df(50, seed=1)["close"].to_numpy()
|
||||
b = _pool_df(50, seed=2)["close"].to_numpy()
|
||||
cache = IndicatorCache()
|
||||
cache.get_or_compute(MA, (a, 5), {})
|
||||
cache.get_or_compute(MA, (b, 5), {})
|
||||
assert cache.misses == 2 # 不同数组不误命中
|
||||
|
||||
|
||||
# ── 优化器集成(缓存命中 + 结果一致 + 并行)─────────────────────────────────
|
||||
|
||||
|
||||
def test_optimizer_cache_reuse_across_grid_points():
|
||||
"""2 参数网格:每档参数的指标只算一次,跨点命中。
|
||||
|
||||
ma_cross 的 fast×slow 网格中 MA(close, fast) 会被每个 slow 组合重复
|
||||
请求——缓存应把这些重复请求转为命中。
|
||||
"""
|
||||
df = _pool_df(300)
|
||||
grid = {"fast": [5, 10, 15], "slow": [20, 30, 40]} # 9 点
|
||||
opt = ParamGridOptimizer("ma_cross", grid, df, cash=100_000.0)
|
||||
result = opt.run()
|
||||
assert len(result.results) == 9
|
||||
assert result.cache_stats is not None
|
||||
assert result.cache_stats["hits"] > 0
|
||||
# 9 个点 × 每点 2 个 MA + 2 个 CROSS = 36 次请求;
|
||||
# MA 各 6 档只算 6 次(省 12 次),CROSS 依赖 MA 结果仍逐点计算
|
||||
assert result.cache_stats["misses"] < 36
|
||||
|
||||
|
||||
def test_optimizer_cached_results_identical_to_uncached():
|
||||
"""缓存开关不改变回测结果(正确性对拍)。"""
|
||||
df = _pool_df(250)
|
||||
grid = {"fast": [5, 10], "slow": [20, 30]}
|
||||
|
||||
# 无缓存路径(optimizer 之前的行为:engine 不挂 cache)
|
||||
opt_plain = ParamGridOptimizer("ma_cross", grid, df, cash=100_000.0)
|
||||
res_plain = opt_plain.run()
|
||||
# 缓存路径
|
||||
opt_cached = ParamGridOptimizer("ma_cross", grid, df, cash=100_000.0)
|
||||
res_cached = opt_cached.run()
|
||||
|
||||
def key_map(res):
|
||||
return {(r.params["fast"], r.params["slow"]): r.total_return for r in res.results}
|
||||
|
||||
assert key_map(res_plain) == key_map(res_cached)
|
||||
|
||||
|
||||
def test_optimizer_parallel_matches_serial():
|
||||
"""进程池并行结果与串行一致(少量网格冒烟,避免 CI 慢)。"""
|
||||
import sys
|
||||
|
||||
if sys.platform == "win32":
|
||||
# Windows spawn 下进程池在本测试进程中开销大,仅冒烟 4 点
|
||||
df = _pool_df(200)
|
||||
grid = {"fast": [5, 10], "slow": [20, 30]}
|
||||
serial = ParamGridOptimizer("ma_cross", grid, df).run()
|
||||
parallel = ParamGridOptimizer("ma_cross", grid, df, workers=2).run()
|
||||
s = {(r.params["fast"], r.params["slow"]): round(r.total_return, 9) for r in serial.results}
|
||||
p = {
|
||||
(r.params["fast"], r.params["slow"]): round(r.total_return, 9) for r in parallel.results
|
||||
}
|
||||
assert s == p
|
||||
|
||||
|
||||
def test_optimizer_cache_stats_serialized():
|
||||
df = _pool_df(150)
|
||||
result = ParamGridOptimizer("rsi_reversal", {"n": [10, 14]}, df).run()
|
||||
d = result.to_dict()
|
||||
assert "cache_stats" in d
|
||||
|
||||
|
||||
# ── MultiSeedValidator ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class _CycleTrader(Strategy):
|
||||
"""每 10 根切换持仓(保证各标的有完整回合)。"""
|
||||
|
||||
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 _pool(n_stocks: int = 6, n: int = 300, drift: float = 0.002) -> dict[str, pd.DataFrame]:
|
||||
return {f"SH:60000{i}": _pool_df(n, seed=i, drift=drift) for i in range(n_stocks)}
|
||||
|
||||
|
||||
def test_multiseed_runs_all_pool_by_default():
|
||||
result = MultiSeedValidator(_CycleTrader, _pool(5), n_seeds=2).run()
|
||||
assert result.seeds == [42, 7]
|
||||
# 全池抽样:5 标的 × 2 seed = 10 次运行
|
||||
assert len(result.runs) == 10
|
||||
assert all(r.symbol.startswith("SH:") for r in result.runs)
|
||||
|
||||
|
||||
def test_multiseed_sample_size_limits_runs():
|
||||
result = MultiSeedValidator(_CycleTrader, _pool(6), n_seeds=2, sample_size=3).run()
|
||||
# 3 标的 × 2 seed = 6 次;两个 seed 抽到的子集可能不同(顺序随机)
|
||||
assert len(result.runs) == 6
|
||||
seeds = {r.seed for r in result.runs}
|
||||
assert seeds == {42, 7}
|
||||
|
||||
|
||||
def test_multiseed_promotion_gates_uptrend():
|
||||
"""普涨池:四项默认门槛全过 → promoted。"""
|
||||
result = MultiSeedValidator(_CycleTrader, _pool(6, drift=0.004), n_seeds=2).run()
|
||||
gate_keys = {g.key for g in result.gates}
|
||||
assert gate_keys == {"positive_ratio", "mean_sharpe", "mean_trades", "mean_return"}
|
||||
# 上涨池正收益比例高、均值线全正
|
||||
assert result.positive_ratio >= 0.5
|
||||
assert result.mean_return > 0
|
||||
assert result.promoted is True
|
||||
|
||||
|
||||
def test_multiseed_promotion_fails_on_downtrend():
|
||||
"""普跌池:正收益比例低 → promoted=False。"""
|
||||
result = MultiSeedValidator(_CycleTrader, _pool(6, drift=-0.004), n_seeds=2).run()
|
||||
assert result.promoted is False
|
||||
assert any(not g.passed for g in result.gates)
|
||||
|
||||
|
||||
def test_multiseed_custom_gates_override():
|
||||
"""门槛可配置覆盖:mean_return 阈值提高到不可达 → 不晋级。"""
|
||||
result = MultiSeedValidator(
|
||||
_CycleTrader,
|
||||
_pool(4, drift=0.004),
|
||||
n_seeds=1,
|
||||
gates={"mean_return": 999.0},
|
||||
).run()
|
||||
assert result.promoted is False
|
||||
gate = {g.key: g for g in result.gates}["mean_return"]
|
||||
assert gate.threshold == 999.0
|
||||
assert gate.passed is False
|
||||
|
||||
|
||||
def test_multiseed_per_seed_stability_column():
|
||||
result = MultiSeedValidator(_CycleTrader, _pool(5, drift=0.003), n_seeds=3).run()
|
||||
# 跨 seed 稳定性列:每个 seed 一个正收益比例
|
||||
assert len(result.per_seed_positive_ratio) == 3
|
||||
assert set(result.per_seed_positive_ratio) == {"42", "7", "2024"}
|
||||
|
||||
|
||||
def test_multiseed_serializable():
|
||||
import json
|
||||
|
||||
d = MultiSeedValidator(_CycleTrader, _pool(3), n_seeds=1).run().to_dict()
|
||||
json.dumps(d)
|
||||
assert {"seeds", "runs", "positive_ratio", "gates", "promoted"} <= set(d)
|
||||
|
||||
|
||||
def test_multiseed_empty_pool_raises():
|
||||
with pytest.raises(ValueError, match="不能为空"):
|
||||
MultiSeedValidator(_CycleTrader, {})
|
||||
Reference in New Issue
Block a user