Files
easy_tdx_max/tests/unit/test_optimizer.py
T
Justin Gu 87fafe9131 feat(backtest): 参数网格寻优(optimizer + 前端寻优页)
对单个策略的 1-2 个参数做网格搜索,遍历用户指定的取值列表笛卡尔积,
每个组合跑一次回测,按 total_return 排序,返回排名表 + 热力图。

后端:
- ParamGridOptimizer(backtest/optimizer.py):itertools.product 遍历网格,
  每点 entry.build(params) + BacktestEngine.run(df),复用同一 DataFrame
- 网格大小上限 200 防组合爆炸,单点失败容错(跳过不中断)
- 2 参数时生成热力图矩阵(x/y 轴取值 + cell 收益率)
- POST /backtest/optimize/run/async 端点(后台任务)
- OptimizeBacktestRequest schema(param_grid 1-2 参数)

前端(/optimize 寻优页):
- ParamGridPicker:勾选 1-2 个寻优参数,逗号分隔填取值列表
- OptimizeResultTable:网格点排名表(按收益降序,最优高亮)
- OptimizeHeatmap:2 参数热力图(ECharts heatmap,绿→红映射收益)
- 最优点「查看」按钮跳转单标的页用该参数回测

测试:821 passed(+10 寻优器单测 + 3 寻优路由测试)
2026-07-03 03:55:35 +08:00

154 lines
5.2 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.
"""单元测试:参数网格寻优器."""
from __future__ import annotations
import numpy as np
import pandas as pd
import pytest
from easy_tdx.backtest.optimizer import (
GridPointResult,
OptimizeResult,
ParamGridOptimizer,
)
def _make_df(n: int = 150, seed: int = 42) -> pd.DataFrame:
"""生成带趋势的合成 OHLCV(确保均线策略能产生交易)。"""
rng = np.random.default_rng(seed)
close = 10.0 + np.cumsum(rng.normal(0, 0.3, n) + 0.05)
return pd.DataFrame(
{
"datetime": pd.date_range("2024-01-01", periods=n, freq="B"),
"open": close - 0.1,
"high": close + 0.2,
"low": close - 0.2,
"close": close,
"vol": rng.integers(1000, 10000, n).astype(float),
"amount": close * 5000,
}
)
class TestParamGridOptimizer:
"""寻优器核心逻辑."""
def test_grid_enumeration(self) -> None:
"""网格点数应等于笛卡尔积大小。"""
opt = ParamGridOptimizer(
strategy_name="ma_cross",
param_grid={"fast": [5, 10], "slow": [15, 20, 30]},
df=_make_df(),
)
result = opt.run()
assert len(result.results) == 6 # 2 × 3
def test_results_sorted_by_return_descending(self) -> None:
"""结果应按 total_return 降序排列。"""
opt = ParamGridOptimizer(
strategy_name="ma_cross",
param_grid={"fast": [5, 10, 20], "slow": [15, 20, 30]},
df=_make_df(),
)
result = opt.run()
returns = [r.total_return for r in result.results]
assert returns == sorted(returns, reverse=True)
def test_best_is_first_result(self) -> None:
"""best 应是 results[0]。"""
opt = ParamGridOptimizer(
strategy_name="ma_cross",
param_grid={"fast": [5, 10], "slow": [20, 30]},
df=_make_df(),
)
result = opt.run()
assert result.best is not None
assert result.best.params == result.results[0].params
assert result.best.total_return == result.results[0].total_return
def test_heatmap_2_params(self) -> None:
"""2 参数应生成热力图矩阵。"""
opt = ParamGridOptimizer(
strategy_name="ma_cross",
param_grid={"fast": [5, 10, 20], "slow": [15, 20, 30]},
df=_make_df(),
)
result = opt.run()
assert result.heatmap is not None
assert result.heatmap["x_name"] == "fast"
assert result.heatmap["y_name"] == "slow"
assert len(result.heatmap["x"]) == 3
assert len(result.heatmap["y"]) == 3
assert len(result.heatmap["data"]) == 9 # 3×3
def test_no_heatmap_for_single_param(self) -> None:
"""1 参数时 heatmap 应为 None。"""
opt = ParamGridOptimizer(
strategy_name="rsi_reversal",
param_grid={"n": [7, 14, 21]},
df=_make_df(),
)
result = opt.run()
assert result.heatmap is None
assert len(result.results) == 3
def test_grid_size_limit_exceeded(self) -> None:
"""超过 MAX_GRID_POINTS 应抛 ValueError。"""
big_grid = {f"p{i}": list(range(5)) for i in range(6)} # 5^6 = 15625
with pytest.raises(ValueError, match="超过上限"):
ParamGridOptimizer(
strategy_name="ma_cross",
param_grid=big_grid,
df=_make_df(),
)
def test_empty_value_list_rejected(self) -> None:
"""空取值列表应抛 ValueError。"""
with pytest.raises(ValueError, match="空取值列表"):
ParamGridOptimizer(
strategy_name="ma_cross",
param_grid={"fast": [], "slow": [20]},
df=_make_df(),
)
def test_to_dict_serializable(self) -> None:
"""to_dict 应返回 JSON 兼容结构。"""
opt = ParamGridOptimizer(
strategy_name="ma_cross",
param_grid={"fast": [5, 10], "slow": [20, 30]},
df=_make_df(),
)
result = opt.run()
d = result.to_dict()
assert d["strategy"] == "ma_cross"
assert d["param_names"] == ["fast", "slow"]
assert len(d["results"]) == 4
assert d["best"] is not None
assert "params" in d["best"]
assert "total_return" in d["best"]
def test_to_dict_cleans_nan(self) -> None:
"""NaN 指标应被清洗为 NoneJSON 兼容)。"""
# 构造含 NaN 的结果(无交易的参数组合 sharpe 可能 NaN)
result = OptimizeResult(
strategy="test",
param_names=["n"],
results=[
GridPointResult(params={"n": 1}, total_return=float("nan"), sharpe=float("inf")),
],
)
d = result.to_dict()
assert d["results"][0]["total_return"] is None
assert d["results"][0]["sharpe"] is None
def test_unknown_strategy_raises(self) -> None:
"""未知策略应在 run() 时抛 KeyError。"""
opt = ParamGridOptimizer(
strategy_name="nope",
param_grid={"x": [1]},
df=_make_df(),
)
with pytest.raises(KeyError):
opt.run()