From 87fafe91310623712235c2e0913eb106d9a62331 Mon Sep 17 00:00:00 2001 From: Justin Gu <97915@qq.com> Date: Fri, 3 Jul 2026 03:55:35 +0800 Subject: [PATCH] =?UTF-8?q?feat(backtest):=20=E5=8F=82=E6=95=B0=E7=BD=91?= =?UTF-8?q?=E6=A0=BC=E5=AF=BB=E4=BC=98=EF=BC=88optimizer=20+=20=E5=89=8D?= =?UTF-8?q?=E7=AB=AF=E5=AF=BB=E4=BC=98=E9=A1=B5=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 对单个策略的 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 寻优路由测试) --- src/easy_tdx/backtest/optimizer.py | 248 ++++++++++++++++++ src/easy_tdx/web/backtest_schemas.py | 41 +++ src/easy_tdx/web/routers/backtest.py | 75 ++++++ tests/unit/test_optimizer.py | 153 +++++++++++ tests/unit/test_web_backtest.py | 90 +++++++ web-ui/src/App.vue | 1 + web-ui/src/api.ts | 14 + web-ui/src/components/OptimizeHeatmap.vue | 88 +++++++ web-ui/src/components/OptimizeResultTable.vue | 101 +++++++ web-ui/src/components/ParamGridPicker.vue | 126 +++++++++ web-ui/src/echarts-setup.ts | 9 +- web-ui/src/router.ts | 4 +- web-ui/src/stores/backtest.ts | 47 +++- web-ui/src/types.ts | 45 +++- web-ui/src/views/OptimizeView.vue | 248 ++++++++++++++++++ 15 files changed, 1284 insertions(+), 6 deletions(-) create mode 100644 src/easy_tdx/backtest/optimizer.py create mode 100644 tests/unit/test_optimizer.py create mode 100644 web-ui/src/components/OptimizeHeatmap.vue create mode 100644 web-ui/src/components/OptimizeResultTable.vue create mode 100644 web-ui/src/components/ParamGridPicker.vue create mode 100644 web-ui/src/views/OptimizeView.vue diff --git a/src/easy_tdx/backtest/optimizer.py b/src/easy_tdx/backtest/optimizer.py new file mode 100644 index 0000000..62901bb --- /dev/null +++ b/src/easy_tdx/backtest/optimizer.py @@ -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} diff --git a/src/easy_tdx/web/backtest_schemas.py b/src/easy_tdx/web/backtest_schemas.py index 3b43819..b1057dc 100644 --- a/src/easy_tdx/web/backtest_schemas.py +++ b/src/easy_tdx/web/backtest_schemas.py @@ -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 + + # ── 响应模型 ─────────────────────────────────────────────────────────────────── diff --git a/src/easy_tdx/web/routers/backtest.py b/src/easy_tdx/web/routers/backtest.py index 55139b7..3c5515a 100644 --- a/src/easy_tdx/web/routers/backtest.py +++ b/src/easy_tdx/web/routers/backtest.py @@ -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 diff --git a/tests/unit/test_optimizer.py b/tests/unit/test_optimizer.py new file mode 100644 index 0000000..f6cff24 --- /dev/null +++ b/tests/unit/test_optimizer.py @@ -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() diff --git a/tests/unit/test_web_backtest.py b/tests/unit/test_web_backtest.py index 9cb0e27..cb97709 100644 --- a/tests/unit/test_web_backtest.py +++ b/tests/unit/test_web_backtest.py @@ -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 diff --git a/web-ui/src/App.vue b/web-ui/src/App.vue index adf2acc..b1723ad 100644 --- a/web-ui/src/App.vue +++ b/web-ui/src/App.vue @@ -9,6 +9,7 @@
diff --git a/web-ui/src/api.ts b/web-ui/src/api.ts index b7aad8f..a694dba 100644 --- a/web-ui/src/api.ts +++ b/web-ui/src/api.ts @@ -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 { + 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 { const resp = await fetch(`${BASE}/backtest/tasks/${taskId}`) diff --git a/web-ui/src/components/OptimizeHeatmap.vue b/web-ui/src/components/OptimizeHeatmap.vue new file mode 100644 index 0000000..55ff99d --- /dev/null +++ b/web-ui/src/components/OptimizeHeatmap.vue @@ -0,0 +1,88 @@ + + + + + diff --git a/web-ui/src/components/OptimizeResultTable.vue b/web-ui/src/components/OptimizeResultTable.vue new file mode 100644 index 0000000..1e30a17 --- /dev/null +++ b/web-ui/src/components/OptimizeResultTable.vue @@ -0,0 +1,101 @@ + + + + + diff --git a/web-ui/src/components/ParamGridPicker.vue b/web-ui/src/components/ParamGridPicker.vue new file mode 100644 index 0000000..f225d0f --- /dev/null +++ b/web-ui/src/components/ParamGridPicker.vue @@ -0,0 +1,126 @@ + + + + + diff --git a/web-ui/src/echarts-setup.ts b/web-ui/src/echarts-setup.ts index ec15116..e7fea42 100644 --- a/web-ui/src/echarts-setup.ts +++ b/web-ui/src/echarts-setup.ts @@ -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股惯例:红涨绿跌 diff --git a/web-ui/src/router.ts b/web-ui/src/router.ts index 610ba2c..4040e04 100644 --- a/web-ui/src/router.ts +++ b/web-ui/src/router.ts @@ -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({ diff --git a/web-ui/src/stores/backtest.ts b/web-ui/src/stores/backtest.ts index 0ac3a27..492df28 100644 --- a/web-ui/src/stores/backtest.ts +++ b/web-ui/src/stores/backtest.ts @@ -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(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, } }) diff --git a/web-ui/src/types.ts b/web-ui/src/types.ts index 0540c6d..8822ef9 100644 --- a/web-ui/src/types.ts +++ b/web-ui/src/types.ts @@ -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> + ohlcv?: Bar[] + symbol?: string + category?: Category + count?: number + start_date?: string + end_date?: string +} + +export interface GridPointResult { + params: Record + 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 + y: Array + 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 { diff --git a/web-ui/src/views/OptimizeView.vue b/web-ui/src/views/OptimizeView.vue new file mode 100644 index 0000000..b32785c --- /dev/null +++ b/web-ui/src/views/OptimizeView.vue @@ -0,0 +1,248 @@ + + + + +