mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 14:34:18 +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,248 @@
|
||||
"""参数网格寻优器。
|
||||
|
||||
对单个策略的 1-2 个参数做网格搜索:遍历用户指定的取值列表的笛卡尔积,
|
||||
每个组合跑一次回测,按 total_return 排序,返回排名 + 热力图矩阵。
|
||||
|
||||
设计镜像 :class:`~easy_tdx.backtest.combo.CombinationRunner` 的枚举/排序模式,
|
||||
但遍历的是参数组合而非策略组合。网格大小硬上限 ``MAX_GRID_POINTS`` 防组合爆炸。
|
||||
|
||||
用法::
|
||||
|
||||
from easy_tdx.backtest.optimizer import ParamGridOptimizer
|
||||
|
||||
opt = ParamGridOptimizer(
|
||||
strategy_name="ma_cross",
|
||||
param_grid={"fast": [5, 10, 20], "slow": [10, 20, 30]},
|
||||
df=df,
|
||||
cash=100000,
|
||||
)
|
||||
result = opt.run()
|
||||
print(result.best.params, result.best.total_return)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import itertools
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from easy_tdx.backtest.engine import BacktestEngine
|
||||
from easy_tdx.backtest.types import BacktestResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 网格点上限:防止组合爆炸。3 参数各 6 值 = 216 已接近上限。
|
||||
MAX_GRID_POINTS = 200
|
||||
|
||||
|
||||
@dataclass
|
||||
class GridPointResult:
|
||||
"""单个网格点的回测结果摘要。
|
||||
|
||||
Attributes:
|
||||
params: 该点的参数取值(如 {"fast": 10, "slow": 20})
|
||||
total_return: 总收益率
|
||||
sharpe: 夏普比率
|
||||
max_drawdown: 最大回撤
|
||||
total_trades: 总交易笔数
|
||||
win_rate: 胜率(0-1)
|
||||
profit_factor: 盈亏比
|
||||
"""
|
||||
|
||||
params: dict[str, Any]
|
||||
total_return: float = 0.0
|
||||
sharpe: float = 0.0
|
||||
max_drawdown: float = 0.0
|
||||
total_trades: int = 0
|
||||
win_rate: float = 0.0
|
||||
profit_factor: float = 0.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class OptimizeResult:
|
||||
"""网格寻优完整结果。
|
||||
|
||||
Attributes:
|
||||
strategy: 策略名
|
||||
param_names: 寻优的参数名列表(1-2 个,决定热力图维度)
|
||||
results: 所有网格点结果,按 total_return 降序排列
|
||||
best: 最优点(results[0] 的引用)
|
||||
heatmap: 2 参数时的热力图矩阵(x/y 轴取值 + cell 收益率);1 参数或空时为 None
|
||||
"""
|
||||
|
||||
strategy: str
|
||||
param_names: list[str]
|
||||
results: list[GridPointResult]
|
||||
best: GridPointResult | None = None
|
||||
heatmap: dict[str, Any] | None = None
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
"""序列化为 JSON 兼容字典。"""
|
||||
import math
|
||||
|
||||
def clean(v: Any) -> Any:
|
||||
if isinstance(v, float) and not math.isfinite(v):
|
||||
return None
|
||||
return v
|
||||
|
||||
return {
|
||||
"strategy": self.strategy,
|
||||
"param_names": self.param_names,
|
||||
"results": [
|
||||
{
|
||||
"params": r.params,
|
||||
"total_return": clean(r.total_return),
|
||||
"sharpe": clean(r.sharpe),
|
||||
"max_drawdown": clean(r.max_drawdown),
|
||||
"total_trades": r.total_trades,
|
||||
"win_rate": clean(r.win_rate),
|
||||
"profit_factor": clean(r.profit_factor),
|
||||
}
|
||||
for r in self.results
|
||||
],
|
||||
"best": (
|
||||
{
|
||||
"params": self.best.params,
|
||||
"total_return": clean(self.best.total_return),
|
||||
"sharpe": clean(self.best.sharpe),
|
||||
"max_drawdown": clean(self.best.max_drawdown),
|
||||
"total_trades": self.best.total_trades,
|
||||
"win_rate": clean(self.best.win_rate),
|
||||
"profit_factor": clean(self.best.profit_factor),
|
||||
}
|
||||
if self.best
|
||||
else None
|
||||
),
|
||||
"heatmap": self.heatmap,
|
||||
}
|
||||
|
||||
|
||||
class ParamGridOptimizer:
|
||||
"""参数网格寻优器。
|
||||
|
||||
遍历 ``param_grid`` 的笛卡尔积,每个组合实例化策略 + 跑回测,
|
||||
收集指标后排序。复用同一份 DataFrame(引擎无跨 run 缓存,安全)。
|
||||
|
||||
Args:
|
||||
strategy_name: 策略名(从注册表解析)
|
||||
param_grid: 参数取值网格,如 {"fast": [5,10,20], "slow": [10,20,30]}
|
||||
df: OHLCV DataFrame(所有网格点共用)
|
||||
cash: 初始资金
|
||||
commission: 佣金率
|
||||
min_commission: 最低佣金
|
||||
stamp_tax: 印花税
|
||||
slippage: 滑点
|
||||
execution: 成交模式
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
strategy_name: str,
|
||||
param_grid: dict[str, list[Any]],
|
||||
df: pd.DataFrame,
|
||||
cash: float = 100_000.0,
|
||||
commission: float = 0.0003,
|
||||
min_commission: float = 5.0,
|
||||
stamp_tax: float = 0.001,
|
||||
slippage: float = 0.0,
|
||||
execution: str = "next_open",
|
||||
) -> None:
|
||||
size = 1
|
||||
for vals in param_grid.values():
|
||||
size *= len(vals)
|
||||
if size > MAX_GRID_POINTS:
|
||||
raise ValueError(f"网格大小 {size} 超过上限 {MAX_GRID_POINTS},请减少参数取值数量")
|
||||
if size == 0:
|
||||
raise ValueError("param_grid 不能有空取值列表")
|
||||
|
||||
self._strategy_name = strategy_name
|
||||
self._param_grid = param_grid
|
||||
self._df = df
|
||||
self._cash = cash
|
||||
self._commission = commission
|
||||
self._min_commission = min_commission
|
||||
self._stamp_tax = stamp_tax
|
||||
self._slippage = slippage
|
||||
self._execution = execution
|
||||
|
||||
def run(self) -> OptimizeResult:
|
||||
"""执行网格寻优,返回排序后的结果。"""
|
||||
# 延迟导入避免循环依赖
|
||||
from easy_tdx.backtest.strategies import get_registry
|
||||
|
||||
entry = get_registry().get(self._strategy_name)
|
||||
param_names = list(self._param_grid.keys())
|
||||
value_lists = [self._param_grid[name] for name in param_names]
|
||||
|
||||
results: list[GridPointResult] = []
|
||||
for combo in itertools.product(*value_lists):
|
||||
params = dict(zip(param_names, combo, strict=True))
|
||||
try:
|
||||
strategy = entry.build(params)
|
||||
engine = BacktestEngine(
|
||||
strategy=strategy,
|
||||
cash=self._cash,
|
||||
commission=self._commission,
|
||||
min_commission=self._min_commission,
|
||||
stamp_tax=self._stamp_tax,
|
||||
slippage=self._slippage,
|
||||
execution=self._execution,
|
||||
)
|
||||
bt_result: BacktestResult = engine.run(self._df)
|
||||
perf = bt_result.performance
|
||||
results.append(
|
||||
GridPointResult(
|
||||
params=params,
|
||||
total_return=perf.get("total_return", 0.0),
|
||||
sharpe=perf.get("sharpe", 0.0),
|
||||
max_drawdown=perf.get("max_drawdown", 0.0),
|
||||
total_trades=int(perf.get("total_trades", 0)),
|
||||
win_rate=perf.get("win_rate", 0.0),
|
||||
profit_factor=perf.get("profit_factor", 0.0),
|
||||
)
|
||||
)
|
||||
except Exception: # noqa: BLE001 — 单点失败不中断整个网格
|
||||
logger.warning("网格点 %s 回测失败,跳过", params, exc_info=True)
|
||||
|
||||
# 按 total_return 降序
|
||||
results.sort(key=lambda r: r.total_return, reverse=True)
|
||||
|
||||
best = results[0] if results else None
|
||||
heatmap = self._build_heatmap(results, param_names) if len(param_names) == 2 else None
|
||||
|
||||
return OptimizeResult(
|
||||
strategy=self._strategy_name,
|
||||
param_names=param_names,
|
||||
results=results,
|
||||
best=best,
|
||||
heatmap=heatmap,
|
||||
)
|
||||
|
||||
def _build_heatmap(
|
||||
self,
|
||||
results: list[GridPointResult],
|
||||
param_names: list[str],
|
||||
) -> dict[str, Any]:
|
||||
"""2 参数时构建热力图矩阵:x=参数1取值,y=参数2取值,cell=total_return。
|
||||
|
||||
Returns:
|
||||
{"x": [...], "y": [...], "data": [[x_idx, y_idx, value], ...]}
|
||||
"""
|
||||
x_name, y_name = param_names
|
||||
x_vals = sorted(set(self._param_grid[x_name]))
|
||||
y_vals = sorted(set(self._param_grid[y_name]))
|
||||
x_idx = {v: i for i, v in enumerate(x_vals)}
|
||||
y_idx = {v: i for i, v in enumerate(y_vals)}
|
||||
|
||||
data: list[list[Any]] = []
|
||||
for r in results:
|
||||
x = r.params.get(x_name)
|
||||
y = r.params.get(y_name)
|
||||
if x not in x_idx or y not in y_idx:
|
||||
continue
|
||||
data.append([x_idx[x], y_idx[y], r.total_return])
|
||||
|
||||
return {"x_name": x_name, "y_name": y_name, "x": x_vals, "y": y_vals, "data": data}
|
||||
@@ -106,6 +106,47 @@ class PortfolioBacktestRequest(BaseModel):
|
||||
return self
|
||||
|
||||
|
||||
class OptimizeBacktestRequest(BaseModel):
|
||||
"""参数网格寻优请求。
|
||||
|
||||
在单个标的上,对策略的 1-2 个参数做网格搜索。数据来源与单标的回测一致
|
||||
(ohlcv 内联或 symbol 取行情)。param_grid 指定寻优参数及其取值列表,
|
||||
网格大小(各取值数乘积)上限 200。
|
||||
"""
|
||||
|
||||
strategy: str = Field(..., description="策略名")
|
||||
cash: float = Field(default=100000.0, gt=0)
|
||||
commission: float = Field(default=0.0003, ge=0, le=0.01)
|
||||
slippage: float = Field(default=0.0, ge=0, le=0.05)
|
||||
execution: Literal["next_open", "next_close", "this_close", "worst", "best"] = Field(
|
||||
default="next_open"
|
||||
)
|
||||
param_grid: dict[str, list[int | float | str]] = Field(
|
||||
...,
|
||||
min_length=1,
|
||||
max_length=2,
|
||||
description='参数取值网格,如 {"fast":[5,10,20], "slow":[15,20,30]}',
|
||||
)
|
||||
|
||||
# 数据来源 A:内联 OHLCV
|
||||
ohlcv: list[dict[str, Any]] | None = Field(default=None, max_length=2000)
|
||||
|
||||
# 数据来源 B:按标的取行情
|
||||
symbol: str | None = Field(default=None, pattern=r"^(SZ|SH|BJ):\d{6}$")
|
||||
category: Literal["DAY", "WEEK", "MONTH", "MIN_5", "MIN_15", "MIN_30", "MIN_60"] = Field(
|
||||
default="DAY"
|
||||
)
|
||||
count: int = Field(default=250, ge=20, le=800)
|
||||
start_date: str | None = Field(default=None)
|
||||
end_date: str | None = Field(default=None)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_data_source(self) -> OptimizeBacktestRequest:
|
||||
if self.ohlcv is None and self.symbol is None:
|
||||
raise ValueError("必须提供 ohlcv 或 symbol 之一")
|
||||
return self
|
||||
|
||||
|
||||
# ── 响应模型 ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
|
||||
@@ -19,6 +19,7 @@ from fastapi import APIRouter, Depends
|
||||
from easy_tdx.web.backtest_schemas import (
|
||||
BacktestRequest,
|
||||
BacktestResultResponse,
|
||||
OptimizeBacktestRequest,
|
||||
PortfolioBacktestRequest,
|
||||
StrategySchemaResponse,
|
||||
TaskStateResponse,
|
||||
@@ -153,6 +154,49 @@ async def run_portfolio_backtest_async(
|
||||
return TaskSubmitResponse(task_id=task_id, status=status)
|
||||
|
||||
|
||||
# ── 参数网格寻优 ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/backtest/optimize/run/async", response_model=TaskSubmitResponse, status_code=202)
|
||||
async def run_optimize_async(
|
||||
req: OptimizeBacktestRequest,
|
||||
client: Any = Depends(get_client),
|
||||
) -> TaskSubmitResponse:
|
||||
"""提交参数网格寻优后台任务。
|
||||
|
||||
在单个标的上对策略参数做网格搜索。数据获取支持内联 ohlcv 或按 symbol 取行情。
|
||||
通过 GET /backtest/tasks/{task_id} 轮询结果。
|
||||
"""
|
||||
# 1. 取数据
|
||||
if req.ohlcv is not None:
|
||||
df = _ohlcv_to_df(req.ohlcv)
|
||||
desc_bars = f"{len(df)} 根"
|
||||
elif req.symbol is not None:
|
||||
df = await _fetch_bars(client, req.symbol, req.category, 800)
|
||||
desc_bars = f"{req.symbol}"
|
||||
if req.start_date or req.end_date:
|
||||
df = _filter_df_by_date(df, req.start_date, req.end_date)
|
||||
else:
|
||||
raise ValueError("必须提供 ohlcv 或 symbol")
|
||||
|
||||
# 2. 捕获快照
|
||||
snapshot = req.model_copy()
|
||||
grid_size = 1
|
||||
for vals in snapshot.param_grid.values():
|
||||
grid_size *= len(vals)
|
||||
description = f"{snapshot.strategy} 寻优 | {desc_bars} | {grid_size}点"
|
||||
|
||||
# 3. 提交后台任务
|
||||
runner = get_runner()
|
||||
task_id = runner.submit(
|
||||
lambda: _run_optimize(df, snapshot),
|
||||
description=description,
|
||||
)
|
||||
state = runner.get(task_id)
|
||||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||||
return TaskSubmitResponse(task_id=task_id, status=status)
|
||||
|
||||
|
||||
# ── 内部实现 ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -294,6 +338,37 @@ async def _fetch_portfolio_bars(
|
||||
return stock_data_list
|
||||
|
||||
|
||||
def _run_optimize(df: pd.DataFrame, req: OptimizeBacktestRequest) -> dict[str, Any]:
|
||||
"""执行参数网格寻优并返回清洗后的结果字典(后台线程内调用)。"""
|
||||
from easy_tdx.backtest.optimizer import ParamGridOptimizer
|
||||
|
||||
optimizer = ParamGridOptimizer(
|
||||
strategy_name=req.strategy,
|
||||
param_grid=req.param_grid,
|
||||
df=df,
|
||||
cash=req.cash,
|
||||
commission=req.commission,
|
||||
slippage=req.slippage,
|
||||
execution=req.execution,
|
||||
)
|
||||
result = optimizer.run()
|
||||
return result.to_dict()
|
||||
|
||||
|
||||
def _filter_df_by_date(df: pd.DataFrame, start: str | None, end: str | None) -> pd.DataFrame:
|
||||
"""按日期范围过滤 DataFrame(闭区间,比较 YYYY-MM-DD)。"""
|
||||
if not start and not end:
|
||||
return df
|
||||
dt_col = "datetime" if "datetime" in df.columns else "date"
|
||||
dt_str = df[dt_col].astype(str).str.slice(0, 10)
|
||||
mask = pd.Series(True, index=df.index)
|
||||
if start:
|
||||
mask &= dt_str >= start
|
||||
if end:
|
||||
mask &= dt_str <= end
|
||||
return df[mask].reset_index(drop=True)
|
||||
|
||||
|
||||
def _now() -> float:
|
||||
"""获取当前时间戳(隔离 import,便于测试)。"""
|
||||
import time
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -9,6 +9,7 @@
|
||||
<nav class="app-nav">
|
||||
<RouterLink to="/" active-class="active">单标的回测</RouterLink>
|
||||
<RouterLink to="/portfolio" active-class="active">组合回测</RouterLink>
|
||||
<RouterLink to="/optimize" active-class="active">参数寻优</RouterLink>
|
||||
</nav>
|
||||
</header>
|
||||
<main class="app-main">
|
||||
|
||||
@@ -7,6 +7,7 @@ import type {
|
||||
BacktestResult,
|
||||
Bar,
|
||||
Category,
|
||||
OptimizeBacktestRequest,
|
||||
PortfolioBacktestRequest,
|
||||
StrategiesResponse,
|
||||
TaskState,
|
||||
@@ -124,6 +125,19 @@ export async function submitPortfolioTask(
|
||||
return (await resp.json()) as TaskSubmitResponse
|
||||
}
|
||||
|
||||
/** 提交参数网格寻优后台任务,返回 task_id。 */
|
||||
export async function submitOptimizeTask(
|
||||
req: OptimizeBacktestRequest,
|
||||
): Promise<TaskSubmitResponse> {
|
||||
const resp = await fetch(`${BASE}/backtest/optimize/run/async`, {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify(req),
|
||||
})
|
||||
if (!resp.ok) await throwError(resp)
|
||||
return (await resp.json()) as TaskSubmitResponse
|
||||
}
|
||||
|
||||
/** 查询后台任务状态(轮询用)。 */
|
||||
export async function fetchTask(taskId: string): Promise<TaskState> {
|
||||
const resp = await fetch(`${BASE}/backtest/tasks/${taskId}`)
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
<script setup lang="ts">
|
||||
// 2 参数寻优热力图(ECharts heatmap)。x=参数1取值,y=参数2取值,cell=total_return。
|
||||
|
||||
import { onBeforeUnmount, onMounted, ref, watch } from 'vue'
|
||||
|
||||
import echarts from '../echarts-setup'
|
||||
import type { OptimizeHeatmap } from '../types'
|
||||
|
||||
const props = defineProps<{
|
||||
heatmap: OptimizeHeatmap
|
||||
}>()
|
||||
|
||||
const container = ref<HTMLDivElement>()
|
||||
let chart: echarts.ECharts | null = null
|
||||
|
||||
function render() {
|
||||
if (!container.value) return
|
||||
chart ??= echarts.init(container.value, 'dark')
|
||||
chart.setOption(buildOption(), true)
|
||||
}
|
||||
|
||||
function buildOption(): echarts.EChartsCoreOption {
|
||||
const { x, y, data, x_name, y_name } = props.heatmap
|
||||
// 计算 visualMap 范围
|
||||
const values = data.map((d) => d[2]).filter((v): v is number => v !== null)
|
||||
const min = values.length ? Math.min(...values) : 0
|
||||
const max = values.length ? Math.max(...values) : 1
|
||||
return {
|
||||
backgroundColor: 'transparent',
|
||||
tooltip: {
|
||||
position: 'top',
|
||||
formatter: (p: { data: [number, number, number | null] }) => {
|
||||
const xv = x[p.data[0]]
|
||||
const yv = y[p.data[1]]
|
||||
const ret = p.data[2]
|
||||
const retStr = ret !== null ? `${(ret * 100).toFixed(2)}%` : '-'
|
||||
return `${x_name}=${xv}, ${y_name}=${yv}<br/>收益: ${retStr}`
|
||||
},
|
||||
},
|
||||
grid: { left: '12%', right: '5%', top: 20, bottom: 60 },
|
||||
xAxis: { type: 'category', data: x.map(String), name: x_name, splitArea: { show: true } },
|
||||
yAxis: { type: 'category', data: y.map(String), name: y_name, splitArea: { show: true } },
|
||||
visualMap: {
|
||||
min,
|
||||
max,
|
||||
calculable: true,
|
||||
orient: 'horizontal',
|
||||
left: 'center',
|
||||
bottom: 0,
|
||||
formatter: (v: number) => `${(v * 100).toFixed(0)}%`,
|
||||
inRange: { color: ['#18a058', '#2a2e3a', '#ef4146'] }, // 绿(低)→暗→红(高)
|
||||
},
|
||||
series: [
|
||||
{
|
||||
type: 'heatmap',
|
||||
data,
|
||||
label: { show: false },
|
||||
emphasis: { itemStyle: { shadowBlur: 10, shadowColor: 'rgba(0,0,0,0.5)' } },
|
||||
},
|
||||
],
|
||||
}
|
||||
}
|
||||
|
||||
function resize() {
|
||||
chart?.resize()
|
||||
}
|
||||
onMounted(() => {
|
||||
render()
|
||||
window.addEventListener('resize', resize)
|
||||
})
|
||||
onBeforeUnmount(() => {
|
||||
window.removeEventListener('resize', resize)
|
||||
chart?.dispose()
|
||||
chart = null
|
||||
})
|
||||
watch(() => props.heatmap, render)
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<div ref="container" class="heatmap-chart"></div>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.heatmap-chart {
|
||||
width: 100%;
|
||||
height: 380px;
|
||||
}
|
||||
</style>
|
||||
@@ -0,0 +1,101 @@
|
||||
<script setup lang="ts">
|
||||
// 网格点排名表,按 total_return 降序,最优高亮。
|
||||
|
||||
import type { GridPointResult } from '../types'
|
||||
|
||||
defineProps<{
|
||||
results: GridPointResult[]
|
||||
bestIndex?: number
|
||||
}>()
|
||||
|
||||
defineEmits<{ select: [params: Record<string, number | string>] }>()
|
||||
|
||||
function pct(v: number | null): string {
|
||||
return v !== null && Number.isFinite(v) ? `${(v * 100).toFixed(2)}%` : '-'
|
||||
}
|
||||
function num(v: number | null, d = 2): string {
|
||||
return v !== null && Number.isFinite(v) ? v.toFixed(d) : '-'
|
||||
}
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<table class="opt-table">
|
||||
<thead>
|
||||
<tr>
|
||||
<th>#</th>
|
||||
<th>参数</th>
|
||||
<th class="num">总收益</th>
|
||||
<th class="num">夏普</th>
|
||||
<th class="num">最大回撤</th>
|
||||
<th class="num">交易数</th>
|
||||
<th class="num">胜率</th>
|
||||
<th></th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr v-for="(r, i) in results" :key="i" :class="{ best: i === bestIndex }">
|
||||
<td class="rank">{{ i + 1 }}</td>
|
||||
<td class="params">{{ JSON.stringify(r.params) }}</td>
|
||||
<td class="num" :class="r.total_return !== null && r.total_return > 0 ? 'pos' : 'neg'">
|
||||
{{ pct(r.total_return) }}
|
||||
</td>
|
||||
<td class="num">{{ num(r.sharpe) }}</td>
|
||||
<td class="num neg">{{ pct(r.max_drawdown) }}</td>
|
||||
<td class="num">{{ r.total_trades }}</td>
|
||||
<td class="num">{{ pct(r.win_rate) }}</td>
|
||||
<td><button class="view-btn" @click="$emit('select', r.params)">查看</button></td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.opt-table {
|
||||
width: 100%;
|
||||
border-collapse: collapse;
|
||||
font-size: 13px;
|
||||
}
|
||||
.opt-table th,
|
||||
.opt-table td {
|
||||
padding: 6px 10px;
|
||||
border-bottom: 1px solid var(--border);
|
||||
text-align: left;
|
||||
}
|
||||
.opt-table th {
|
||||
color: var(--text-dim);
|
||||
font-size: 12px;
|
||||
position: sticky;
|
||||
top: 0;
|
||||
background: var(--bg-panel);
|
||||
}
|
||||
.num {
|
||||
text-align: right;
|
||||
font-family: var(--font-mono);
|
||||
}
|
||||
.params {
|
||||
font-family: var(--font-mono);
|
||||
font-size: 12px;
|
||||
color: var(--text-muted);
|
||||
}
|
||||
.rank {
|
||||
color: var(--text-dim);
|
||||
width: 32px;
|
||||
}
|
||||
.best {
|
||||
background: rgba(74, 158, 255, 0.08);
|
||||
}
|
||||
.best .rank {
|
||||
color: var(--accent);
|
||||
font-weight: 700;
|
||||
}
|
||||
.pos {
|
||||
color: var(--up);
|
||||
}
|
||||
.neg {
|
||||
color: var(--down);
|
||||
}
|
||||
.view-btn {
|
||||
font-size: 11px;
|
||||
padding: 2px 8px;
|
||||
}
|
||||
</style>
|
||||
@@ -0,0 +1,126 @@
|
||||
<script setup lang="ts">
|
||||
// 寻优参数选择:从策略参数里勾选 1-2 个,各填取值列表(逗号分隔)。
|
||||
|
||||
import { computed, ref, watch } from 'vue'
|
||||
|
||||
import type { StrategySchema } from '../types'
|
||||
|
||||
const props = defineProps<{
|
||||
strategy: StrategySchema | null
|
||||
modelValue: Record<string, Array<number | string>>
|
||||
}>()
|
||||
const emit = defineEmits<{ 'update:modelValue': [value: Record<string, Array<number | string>>] }>()
|
||||
|
||||
// 每个参数的取值输入框原始文本
|
||||
const inputs = ref<Record<string, string>>({})
|
||||
|
||||
// 选中要寻优的参数
|
||||
const selected = ref<Set<string>>(new Set())
|
||||
|
||||
function toggle(name: string) {
|
||||
if (selected.value.has(name)) {
|
||||
selected.value.delete(name)
|
||||
} else {
|
||||
if (selected.value.size >= 2) return // 最多 2 个
|
||||
selected.value.add(name)
|
||||
}
|
||||
// 触发响应式
|
||||
selected.value = new Set(selected.value)
|
||||
syncOutputs()
|
||||
}
|
||||
|
||||
function syncOutputs() {
|
||||
const out: Record<string, Array<number | string>> = {}
|
||||
for (const name of selected.value) {
|
||||
const raw = inputs.value[name] ?? ''
|
||||
out[name] = raw
|
||||
.split(/[,,\s]+/)
|
||||
.map((s) => s.trim())
|
||||
.filter(Boolean)
|
||||
.map((s) => {
|
||||
const n = Number(s)
|
||||
return Number.isFinite(n) ? n : s
|
||||
})
|
||||
}
|
||||
emit('update:modelValue', out)
|
||||
}
|
||||
|
||||
function onInput(name: string, val: string) {
|
||||
inputs.value[name] = val
|
||||
syncOutputs()
|
||||
}
|
||||
|
||||
// 切换策略时清空选择
|
||||
watch(
|
||||
() => props.strategy?.name,
|
||||
() => {
|
||||
selected.value = new Set()
|
||||
inputs.value = {}
|
||||
syncOutputs()
|
||||
},
|
||||
)
|
||||
|
||||
const gridPoints = computed(() => {
|
||||
const sizes = Array.from(selected.value).map((n) => {
|
||||
const raw = inputs.value[n] ?? ''
|
||||
return raw.split(/[,,\s]+/).filter((s) => s.trim()).length
|
||||
})
|
||||
return sizes.reduce((a, b) => a * b, 1)
|
||||
})
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<div class="grid-picker">
|
||||
<p class="hint">勾选 1-2 个参数寻优,填入取值列表(逗号分隔):</p>
|
||||
<div v-for="p in strategy?.params" :key="p.name" class="param-row">
|
||||
<label class="check">
|
||||
<input
|
||||
type="checkbox"
|
||||
:checked="selected.has(p.name)"
|
||||
:disabled="!selected.has(p.name) && selected.size >= 2"
|
||||
@change="toggle(p.name)"
|
||||
/>
|
||||
<span>{{ p.label }}({{ p.name }})</span>
|
||||
</label>
|
||||
<input
|
||||
v-if="selected.has(p.name)"
|
||||
:value="inputs[p.name] ?? ''"
|
||||
:placeholder="`如 ${p.default}, ${p.default}, ...`"
|
||||
class="values-input"
|
||||
@input="onInput(p.name, ($event.target as HTMLInputElement).value)"
|
||||
/>
|
||||
</div>
|
||||
<p class="grid-size">网格点数:{{ gridPoints }}(上限 200)</p>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.hint {
|
||||
color: var(--text-muted);
|
||||
font-size: 12px;
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.param-row {
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.check {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
font-size: 13px;
|
||||
color: var(--text);
|
||||
margin-bottom: 4px;
|
||||
}
|
||||
.check input[type='checkbox'] {
|
||||
width: auto;
|
||||
}
|
||||
.values-input {
|
||||
margin-top: 4px;
|
||||
font-family: var(--font-mono);
|
||||
}
|
||||
.grid-size {
|
||||
color: var(--text-dim);
|
||||
font-size: 11px;
|
||||
margin-top: 4px;
|
||||
}
|
||||
</style>
|
||||
@@ -1,8 +1,8 @@
|
||||
// ECharts 按需引入。只注册 MVP-A 用到的图表类型,避免全量引入(~1MB → ~400KB)。
|
||||
// 用到的:candlestick(K线)、line(净值/回撤曲线)、markPoint(买卖点标注)。
|
||||
// ECharts 按需引入。只注册用到的图表类型,避免全量引入(~1MB → ~400KB)。
|
||||
// 用到的:candlestick(K线)、line(净值/回撤曲线)、markPoint(买卖点标注)、heatmap(寻优热力图)。
|
||||
|
||||
import * as echarts from 'echarts/core'
|
||||
import { BarChart, CandlestickChart, LineChart } from 'echarts/charts'
|
||||
import { BarChart, CandlestickChart, HeatmapChart, LineChart } from 'echarts/charts'
|
||||
import {
|
||||
DataZoomComponent,
|
||||
GridComponent,
|
||||
@@ -10,6 +10,7 @@ import {
|
||||
MarkPointComponent,
|
||||
TitleComponent,
|
||||
TooltipComponent,
|
||||
VisualMapComponent,
|
||||
} from 'echarts/components'
|
||||
import { CanvasRenderer } from 'echarts/renderers'
|
||||
|
||||
@@ -18,12 +19,14 @@ echarts.use([
|
||||
CandlestickChart,
|
||||
LineChart,
|
||||
BarChart,
|
||||
HeatmapChart,
|
||||
GridComponent,
|
||||
TooltipComponent,
|
||||
LegendComponent,
|
||||
TitleComponent,
|
||||
DataZoomComponent,
|
||||
MarkPointComponent,
|
||||
VisualMapComponent,
|
||||
])
|
||||
|
||||
// A股惯例:红涨绿跌
|
||||
|
||||
@@ -1,12 +1,14 @@
|
||||
import { createRouter, createWebHistory } from 'vue-router'
|
||||
|
||||
import BacktestView from './views/BacktestView.vue'
|
||||
import OptimizeView from './views/OptimizeView.vue'
|
||||
import PortfolioView from './views/PortfolioView.vue'
|
||||
|
||||
// 单标的回测(/)+ 组合回测(/portfolio)。参数寻优/结果对比留待 Phase 4-5。
|
||||
// 单标的回测(/)+ 组合回测(/portfolio)+ 参数寻优(/optimize)。
|
||||
const routes = [
|
||||
{ path: '/', name: 'backtest', component: BacktestView },
|
||||
{ path: '/portfolio', name: 'portfolio', component: PortfolioView },
|
||||
{ path: '/optimize', name: 'optimize', component: OptimizeView },
|
||||
]
|
||||
|
||||
export const router = createRouter({
|
||||
|
||||
@@ -4,13 +4,22 @@
|
||||
import { defineStore } from 'pinia'
|
||||
import { computed, ref } from 'vue'
|
||||
|
||||
import { fetchStrategies, formatError, runBacktest, submitPortfolioTask, fetchTask } from '../api'
|
||||
import {
|
||||
fetchStrategies,
|
||||
formatError,
|
||||
runBacktest,
|
||||
submitPortfolioTask,
|
||||
submitOptimizeTask,
|
||||
fetchTask,
|
||||
} from '../api'
|
||||
import type {
|
||||
BacktestRequest,
|
||||
BacktestResult,
|
||||
Bar,
|
||||
PortfolioBacktestRequest,
|
||||
PortfolioResult,
|
||||
OptimizeBacktestRequest,
|
||||
OptimizeResult,
|
||||
StrategySchema,
|
||||
} from '../types'
|
||||
|
||||
@@ -105,6 +114,39 @@ export const useBacktestStore = defineStore('backtest', () => {
|
||||
error.value = ''
|
||||
}
|
||||
|
||||
// ── 参数网格寻优(Phase 4) ─────────────────────────────────────────────
|
||||
const optimizeResult = ref<OptimizeResult | null>(null)
|
||||
const optimizeRunning = ref(false)
|
||||
|
||||
/** 提交寻优后台任务并轮询直到完成。 */
|
||||
async function runOptimize(req: OptimizeBacktestRequest) {
|
||||
optimizeRunning.value = true
|
||||
error.value = ''
|
||||
optimizeResult.value = null
|
||||
try {
|
||||
const { task_id } = await submitOptimizeTask(req)
|
||||
const start = Date.now()
|
||||
// eslint-disable-next-line no-constant-condition
|
||||
while (true) {
|
||||
const state = await fetchTask(task_id)
|
||||
if (state.status === 'done' && state.result) {
|
||||
optimizeResult.value = state.result as OptimizeResult
|
||||
break
|
||||
}
|
||||
if (state.status === 'failed') {
|
||||
throw new Error(state.error || '寻优失败')
|
||||
}
|
||||
if (Date.now() - start > 180_000) throw new Error('寻优超时(180s)')
|
||||
await new Promise((r) => setTimeout(r, 400))
|
||||
}
|
||||
} catch (e) {
|
||||
error.value = formatError(e)
|
||||
optimizeResult.value = null
|
||||
} finally {
|
||||
optimizeRunning.value = false
|
||||
}
|
||||
}
|
||||
|
||||
return {
|
||||
// state
|
||||
strategies,
|
||||
@@ -116,6 +158,8 @@ export const useBacktestStore = defineStore('backtest', () => {
|
||||
error,
|
||||
portfolioResult,
|
||||
portfolioRunning,
|
||||
optimizeResult,
|
||||
optimizeRunning,
|
||||
// getters
|
||||
hasBars,
|
||||
// actions
|
||||
@@ -125,5 +169,6 @@ export const useBacktestStore = defineStore('backtest', () => {
|
||||
clearResult,
|
||||
runPortfolio,
|
||||
clearPortfolio,
|
||||
runOptimize,
|
||||
}
|
||||
})
|
||||
|
||||
+44
-1
@@ -130,7 +130,7 @@ export type TaskStatus = 'pending' | 'running' | 'done' | 'failed'
|
||||
export interface TaskState {
|
||||
task_id: string
|
||||
status: TaskStatus
|
||||
result: BacktestResult | PortfolioResult | null
|
||||
result: BacktestResult | PortfolioResult | OptimizeResult | null
|
||||
error: string | null
|
||||
description: string
|
||||
elapsed: number
|
||||
@@ -163,6 +163,49 @@ export interface PortfolioResult {
|
||||
combined_equity: EquityPoint[]
|
||||
}
|
||||
|
||||
// ── 参数网格寻优(Phase 4) ──────────────────────────────────────────────────
|
||||
|
||||
export interface OptimizeBacktestRequest {
|
||||
strategy: string
|
||||
cash?: number
|
||||
commission?: number
|
||||
slippage?: number
|
||||
execution?: ExecutionMode
|
||||
param_grid: Record<string, Array<number | string>>
|
||||
ohlcv?: Bar[]
|
||||
symbol?: string
|
||||
category?: Category
|
||||
count?: number
|
||||
start_date?: string
|
||||
end_date?: string
|
||||
}
|
||||
|
||||
export interface GridPointResult {
|
||||
params: Record<string, number | string>
|
||||
total_return: number | null
|
||||
sharpe: number | null
|
||||
max_drawdown: number | null
|
||||
total_trades: number
|
||||
win_rate: number | null
|
||||
profit_factor: number | null
|
||||
}
|
||||
|
||||
export interface OptimizeHeatmap {
|
||||
x_name: string
|
||||
y_name: string
|
||||
x: Array<number | string>
|
||||
y: Array<number | string>
|
||||
data: Array<[number, number, number | null]>
|
||||
}
|
||||
|
||||
export interface OptimizeResult {
|
||||
strategy: string
|
||||
param_names: string[]
|
||||
results: GridPointResult[]
|
||||
best: GridPointResult | null
|
||||
heatmap: OptimizeHeatmap | null
|
||||
}
|
||||
|
||||
// ── 错误响应(后端 ApiErrorResponse) ─────────────────────────────────────────
|
||||
|
||||
export interface ApiError {
|
||||
|
||||
@@ -0,0 +1,248 @@
|
||||
<script setup lang="ts">
|
||||
// 参数网格寻优主页面:左配置(选标的 + 策略 + 寻优参数)/ 右报告(排名表 + 热力图)。
|
||||
|
||||
import { computed, onMounted, ref } from 'vue'
|
||||
import { useRouter } from 'vue-router'
|
||||
|
||||
import OptimizeHeatmap from '../components/OptimizeHeatmap.vue'
|
||||
import OptimizeResultTable from '../components/OptimizeResultTable.vue'
|
||||
import ParamGridPicker from '../components/ParamGridPicker.vue'
|
||||
import SymbolPicker from '../components/SymbolPicker.vue'
|
||||
import type { ExecutionMode } from '../types'
|
||||
import { useBacktestStore } from '../stores/backtest'
|
||||
|
||||
const store = useBacktestStore()
|
||||
const router = useRouter()
|
||||
|
||||
const strategy = ref('ma_cross')
|
||||
const paramGrid = ref<Record<string, Array<number | string>>>({})
|
||||
const cash = ref(100000)
|
||||
const execution = ref<ExecutionMode>('next_open')
|
||||
const EXECUTIONS: ExecutionMode[] = ['next_open', 'next_close', 'this_close', 'worst', 'best']
|
||||
|
||||
const selectedStrategy = computed(
|
||||
() => store.strategies.find((s) => s.name === strategy.value) ?? null,
|
||||
)
|
||||
|
||||
onMounted(() => {
|
||||
store.loadStrategies().catch((e) => {
|
||||
store.error = `加载策略列表失败:${e instanceof Error ? e.message : e}`
|
||||
})
|
||||
})
|
||||
|
||||
// 网格点数(前端预校验,提示用户)
|
||||
const gridPoints = computed(() => {
|
||||
const sizes = Object.values(paramGrid.value).map((v) => v.length)
|
||||
return sizes.reduce((a, b) => a * b, 1)
|
||||
})
|
||||
|
||||
async function onRun() {
|
||||
if (!store.hasBars) {
|
||||
store.error = '请先取行情数据'
|
||||
return
|
||||
}
|
||||
if (Object.keys(paramGrid.value).length === 0) {
|
||||
store.error = '请勾选至少 1 个参数并填入取值'
|
||||
return
|
||||
}
|
||||
if (gridPoints.value > 200) {
|
||||
store.error = `网格点数 ${gridPoints.value} 超过上限 200`
|
||||
return
|
||||
}
|
||||
await store.runOptimize({
|
||||
strategy: strategy.value,
|
||||
param_grid: paramGrid.value,
|
||||
cash: cash.value,
|
||||
execution: execution.value,
|
||||
ohlcv: store.ohlcv,
|
||||
})
|
||||
}
|
||||
|
||||
// 点击排名表「查看」→ 跳转单标的页用该参数回测
|
||||
function onViewParams(params: Record<string, number | string>) {
|
||||
// 通过 query 传递参数,单标的页接收后自动填充
|
||||
router.push({
|
||||
path: '/',
|
||||
query: { strategy: strategy.value, params: JSON.stringify(params) },
|
||||
})
|
||||
}
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<div class="optimize-view">
|
||||
<aside class="config-panel">
|
||||
<section class="panel-section">
|
||||
<h3>行情数据</h3>
|
||||
<SymbolPicker />
|
||||
</section>
|
||||
|
||||
<section class="panel-section">
|
||||
<h3>策略</h3>
|
||||
<div class="field">
|
||||
<select v-model="strategy">
|
||||
<option v-for="s in store.strategies" :key="s.name" :value="s.name">
|
||||
{{ s.label }}({{ s.name }})
|
||||
</option>
|
||||
</select>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section class="panel-section">
|
||||
<h3>寻优参数</h3>
|
||||
<ParamGridPicker v-model="paramGrid" :strategy="selectedStrategy" />
|
||||
</section>
|
||||
|
||||
<section class="panel-section">
|
||||
<h3>资金</h3>
|
||||
<div class="field">
|
||||
<label>初始资金</label>
|
||||
<input v-model.number="cash" type="number" min="1000" step="10000" />
|
||||
</div>
|
||||
<div class="field">
|
||||
<label>成交模式</label>
|
||||
<select v-model="execution">
|
||||
<option v-for="e in EXECUTIONS" :key="e" :value="e">{{ e }}</option>
|
||||
</select>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<button
|
||||
class="primary run-btn"
|
||||
:disabled="store.optimizeRunning || !store.hasBars"
|
||||
@click="onRun"
|
||||
>
|
||||
{{ store.optimizeRunning ? '寻优中…' : '开始寻优' }}
|
||||
</button>
|
||||
</aside>
|
||||
|
||||
<main class="report-panel">
|
||||
<div v-if="store.error" class="error-banner">⚠ {{ store.error }}</div>
|
||||
|
||||
<div
|
||||
v-if="!store.optimizeResult && !store.optimizeRunning && !store.error"
|
||||
class="placeholder"
|
||||
>
|
||||
<p>选标的 → 取行情 → 选策略 → 勾选寻优参数 → 开始寻优</p>
|
||||
</div>
|
||||
|
||||
<div v-if="store.optimizeResult" class="report-content">
|
||||
<section class="report-section">
|
||||
<h3>最优结果</h3>
|
||||
<div v-if="store.optimizeResult.best" class="best-summary">
|
||||
<span class="best-params">{{ JSON.stringify(store.optimizeResult.best.params) }}</span>
|
||||
<span class="best-return pos">
|
||||
{{ (store.optimizeResult.best.total_return! * 100).toFixed(2) }}%
|
||||
</span>
|
||||
<span class="best-meta">
|
||||
夏普 {{ store.optimizeResult.best.sharpe?.toFixed(2) }} · 回撤
|
||||
{{ (store.optimizeResult.best.max_drawdown! * 100).toFixed(2) }}%
|
||||
</span>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section v-if="store.optimizeResult.heatmap" class="report-section">
|
||||
<h3>参数热力图({{ store.optimizeResult.heatmap.x_name }} × {{ store.optimizeResult.heatmap.y_name }})</h3>
|
||||
<OptimizeHeatmap :heatmap="store.optimizeResult.heatmap" />
|
||||
</section>
|
||||
|
||||
<section class="report-section">
|
||||
<h3>网格点排名({{ store.optimizeResult.results.length }} 个)</h3>
|
||||
<OptimizeResultTable
|
||||
:results="store.optimizeResult.results"
|
||||
:best-index="0"
|
||||
@select="onViewParams"
|
||||
/>
|
||||
</section>
|
||||
</div>
|
||||
</main>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.optimize-view {
|
||||
display: flex;
|
||||
height: 100%;
|
||||
}
|
||||
.config-panel {
|
||||
width: 320px;
|
||||
flex-shrink: 0;
|
||||
background: var(--bg-panel);
|
||||
border-right: 1px solid var(--border);
|
||||
padding: 16px;
|
||||
overflow-y: auto;
|
||||
}
|
||||
.panel-section {
|
||||
margin-bottom: 20px;
|
||||
padding-bottom: 16px;
|
||||
border-bottom: 1px solid var(--border);
|
||||
}
|
||||
.panel-section:last-of-type {
|
||||
border-bottom: none;
|
||||
}
|
||||
.panel-section h3 {
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
margin-bottom: 12px;
|
||||
}
|
||||
.run-btn {
|
||||
width: 100%;
|
||||
padding: 10px;
|
||||
font-size: 14px;
|
||||
}
|
||||
.report-panel {
|
||||
flex: 1;
|
||||
overflow-y: auto;
|
||||
padding: 16px 20px;
|
||||
}
|
||||
.placeholder {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
height: 100%;
|
||||
color: var(--text-dim);
|
||||
}
|
||||
.error-banner {
|
||||
background: rgba(239, 65, 70, 0.12);
|
||||
border: 1px solid var(--up);
|
||||
color: var(--up);
|
||||
padding: 10px 14px;
|
||||
border-radius: var(--radius);
|
||||
margin-bottom: 16px;
|
||||
font-size: 13px;
|
||||
}
|
||||
.report-section {
|
||||
background: var(--bg-panel);
|
||||
border: 1px solid var(--border);
|
||||
border-radius: var(--radius);
|
||||
padding: 14px 16px;
|
||||
margin-bottom: 16px;
|
||||
}
|
||||
.report-section h3 {
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
color: var(--text-muted);
|
||||
margin-bottom: 12px;
|
||||
}
|
||||
.best-summary {
|
||||
display: flex;
|
||||
align-items: baseline;
|
||||
gap: 16px;
|
||||
}
|
||||
.best-params {
|
||||
font-family: var(--font-mono);
|
||||
font-size: 14px;
|
||||
color: var(--accent);
|
||||
}
|
||||
.best-return {
|
||||
font-size: 22px;
|
||||
font-weight: 700;
|
||||
font-family: var(--font-mono);
|
||||
}
|
||||
.best-meta {
|
||||
color: var(--text-dim);
|
||||
font-size: 12px;
|
||||
}
|
||||
.pos {
|
||||
color: var(--up);
|
||||
}
|
||||
</style>
|
||||
Reference in New Issue
Block a user