mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 18:04:20 +08:00
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 寻优路由测试)
This commit is contained in:
@@ -0,0 +1,153 @@
|
||||
"""单元测试:参数网格寻优器."""
|
||||
|
||||
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 指标应被清洗为 None(JSON 兼容)。"""
|
||||
# 构造含 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()
|
||||
@@ -807,3 +807,93 @@ def test_portfolio_backtest_bad_strategy(client, monkeypatch):
|
||||
break
|
||||
time.sleep(0.05)
|
||||
assert final["status"] == "failed"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Phase 4: 参数网格寻优路由
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_optimize_request_validation():
|
||||
"""寻优请求校验:param_grid 非空、至少 1 个数据源。"""
|
||||
from easy_tdx.web.backtest_schemas import OptimizeBacktestRequest
|
||||
|
||||
one_bar = [
|
||||
{
|
||||
"datetime": "2024-01-01",
|
||||
"open": 1,
|
||||
"high": 1,
|
||||
"low": 1,
|
||||
"close": 1,
|
||||
"vol": 1,
|
||||
"amount": 1,
|
||||
}
|
||||
]
|
||||
req = OptimizeBacktestRequest(
|
||||
strategy="ma_cross",
|
||||
param_grid={"fast": [5, 10], "slow": [20, 30]},
|
||||
ohlcv=one_bar,
|
||||
)
|
||||
assert len(req.param_grid) == 2
|
||||
with pytest.raises(ValueError):
|
||||
OptimizeBacktestRequest(strategy="ma_cross", param_grid={"fast": [5]}) # 缺数据源
|
||||
with pytest.raises(ValueError):
|
||||
OptimizeBacktestRequest( # param_grid > 2 参数
|
||||
strategy="ma_cross",
|
||||
param_grid={"a": [1], "b": [2], "c": [3]},
|
||||
ohlcv=one_bar,
|
||||
)
|
||||
|
||||
|
||||
def test_optimize_endpoint(client, sample_ohlcv):
|
||||
"""POST /backtest/optimize/run/async 端到端(内联数据)。"""
|
||||
resp = client.post(
|
||||
"/api/v1/backtest/optimize/run/async",
|
||||
json={
|
||||
"strategy": "ma_cross",
|
||||
"param_grid": {"fast": [5, 10], "slow": [20, 30]},
|
||||
"cash": 100000,
|
||||
"ohlcv": sample_ohlcv,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 202, resp.text
|
||||
task_id = resp.json()["task_id"]
|
||||
|
||||
final = None
|
||||
for _ in range(200):
|
||||
poll = client.get(f"/api/v1/backtest/tasks/{task_id}")
|
||||
final = poll.json()
|
||||
if final["status"] in ("done", "failed"):
|
||||
break
|
||||
time.sleep(0.05)
|
||||
|
||||
assert final["status"] == "done", final
|
||||
result = final["result"]
|
||||
assert result["strategy"] == "ma_cross"
|
||||
assert result["param_names"] == ["fast", "slow"]
|
||||
assert len(result["results"]) == 4
|
||||
assert result["best"] is not None
|
||||
assert result["heatmap"] is not None
|
||||
assert len(result["heatmap"]["data"]) == 4
|
||||
|
||||
|
||||
def test_optimize_single_param_no_heatmap(client, sample_ohlcv):
|
||||
"""单参数寻优不应返回热力图。"""
|
||||
resp = client.post(
|
||||
"/api/v1/backtest/optimize/run/async",
|
||||
json={
|
||||
"strategy": "rsi_reversal",
|
||||
"param_grid": {"n": [7, 14, 21]},
|
||||
"ohlcv": sample_ohlcv,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 202
|
||||
task_id = resp.json()["task_id"]
|
||||
for _ in range(200):
|
||||
poll = client.get(f"/api/v1/backtest/tasks/{task_id}")
|
||||
final = poll.json()
|
||||
if final["status"] in ("done", "failed"):
|
||||
break
|
||||
time.sleep(0.05)
|
||||
assert final["status"] == "done"
|
||||
assert final["result"]["heatmap"] is None
|
||||
|
||||
Reference in New Issue
Block a user