From ec3309163b90996a9e59223e74e9eca766807631 Mon Sep 17 00:00:00 2001 From: im47cn <67424112+im47cn@users.noreply.github.com> Date: Fri, 10 Jul 2026 11:38:30 +0800 Subject: [PATCH] =?UTF-8?q?feat(optimizer):=20=E5=8F=82=E6=95=B0=E7=BD=91?= =?UTF-8?q?=E6=A0=BC=E6=90=9C=E7=B4=A2=E4=BC=98=E5=8C=96=E5=99=A8=20(?= =?UTF-8?q?=E5=B9=B6=E8=A1=8C=E5=9B=9E=E6=B5=8B=20+=20=E7=9B=AE=E6=A0=87?= =?UTF-8?q?=E6=8E=92=E5=BA=8F=20+=20SSE=20+=20=E5=89=8D=E7=AB=AF=E9=9D=A2?= =?UTF-8?q?=E6=9D=BF)=20(#82)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(optimizer): 参数网格搜索优化器核心 + PanelCache 线程安全 PR2a 第一部分 (后端核心, 无 API/前端): app/backtest/optimizer.py: - expand_param_grid: 校验(类型/范围/选项) + 笛卡尔积展开; 支持显式候选值列表 与 {min,max,step} 范围两种写法; GRID_MAX_COMBINATIONS=2000 硬上限防爆炸。 - StrategyOptimizer.optimize: ThreadPoolExecutor 并行跑各参数组回测, 按目标指标 排序返回最优 + 全排名。objective 统一转越大越好(min 类目标取负), None/inf/ 失败组沉底。支持进度回调 (done/total/best) 与 cancel_event。 - 可选目标含 PR1 新增的 sortino/mc_maxdd_*。 engine.py PanelCache 线程安全: - get_or_compute 整体加锁 — 并行优化器对同一 symbols/日期跑几十组参数时, 首个 线程 compute 面板、其余等待后命中缓存, 同一面板只 scan_parquet+compute_all 一次 (关键性能: 参数不影响面板 key)。max_size 2->4, ttl 180->900s 适配长任务。 测试 24 例: 网格展开/校验(未知参数/越界/选项/空/组合爆炸) + 编排(排序/每组一次/ min方向/失败None沉底/取消/进度/非法目标)。用假 service 注入受控 stats 验证。 * fix(optimizer): 子代理审查修复 — 异常隔离/展示符号/浮点端点/kwargs 校验 两份子代理审查发现 4 个真实 bug: [高] 单组异常拖垮整批: _run_one 未捕获 service.run 异常, fut.result() 会 re-raise 冲出 as_completed → 整个网格搜索崩溃、已完成结果全丢。加了并行后 这个概率显著上升。用 try/except 隔离, 该组记为 error 继续。 [高] min 方向 best_score 符号错误: 内部把 min 目标取负做排序键, 但直接把 取负值当 best_score 返回 → 用户看到 avg_holding_days=-3.0 (负天数)。根因是 排序键与展示值混用。改为分离: 内部 _sort (取负空间, 不外露) + objective_raw (原始展示值)。best_score/进度回调均用原始值。 [中] 浮点累加丢端点: v += step 累积误差使 0.1 步长丢失 max 端点 (0.1+0.1+0.1=0.30000004 > 0.3+1e-9)。改整数计数 lo + i*step。 [中] backtest_kwargs 非法/冲突 key: 展开传给 StrategyBacktestConfig 时若含 非法或保留字段 (symbols 等) 会在 worker 抛 TypeError, 被上述 #1 放大成崩溃。 入口加白名单校验, 提前明确报错。 测试新增 8 例: 符号还原/max_drawdown 负值排序/异常隔离/kwargs 非法+保留/ base_params 合并/浮点端点/去重。原 test_min_direction 只验 params 放过了符号 bug, 现补 best_score 断言。 * feat(optimizer): 参数优化 API SSE 端点 + 前端优化器面板 PR2a 完成 (API + 前端), 接上后端核心: 后端 api/backtest.py: - GET /optimize/stream: 复用 _BacktestJob SSE 框架, 后台线程跑 StrategyOptimizer, progress_cb 推 done/total/best_score, 完成推 best_params/results 排名。param_grid 走 JSON 字符串查询参数 (EventSource 仅支持 GET, 同 params/overrides 惯例)。 - POST /optimize/cancel: 从 query string 复原同一 job_key, cancel_event 停止。 - backtest_kwargs 透传各回测参数 (matching/fees/mode/...) 到每组回测。 前端: - lib/optimizerTask.ts: SSE 客户端 (镜像 backtestTask), 进度/结果/重连/取消。 - pages/backtest/StrategyOptimizer.tsx: 配置面板 (选策略 → 勾选可扫参数设 min/max/step, bool/select 自动全扫; 优化目标下拉; 日期; 组合数实时预估 + 2000 上限提示) + 结果面板 (最优参数高亮 + 排名表: objective/夏普/索提诺/收益/回撤/ 胜率/交易数)。 - Backtest.tsx: 新增 '参数优化' 第三 tab。 测试 test_optimizer_api.py 3 例: job_key 确定性 + 区分 grid/objective + stream 与 cancel 复原同一 key (守护 PR3 C1 类失配)。前端 tsc 无新增类型错误。 注: SSE 端到端需真实日K数据, 本地 mode=none 无法验证实际回测; 但 SSE 管线镜像 已测的 strategy_stream, 优化器核心 24 单测覆盖。 * fix(optimizer): 子代理审查修复 — API 取消/空网格/方向对齐 + 前端重连/切换/展示 两份子代理审查(API + 前端)发现的真实问题: 后端 API: - [中] param_grid 为 null/[]/'' 等合法 JSON 但非网格对象时, 原逻辑跳过线程却不置 job.done -> event_generator 永久空转、job 挂死在表中。改为非空 dict 校验拦下。 - [中] 取消后前端收到 done 而非取消提示: 优化器把 cancel 当每组失败正常返回 dict, done 分支照推'完成'。改为 done 分支先检查 cancel_event, 分流为'优化已取消'。 - [低] direction 空串边界: stream 侧 '' 与 cancel 侧 or None 口径不一致致 job_key 失配(cancel 失效)。stream 加 direction = direction or None 对齐。 前端: - [中] tryReconnectOptimize 是死代码(无调用方)违反 NO DEAD CODE: 接入 useEffect 挂载恢复(镜像 StrategyBacktest 的 tryReconnect), 刷新/切页后恢复未完成优化。 - [中] 切策略后旧结果残留错配(参数列是旧策略): onSelectStrategy 加 clearOptimize。 - [低] 排名表 slice(50) 静默截断: 加'仅显示前50/共N组'提示。 - [低] objective_raw 裸数与 best_score 精度不一: 统一 toFixed(3)。 后端 76 测试通过; 前端 tsc 无新增类型错误。 * refactor(optimizer): job_key 回吐 — 消除 cancel 两侧重算的脆弱契约 采纳子代理审查建议, 从结构上根除整类 job_key 失配 bug: 之前 cancel 需从 query string 逐字段重算 job_key, 必须与 stream 侧完全一致 — 任何默认值/None-空串/类型转换漂移都静默导致取消失效 (PR3 C1、本轮 direction 空串失配都是这个结构的产物)。 改为: stream 首个 SSE 事件 (event: job) 回吐后端算出的 job_key, 前端存下, cancel 直接原样传回按 key 查表。cancel 侧不再重算, 契约漂移无从发生。 - 后端 optimize_stream: 首事件 yield event: job {key}; optimize_cancel 简化为 body.job_key 直接查 _running_jobs (删除 40 行 qs 重算)。 - 前端 optimizerTask: 监听 job 事件存 currentJobKey + localStorage; stopOptimize 改传 {job_key}; done/error/cancel 清理 key。 - 测试: 原 stream/cancel qs 对齐测试已无意义, 改为验证 cancel 按回吐 key 查表 (命中/已完成/未知 key 三态), 用轻量 fake Request 直调 endpoint。 后端 76 测试通过; 前端 tsc 无新增错误。策略回测路径未动 (已合并 + C1 测试守护), 本重构仅限本 PR 新增的优化器路径。 * chore(optimizer): 移除冗余 PanelCache 改动 — main 已独立实现线程安全 rebase 到 main 时发现上游已独立给 PanelCache 加锁 (且 compute 放锁外, 比本 PR 原方案更优), 并新增 asset_type 维度。故本 PR 的 PanelCache 改动 (加锁 + size/ttl bump + docstring) 全部冗余且 docstring 已与 main 实际锁行为不符, 回退到 main 版本。 优化器共享单一 panel key, main 的 PanelCache 已完全够用。 至此本 PR 零 engine.py 改动。 * fix(optimizer): 处理 #82 review 的 6 处问题 作者 review #82 提出的阻塞/改进项, 逐条修复: [阻塞1] 前端构建失败: EmptyState 只接受 title/hint, 误用了 description -> npm run build (tsc -b) 报 TS2322。改为 hint, build 通过。 [阻塞2] 优化没用用户当前策略配置: optimize API/前端只传 strategy_id/param_grid/ objective/日期/mode, 未传 params/overrides。补齐 —— API 新增 params(base_params)/ overrides 两个 query 参数并纳入 job_key; 前端把选中策略的 params_defaults 作为未扫描 参数固定值, buildDefaultOverrides(strategy) 让 basic_filter/信号/风控按当前策略参与。 抽 lib/strategyOverrides.ts 共享 (与策略回测页同口径, 避免重复)。 [阻塞3] 切策略丢失运行中任务控制权: 原 clearOptimize 只清前端状态, 不关 SSE/不 cancel 后端/不清 localStorage -> 后端继续跑但 Stop 消失。改为: 有任务在跑时切策略先 stopOptimize (真正 cancel + 关连接 + 清存储)。 [阻塞4] 停止按钮竞态: job_key 只在收到首个 job 事件后才有, 刚点开始就点停止时前端还没 key, cancel 落空。改为: stopOptimize 标记 cancelRequested, 有 key 立即 POST cancel, 无 key 则保持 SSE 等 job 事件到达时补发 cancel 再关 (关 SSE 不停后端 daemon 线程, 必须 真 POST); 加 5s 兜底。 [改进5] SSE 断线健壮性: 无 data 断线原全靠浏览器自动重连无上限。加 MAX_RECONNECT=5, 超限置 error 停 pending。 [改进6] 组合数校验与后端不一致: 前端 round((hi-lo)/step)+1 会把 min=0/max=1/step=0.6 显示为可运行 3 组, 但后端生成末值 1.2>max 报错。新增 sweepError 与后端 _candidates_for 同口径校验 (步长不整除/越界), 前端提前拦并禁用运行。 后端 156 测试通过; 前端 npm run build 通过。 --- backend/app/api/backtest.py | 208 +++++++++++ backend/app/backtest/optimizer.py | 295 +++++++++++++++ backend/tests/backtest/test_optimizer_api.py | 60 +++ backend/tests/backtest/test_optimizer_grid.py | 125 +++++++ backend/tests/backtest/test_optimizer_run.py | 199 ++++++++++ frontend/src/lib/optimizerTask.ts | 274 ++++++++++++++ frontend/src/lib/strategyOverrides.ts | 24 ++ frontend/src/pages/Backtest.tsx | 21 +- .../src/pages/backtest/StrategyOptimizer.tsx | 348 ++++++++++++++++++ 9 files changed, 1550 insertions(+), 4 deletions(-) create mode 100644 backend/app/backtest/optimizer.py create mode 100644 backend/tests/backtest/test_optimizer_api.py create mode 100644 backend/tests/backtest/test_optimizer_grid.py create mode 100644 backend/tests/backtest/test_optimizer_run.py create mode 100644 frontend/src/lib/optimizerTask.ts create mode 100644 frontend/src/lib/strategyOverrides.ts create mode 100644 frontend/src/pages/backtest/StrategyOptimizer.tsx diff --git a/backend/app/api/backtest.py b/backend/app/api/backtest.py index 6d06d6b..7d89704 100644 --- a/backend/app/api/backtest.py +++ b/backend/app/api/backtest.py @@ -511,3 +511,211 @@ async def strategy_cancel(request: Request): return {"ok": True} return {"ok": False, "message": "任务不存在或已完成"} + +# ══════════════════════════════════════════════════════════════ +# 参数网格优化器 — 复用 _BacktestJob SSE 框架 (多组参数并行回测 + 排序) +# ══════════════════════════════════════════════════════════════ + +# 透传给每组回测的 StrategyBacktestConfig 字段 (作为 backtest_kwargs)。 +_OPT_BT_FIELDS = [ + "matching", "fees_pct", "commission_pct", "stamp_tax_pct", "slippage_bps", + "max_positions", "max_exposure_pct", "initial_capital", "position_sizing", + "mode", "holding_days", +] + + +def _make_opt_job_key(strategy_id, symbols, start, end, param_grid, objective, direction, bt_sig, params=None, overrides=None) -> str: + raw = f"OPT|{strategy_id}|{symbols}|{start}|{end}|{param_grid}|{objective}|{direction}|{bt_sig}|{params}|{overrides}" + return hashlib.md5(raw.encode()).hexdigest()[:12] + + +def _opt_backtest_kwargs( + matching, fees_pct, commission_pct, stamp_tax_pct, slippage_bps, + max_positions, max_exposure_pct, initial_capital, position_sizing, mode, holding_days, +) -> dict: + return { + "matching": matching, + "fees_pct": fees_pct, + "commission_pct": commission_pct, + "stamp_tax_pct": stamp_tax_pct, + "slippage_bps": slippage_bps, + "max_positions": int(max_positions), + "max_exposure_pct": float(max_exposure_pct), + "initial_capital": float(initial_capital), + "position_sizing": position_sizing, + "mode": mode, + "holding_days": int(holding_days), + } + + +@router.get("/optimize/stream") +async def optimize_stream( + request: Request, + strategy_id: str, + param_grid: str, # JSON: {param_id: [values] | {min,max,step}} + objective: str = "sortino", + direction: str | None = None, + max_workers: int = 4, + params: str | None = None, # JSON: 未扫描参数固定为用户当前值 (base_params) + overrides: str | None = None, # JSON: 策略当前的 basic_filter/signals/风控等覆盖 + symbols: str | None = None, + start: str | None = None, + end: str | None = None, + matching: str = "open_t+1", + fees_pct: float = 0.0002, + commission_pct: float | None = None, + stamp_tax_pct: float | None = None, + slippage_bps: float = 5.0, + max_positions: int = 10, + max_exposure_pct: float = 1.0, + initial_capital: float = 1_000_000.0, + position_sizing: str = "equal", + mode: str = "position", + holding_days: int = 5, +): + """SSE 流式参数优化: 并行跑各参数组回测, 按 objective 排序。 + + 事件类型: + - progress: {type: "optimizer_progress", done, total, best_score} + - done: {result} (含 best_params / results 排名) + - error: {message} + """ + from app.backtest.optimizer import OptimizeConfig, StrategyOptimizer + from app.backtest.strategy import StrategyBacktestService + + engine = _get_engine(request) + strategy_engine = request.app.state.strategy_engine + svc = StrategyBacktestService(engine, strategy_engine) + + end_date = date.fromisoformat(end) if end else date.today() + if start: + start_date = date.fromisoformat(start) + else: + earliest = request.app.state.repo.earliest_daily_date() + start_date = earliest or (end_date - timedelta(days=FACTOR_DEFAULT_DAYS)) + + guard_violated = False + if settings.backtest_range_guard and (end_date - start_date).days + 1 > BACKTEST_MAX_SERVER_DAYS: + guard_violated = True + + # 空串归一为 None, 与 cancel 侧 `_get("direction") or None` 口径一致, 避免 job_key 失配。 + direction = direction or None + bt_kwargs = _opt_backtest_kwargs( + matching, fees_pct, commission_pct, stamp_tax_pct, slippage_bps, + max_positions, max_exposure_pct, initial_capital, position_sizing, mode, holding_days, + ) + bt_sig = "|".join(f"{k}={bt_kwargs[k]}" for k in _OPT_BT_FIELDS) + job_key = _make_opt_job_key(strategy_id, symbols, start, end, param_grid, objective, direction, bt_sig, params, overrides) + + _cleanup_stale_jobs() + with _jobs_lock: + job = _running_jobs.get(job_key) + if job is None: + job = _BacktestJob(job_key) + _running_jobs[job_key] = job + is_new = True + else: + is_new = False + + async def event_generator(): + # 首个事件回吐 job_key, 前端存下供 cancel 直接引用 (消除两侧重算契约)。 + yield f"event: job\ndata: {json.dumps({'key': job_key}, ensure_ascii=False)}\n\n" + + if guard_violated: + yield f"event: error\ndata: {json.dumps({'message': BACKTEST_SERVER_GUARD_MESSAGE}, ensure_ascii=False)}\n\n" + return + + if is_new and not job.done: + try: + grid = json.loads(param_grid) + except (json.JSONDecodeError, TypeError): + grid = None + # grid 必须是非空 dict; null/[]/"" 等合法 JSON 但结构错误也在此拦下, + # 否则会跳过线程启动却不置 done -> event_generator 永久空转、job 挂死。 + if not isinstance(grid, dict) or not grid: + job.error = "param_grid 必须是非空的参数网格对象" + job.done = True + job.finish_ts = time.time() + grid = None + + if grid is not None: + # 未扫描参数固定为用户当前值 (base_params); overrides 让策略的 basic_filter/ + # 信号/风控按用户当前配置参与, 保证优化的就是用户实际回测的策略。 + try: + base_params = json.loads(params) if params else {} + except (json.JSONDecodeError, TypeError): + base_params = {} + try: + ov = json.loads(overrides) if overrides else None + except (json.JSONDecodeError, TypeError): + ov = None + ocfg = OptimizeConfig( + strategy_id=strategy_id, + symbols=[s.strip() for s in symbols.split(",") if s.strip()] if symbols else None, + start=start_date, + end=end_date, + param_grid=grid, + objective=objective, + direction=direction, + max_workers=int(max_workers), + base_params=base_params if isinstance(base_params, dict) else {}, + overrides=ov if isinstance(ov, dict) else None, + backtest_kwargs=bt_kwargs, + ) + + def _run_opt(): + try: + opt = StrategyOptimizer(svc, strategy_engine) + job.result = opt.optimize(ocfg, lambda d: job.progress.append(d), job.cancel_event) + job.done = True + job.finish_ts = time.time() + except Exception as e: + job.error = str(e) + job.done = True + job.finish_ts = time.time() + + threading.Thread(target=_run_opt, daemon=True).start() + + cursor = 0 + tick = 0 + try: + while True: + if job.done: + if job.error: + yield f"event: error\ndata: {json.dumps({'message': job.error}, ensure_ascii=False)}\n\n" + elif job.cancel_event.is_set(): + # 取消时优化器把每组记为 cancelled 并正常返回, 需在此分流为取消提示而非"完成"。 + yield f"event: error\ndata: {json.dumps({'message': '优化已取消'}, ensure_ascii=False)}\n\n" + elif job.result is not None: + yield f"event: done\ndata: {json.dumps(job.result, ensure_ascii=False, default=str)}\n\n" + return + tick += 1 + if tick % 4 == 0 and await request.is_disconnected(): + break + while cursor < len(job.progress): + msg = job.progress[cursor] + cursor += 1 + yield f"event: progress\ndata: {json.dumps(msg, ensure_ascii=False, default=str)}\n\n" + await asyncio.sleep(0.5) + except asyncio.CancelledError: + raise + + return StreamingResponse(event_generator(), media_type="text/event-stream") + + +@router.post("/optimize/cancel") +async def optimize_cancel(request: Request): + """取消优化任务 — 前端传 stream 首事件回吐的 job_key, 后端直接查表。 + + 不再让 cancel 侧重算 job_key: 两侧重算必须逐字段一致的脆弱契约(PR3 C1 / direction + 空串失配都源于此)在此彻底消除。stream 首个 SSE 事件把后端算出的 key 回吐给前端, + cancel 原样传回即可。 + """ + body = await request.json() + job_key = body.get("job_key", "") + job = _running_jobs.get(job_key) + if job and not job.done: + job.cancel_event.set() + return {"ok": True} + return {"ok": False, "message": "任务不存在或已完成"} + diff --git a/backend/app/backtest/optimizer.py b/backend/app/backtest/optimizer.py new file mode 100644 index 0000000..fdfbbc9 --- /dev/null +++ b/backend/app/backtest/optimizer.py @@ -0,0 +1,295 @@ +"""参数网格搜索优化器。 + +给定策略 + 参数网格, 遍历所有参数组合各跑一次回测, 按目标指标排序, 返回最优参数。 + +- 参数网格校验对齐 StrategyDef.meta["params"] (类型/范围/选项)。 +- 多线程并行执行, 复用 PanelCache: 同一 symbols/日期的面板只加载一次, 其余组合命中缓存。 +- 支持进度回调 (第 i/N 组完成) 与取消。 +""" +from __future__ import annotations + +import itertools +import logging +import threading +import time +from concurrent.futures import ThreadPoolExecutor, as_completed +from dataclasses import dataclass, field +from datetime import date + +logger = logging.getLogger(__name__) + +# 组合数硬上限 — 防止参数网格爆炸 (每组一次回测, 过大直接拒绝)。 +GRID_MAX_COMBINATIONS = 2000 + +# 需最小化的目标 (值越小越好); 其余默认最大化。 +# 注意: max_drawdown / mc_maxdd_* 为负值, 最大化其带符号值 = 回撤越小越好, 故仍归为 max。 +_MINIMIZE_OBJECTIVES = {"avg_holding_days"} + +# 可选优化目标 (须为 stats 中存在且数值可比的字段)。 +VALID_OBJECTIVES = { + "total_return", "annual_return", "sharpe", "sortino", "calmar", + "win_rate", "profit_factor", "max_drawdown", "mc_maxdd_p50", "mc_maxdd_p95", + "avg_pnl", "median_pnl", "n_trades", "avg_holding_days", +} + + +def _candidates_for(param_id: str, spec, pmeta: dict) -> list: + """从 grid spec 解析某参数的候选值列表并逐个校验。 + + spec 支持三种写法: + - list: 显式候选值 [v1, v2, ...] + - {"values": [...]}: 显式候选值 + - {"min", "max", "step"}: 数值型按步长展开 (含端点) + """ + p_type = pmeta["type"] + + # 解析原始候选值 + if isinstance(spec, list): + raw = spec + elif isinstance(spec, dict) and "values" in spec: + raw = spec["values"] + elif isinstance(spec, dict): + if p_type not in ("float", "int"): + raise ValueError(f"参数 '{param_id}' 为 {p_type} 型, 不支持 min/max/step 展开, 请给候选值列表") + step = spec.get("step") or pmeta.get("step") + if step is None or float(step) <= 0: + raise ValueError(f"参数 '{param_id}' 的 step 必须为正数") + lo = float(spec.get("min", pmeta.get("min", 0))) + hi = float(spec.get("max", pmeta.get("max", 0))) + if hi < lo: + raise ValueError(f"参数 '{param_id}' 的 max < min") + step = float(step) + # 整数计数生成候选, 避免浮点累加误差丢端点 (如 0.1/0.1 步长)。 + n_steps = round((hi - lo) / step) + raw = [round(lo + i * step, 10) for i in range(n_steps + 1)] + else: + raise ValueError(f"参数 '{param_id}' 的网格 spec 必须是列表或 {{min,max,step}} 字典") + + if not raw: + raise ValueError(f"参数 '{param_id}' 的候选值为空") + + # 逐值校验 + 归一化类型 + out = [] + for val in raw: + if p_type in ("float", "int"): + try: + num = float(val) + except (TypeError, ValueError): + raise ValueError(f"参数 '{param_id}' 的候选值 {val!r} 不是数字") from None + if pmeta.get("min") is not None and num < float(pmeta["min"]) - 1e-9: + raise ValueError(f"参数 '{param_id}' 的候选值 {val} 超出范围 (< min {pmeta['min']})") + if pmeta.get("max") is not None and num > float(pmeta["max"]) + 1e-9: + raise ValueError(f"参数 '{param_id}' 的候选值 {val} 超出范围 (> max {pmeta['max']})") + out.append(round(num) if p_type == "int" else num) + elif p_type == "bool": + out.append(bool(val)) + elif p_type == "select": + if val not in pmeta.get("options", []): + raise ValueError(f"参数 '{param_id}' 的候选值 {val!r} 不在 options {pmeta.get('options')} 中") + out.append(val) + else: + out.append(val) + # 去重保序 + seen = set() + uniq = [] + for v in out: + k = (type(v).__name__, v) + if k not in seen: + seen.add(k) + uniq.append(v) + return uniq + + +def _grid_candidates(params_meta: list[dict], param_grid: dict) -> dict[str, list]: + """校验整个 param_grid, 返回 {param_id: [候选值...]}。""" + if not param_grid: + raise ValueError("参数网格为空, 至少需要一个可扫参数") + by_id = {p["id"]: p for p in params_meta} + result: dict[str, list] = {} + for pid, spec in param_grid.items(): + if pid not in by_id: + raise ValueError(f"参数 '{pid}' 在该策略中不存在") + result[pid] = _candidates_for(pid, spec, by_id[pid]) + return result + + +def count_combinations(params_meta: list[dict], param_grid: dict) -> int: + """组合总数 (笛卡尔积), 用于爆炸预判。""" + cands = _grid_candidates(params_meta, param_grid) + total = 1 + for vals in cands.values(): + total *= len(vals) + return total + + +def expand_param_grid(params_meta: list[dict], param_grid: dict) -> list[dict]: + """校验并展开为参数组合列表, 每个组合是 {param_id: value} (仅含被扫参数)。 + + 超过 GRID_MAX_COMBINATIONS 直接拒绝。 + """ + cands = _grid_candidates(params_meta, param_grid) + total = 1 + for vals in cands.values(): + total *= len(vals) + if total > GRID_MAX_COMBINATIONS: + raise ValueError(f"参数组合数 {total} 超过上限 {GRID_MAX_COMBINATIONS}, 请增大 step 或缩小范围") + + keys = list(cands.keys()) + combos = [] + for values in itertools.product(*(cands[k] for k in keys)): + combos.append(dict(zip(keys, values, strict=True))) + return combos + + +def objective_value(stats: dict, objective: str, direction: str) -> float: + """从 stats 提取目标值并转为"越大越好"的可比分数 (None/缺失 -> 最差)。""" + raw = stats.get(objective) + if raw is None: + return float("-inf") + try: + v = float(raw) + except (TypeError, ValueError): + return float("-inf") + if v != v or v in (float("inf"), float("-inf")): # nan/inf + return float("-inf") + return -v if direction == "min" else v + + +def default_direction(objective: str) -> str: + return "min" if objective in _MINIMIZE_OBJECTIVES else "max" + + +# optimize 显式传入的 StrategyBacktestConfig 参数, backtest_kwargs 不得重复覆盖。 +_RESERVED_BT_KEYS = {"strategy_id", "symbols", "start", "end", "params", "overrides"} + + +def _validate_backtest_kwargs(backtest_kwargs: dict) -> None: + """校验 backtest_kwargs 的 key 合法且不与显式参数冲突, 否则会在 worker 线程抛 TypeError。""" + from dataclasses import fields + + from app.backtest.strategy import StrategyBacktestConfig + + valid = {f.name for f in fields(StrategyBacktestConfig)} - _RESERVED_BT_KEYS + for k in backtest_kwargs: + if k in _RESERVED_BT_KEYS: + raise ValueError(f"backtest_kwargs 不能包含 '{k}' (由优化器显式管理)") + if k not in valid: + raise ValueError(f"backtest_kwargs 含非法字段 '{k}', 合法: {sorted(valid)}") + + +@dataclass +class OptimizeConfig: + strategy_id: str + symbols: list[str] | None + start: date + end: date + param_grid: dict + objective: str = "sortino" + direction: str | None = None # None -> 由 objective 推断 + max_workers: int = 4 + base_params: dict = field(default_factory=dict) # 不扫的固定策略参数 + overrides: dict | None = None + backtest_kwargs: dict = field(default_factory=dict) # matching/fees/mode/initial_capital 等 + + +class StrategyOptimizer: + """遍历参数组合并行回测, 按目标排序。""" + + def __init__(self, service, strategy_engine) -> None: + self.service = service + self.strategy_engine = strategy_engine + + def optimize( + self, + cfg: OptimizeConfig, + progress_cb=None, + cancel_event: threading.Event | None = None, + ) -> dict: + from app.backtest.strategy import StrategyBacktestConfig + + t0 = time.perf_counter() + if cfg.objective not in VALID_OBJECTIVES: + raise ValueError(f"不支持的优化目标 '{cfg.objective}', 可选: {sorted(VALID_OBJECTIVES)}") + direction = cfg.direction or default_direction(cfg.objective) + _validate_backtest_kwargs(cfg.backtest_kwargs) + + s = self.strategy_engine.get(cfg.strategy_id) # 可能抛 ValueError + params_meta = s.meta.get("params", []) + combos = expand_param_grid(params_meta, cfg.param_grid) + n_total = len(combos) + + results: list[dict] = [] + done = 0 + lock = threading.Lock() + + def _run_one(idx: int, combo: dict) -> dict | None: + if cancel_event is not None and cancel_event.is_set(): + return None + # 单组异常必须隔离: 加了并行后, 一组抛异常若冒泡会拖垮整批 (丢弃全部已完成结果)。 + try: + merged = {**cfg.base_params, **combo} + bt_cfg = StrategyBacktestConfig( + strategy_id=cfg.strategy_id, + symbols=cfg.symbols, + start=cfg.start, + end=cfg.end, + params=merged, + overrides=cfg.overrides, + **cfg.backtest_kwargs, + ) + res = self.service.run(bt_cfg, cancel_event=cancel_event) + except Exception as e: # 隔离单组失败, 记录后继续, 不拖垮整批 + logger.warning("参数组 %s 回测异常: %r", combo, e) + return {"params": combo, "error": repr(e), "objective_raw": None, "_sort": float("-inf")} + if res.error: + return {"params": combo, "error": res.error, "objective_raw": None, "_sort": float("-inf")} + # _sort: 内部排序键 (统一"越大越好"); objective_raw: 原始展示值 (不受方向取负污染)。 + return { + "params": combo, + "objective_raw": res.stats.get(cfg.objective), + "_sort": objective_value(res.stats, cfg.objective, direction), + "stats": res.stats, + } + + def _best_raw() -> float | None: + if not results: + return None + top = max(results, key=lambda x: x["_sort"]) + return None if top["_sort"] == float("-inf") else top.get("objective_raw") + + max_workers = max(1, min(int(cfg.max_workers), n_total)) + with ThreadPoolExecutor(max_workers=max_workers) as pool: + futures = {pool.submit(_run_one, i, c): i for i, c in enumerate(combos)} + for fut in as_completed(futures): + r = fut.result() # _run_one 内部已兜底, 不会 re-raise 业务异常 + with lock: + done += 1 + if r is not None: + results.append(r) + if progress_cb is not None: + br = _best_raw() + progress_cb({ + "type": "optimizer_progress", + "done": done, + "total": n_total, + "best_score": round(br, 4) if br is not None else None, + }) + + # 排序: 内部 _sort 降序 (越大越好); -inf (失败/无效) 沉底。展示层用 objective_raw。 + ranked = sorted(results, key=lambda x: x["_sort"], reverse=True) + for i, r in enumerate(ranked): + r["rank"] = i + 1 + r.pop("_sort", None) # 不外露内部排序键, 避免展示层误用取负值 + + best = ranked[0] if ranked and ranked[0].get("objective_raw") is not None else None + best_raw = best["objective_raw"] if best else None + return { + "objective": cfg.objective, + "direction": direction, + "n_combinations": n_total, + "n_completed": len(results), + "best_params": best["params"] if best else None, + "best_score": round(best_raw, 4) if best_raw is not None else None, + "results": ranked, + "elapsed_ms": round((time.perf_counter() - t0) * 1000, 1), + } diff --git a/backend/tests/backtest/test_optimizer_api.py b/backend/tests/backtest/test_optimizer_api.py new file mode 100644 index 0000000..d21dc9a --- /dev/null +++ b/backend/tests/backtest/test_optimizer_api.py @@ -0,0 +1,60 @@ +"""优化器 API job_key 契约测试 — 守护 stream 与 cancel 的 key 对齐 (仿 PR3 C1 教训)。""" +from __future__ import annotations + +from app.api.backtest import _OPT_BT_FIELDS, _make_opt_job_key, _opt_backtest_kwargs + + +def _sig(bt: dict) -> str: + return "|".join(f"{k}={bt[k]}" for k in _OPT_BT_FIELDS) + + +def test_job_key_deterministic(): + bt = _opt_backtest_kwargs("open_t+1", 0.0002, None, None, 5.0, 10, 1.0, 1e6, "equal", "position", 5) + sig = _sig(bt) + k1 = _make_opt_job_key("s", None, None, None, '{"p":[1,2]}', "sortino", None, sig) + k2 = _make_opt_job_key("s", None, None, None, '{"p":[1,2]}', "sortino", None, sig) + assert k1 == k2 + + +def test_job_key_distinguishes_grid_and_objective(): + bt = _opt_backtest_kwargs("open_t+1", 0.0002, None, None, 5.0, 10, 1.0, 1e6, "equal", "position", 5) + sig = _sig(bt) + base = _make_opt_job_key("s", None, None, None, '{"p":[1,2]}', "sortino", None, sig) + assert base != _make_opt_job_key("s", None, None, None, '{"p":[1,3]}', "sortino", None, sig) # grid 不同 + assert base != _make_opt_job_key("s", None, None, None, '{"p":[1,2]}', "sharpe", None, sig) # objective 不同 + + +def test_cancel_looks_up_job_by_echoed_key(): + """重构后: cancel 直接用 stream 回吐的 job_key 查表, 不再重算参数。 + + 这消除了'两侧重算必须逐字段一致'的脆弱契约 (PR3 C1 / direction 空串失配的根因)。 + """ + import asyncio + + from app.api.backtest import _BacktestJob, _running_jobs, optimize_cancel + + class _Req: + def __init__(self, body): + self._body = body + async def json(self): + return self._body + + key = "optkey_test_1" + job = _BacktestJob(key) + _running_jobs[key] = job + try: + # 用回吐的 key 取消 → 命中并 set cancel_event + res = asyncio.run(optimize_cancel(_Req({"job_key": key}))) + assert res["ok"] is True + assert job.cancel_event.is_set() + + # 已完成任务再取消 → ok False + job.done = True + res2 = asyncio.run(optimize_cancel(_Req({"job_key": key}))) + assert res2["ok"] is False + + # 未知 key → ok False, 不抛异常 + res3 = asyncio.run(optimize_cancel(_Req({"job_key": "nonexistent"}))) + assert res3["ok"] is False + finally: + _running_jobs.pop(key, None) diff --git a/backend/tests/backtest/test_optimizer_grid.py b/backend/tests/backtest/test_optimizer_grid.py new file mode 100644 index 0000000..c1527b6 --- /dev/null +++ b/backend/tests/backtest/test_optimizer_grid.py @@ -0,0 +1,125 @@ +"""参数网格展开与校验测试 — 优化器的纯逻辑核心。 + +被测: +- expand_param_grid(params_meta, param_grid): 校验 + 笛卡尔积 -> 参数组合列表 +- count_combinations(params_meta, param_grid): 组合数 (不真正展开, 用于爆炸预判) +""" +from __future__ import annotations + +import pytest + +from app.backtest.optimizer import ( + GRID_MAX_COMBINATIONS, + count_combinations, + expand_param_grid, +) + +# 模拟一个策略的 params meta (对齐 StrategyDef.meta["params"] 结构) +PARAMS_META = [ + {"id": "ma_proximity", "type": "float", "default": 0.02, "min": 0.01, "max": 0.05, "step": 0.005}, + {"id": "min_boards", "type": "int", "default": 2, "min": 1, "max": 20, "step": 1}, + {"id": "use_ma20", "type": "bool", "default": True}, + {"id": "fill", "type": "select", "default": "close_t", "options": ["close_t", "open_t+1"]}, +] + + +# --------------------------------------------------------------- +# 显式候选值列表 +# --------------------------------------------------------------- + +def test_explicit_value_lists_cartesian_product(): + grid = {"ma_proximity": [0.01, 0.02], "min_boards": [2, 3, 4]} + combos = expand_param_grid(PARAMS_META, grid) + assert len(combos) == 6 # 2 x 3 + assert {"ma_proximity": 0.01, "min_boards": 2} in combos + assert {"ma_proximity": 0.02, "min_boards": 4} in combos + + +def test_single_param_sweep(): + combos = expand_param_grid(PARAMS_META, {"min_boards": [1, 5, 10]}) + assert combos == [{"min_boards": 1}, {"min_boards": 5}, {"min_boards": 10}] + + +def test_bool_and_select_sweep(): + grid = {"use_ma20": [True, False], "fill": ["close_t", "open_t+1"]} + combos = expand_param_grid(PARAMS_META, grid) + assert len(combos) == 4 + + +# --------------------------------------------------------------- +# 范围 spec {min,max,step} 自动展开 +# --------------------------------------------------------------- + +def test_range_spec_expands_by_step(): + combos = expand_param_grid(PARAMS_META, {"ma_proximity": {"min": 0.01, "max": 0.03, "step": 0.01}}) + vals = sorted(c["ma_proximity"] for c in combos) + assert vals == [0.01, 0.02, 0.03] # 含端点 + + +def test_range_spec_float_keeps_endpoint_despite_accumulation(): + """0.1 步长的浮点累加易丢端点 (0.1+0.1+0.1=0.30000004); 整数计数必须保住 0.3。""" + meta = [{"id": "p", "type": "float", "default": 0.2, "min": 0.1, "max": 0.3, "step": 0.1}] + combos = expand_param_grid(meta, {"p": {"min": 0.1, "max": 0.3, "step": 0.1}}) + vals = sorted(c["p"] for c in combos) + assert vals == [0.1, 0.2, 0.3] # 含端点 0.3, 不丢 + + +def test_duplicate_values_folded(): + combos = expand_param_grid(PARAMS_META, {"ma_proximity": [0.02, 0.02, 0.03]}) + vals = sorted(c["ma_proximity"] for c in combos) + assert vals == [0.02, 0.03] # 去重 + + +def test_range_spec_int_yields_ints(): + combos = expand_param_grid(PARAMS_META, {"min_boards": {"min": 1, "max": 4, "step": 1}}) + vals = sorted(c["min_boards"] for c in combos) + assert vals == [1, 2, 3, 4] + assert all(isinstance(v, int) for v in vals) + + +# --------------------------------------------------------------- +# 校验: 拒绝非法 grid +# --------------------------------------------------------------- + +def test_unknown_param_rejected(): + with pytest.raises(ValueError, match="不存在"): + expand_param_grid(PARAMS_META, {"nonexistent": [1, 2]}) + + +def test_value_out_of_range_rejected(): + with pytest.raises(ValueError, match=r"超出范围|范围"): + expand_param_grid(PARAMS_META, {"ma_proximity": [0.01, 0.99]}) + + +def test_select_value_not_in_options_rejected(): + with pytest.raises(ValueError, match=r"options|选项"): + expand_param_grid(PARAMS_META, {"fill": ["close_t", "bad_value"]}) + + +def test_empty_grid_rejected(): + with pytest.raises(ValueError, match=r"空|至少"): + expand_param_grid(PARAMS_META, {}) + + +def test_combination_explosion_rejected(): + # 构造超过硬上限的组合 + big = {"ma_proximity": {"min": 0.01, "max": 0.05, "step": 0.001}} # 41 个 + # 单参数 41 个不会爆; 用多参数放大 + grid = { + "ma_proximity": {"min": 0.01, "max": 0.05, "step": 0.001}, # 41 + "min_boards": {"min": 1, "max": 20, "step": 1}, # 20 + } # 41 x 20 = 820, 仍 < 2000; 再加一维 + grid["use_ma20"] = [True, False] # x2 = 1640 + # 到这仍 < 2000, 断言 count 正确 + assert count_combinations(PARAMS_META, grid) == 1640 + assert count_combinations(PARAMS_META, big) == 41 + # 显式超限 + huge = { + "ma_proximity": {"min": 0.01, "max": 0.05, "step": 0.001}, # 41 + "min_boards": {"min": 1, "max": 20, "step": 1}, # 20 + "fill": ["close_t", "open_t+1"], # 2 + "use_ma20": [True, False], # 2 + } # 41x20x2x2 = 3280 > 2000 + assert count_combinations(PARAMS_META, huge) > GRID_MAX_COMBINATIONS + with pytest.raises(ValueError, match=r"组合数|上限|超过"): + expand_param_grid(PARAMS_META, huge) diff --git a/backend/tests/backtest/test_optimizer_run.py b/backend/tests/backtest/test_optimizer_run.py new file mode 100644 index 0000000..8008c4b --- /dev/null +++ b/backend/tests/backtest/test_optimizer_run.py @@ -0,0 +1,199 @@ +"""优化器编排测试 — 用假 service 注入受控 stats, 验证排序/取消/进度/目标方向。""" +from __future__ import annotations + +import threading +from dataclasses import dataclass +from datetime import date + +import pytest + +from app.backtest.optimizer import OptimizeConfig, StrategyOptimizer + +# ---- 假 StrategyDef / 引擎 / service ---- + +@dataclass +class _FakeDef: + meta: dict + + +class _FakeEngine: + def __init__(self, params_meta): + self._def = _FakeDef(meta={"params": params_meta}) + + def get(self, strategy_id): + return self._def + + +@dataclass +class _FakeResult: + stats: dict + error: str | None = None + + +class _FakeService: + """run() 依据 params 返回受控 stats: sortino = ma_proximity 的映射, 便于校验排序。""" + + def __init__(self, score_fn): + self.score_fn = score_fn + self.calls = [] + self._lock = threading.Lock() + + def run(self, config, progress_cb=None, cancel_event=None): + with self._lock: + self.calls.append(dict(config.params or {})) + return self.score_fn(config.params or {}) + + +PARAMS_META = [ + {"id": "ma_proximity", "type": "float", "default": 0.02, "min": 0.01, "max": 0.05, "step": 0.005}, +] + + +def _optimizer(score_fn): + return StrategyOptimizer(_FakeService(score_fn), _FakeEngine(PARAMS_META)) + + +def _cfg(**kw): + base = dict( + strategy_id="s", symbols=None, start=date(2024, 1, 1), end=date(2024, 6, 1), + param_grid={"ma_proximity": [0.01, 0.02, 0.03]}, objective="sortino", max_workers=4, + ) + base.update(kw) + return OptimizeConfig(**base) + + +def test_ranks_best_by_objective_max(): + # sortino 随 ma_proximity 递增 -> 最大值应为 0.03 + def score(p): + return _FakeResult(stats={"sortino": p["ma_proximity"] * 100}) + out = _optimizer(score).optimize(_cfg()) + assert out["best_params"] == {"ma_proximity": 0.03} + assert out["best_score"] == 3.0 + assert out["n_combinations"] == 3 + assert out["n_completed"] == 3 + assert [r["rank"] for r in out["results"]] == [1, 2, 3] + assert out["results"][0]["params"] == {"ma_proximity": 0.03} + + +def test_all_combos_executed_once(): + def score(p): + return _FakeResult(stats={"sortino": 1.0}) + opt = _optimizer(score) + out = opt.optimize(_cfg(param_grid={"ma_proximity": [0.01, 0.02, 0.03, 0.04, 0.05]})) + assert out["n_combinations"] == 5 + # 每组恰跑一次 + ran = sorted(c["ma_proximity"] for c in opt.service.calls) + assert ran == [0.01, 0.02, 0.03, 0.04, 0.05] + + +def test_min_direction_objective_restores_display_sign(): + # avg_holding_days 是 min 方向: 最小者最优, 且 best_score 必须是原始正值 (非内部取负值) + def score(p): + return _FakeResult(stats={"avg_holding_days": p["ma_proximity"] * 100}) + out = _optimizer(score).optimize(_cfg(objective="avg_holding_days")) + assert out["best_params"] == {"ma_proximity": 0.01} + # min 方向: 最优 avg_holding_days = 0.01*100 = 1.0, 用户应看到 +1.0 而非 -1.0 + assert out["best_score"] == 1.0 + # results 不应外露内部排序键 _sort + assert all("_sort" not in r for r in out["results"]) + assert out["results"][0]["objective_raw"] == 1.0 + + +def test_max_drawdown_objective_prefers_smaller_drawdown(): + # max_drawdown 为负值, max 方向: -0.1 (回撤更小) 应优于 -0.3 + def score(p): + dd = {0.01: -0.1, 0.02: -0.3, 0.03: -0.2}[p["ma_proximity"]] + return _FakeResult(stats={"max_drawdown": dd}) + out = _optimizer(score).optimize(_cfg(objective="max_drawdown")) + assert out["best_params"] == {"ma_proximity": 0.01} + assert out["best_score"] == -0.1 # 展示原始负值 + + +def test_service_exception_isolated_not_crashing_batch(): + # 某组 service.run 抛异常 -> 应记为该组失败, 其余组正常完成, 不拖垮整批 + def score(p): + if p["ma_proximity"] == 0.02: + raise KeyError("boom") + return _FakeResult(stats={"sortino": p["ma_proximity"] * 100}) + out = _optimizer(score).optimize(_cfg()) + assert out["n_completed"] == 3 # 三组都有结果记录 (含失败组) + assert out["best_params"] == {"ma_proximity": 0.03} # 最优组不受影响 + failed = [r for r in out["results"] if r.get("error")] + assert len(failed) == 1 + assert "boom" in failed[0]["error"] + + +def test_backtest_kwargs_illegal_key_rejected(): + def score(p): + return _FakeResult(stats={"sortino": 1.0}) + with pytest.raises(ValueError, match=r"非法字段|不能包含"): + _optimizer(score).optimize(_cfg(backtest_kwargs={"bad_field": 1})) + + +def test_backtest_kwargs_reserved_key_rejected(): + def score(p): + return _FakeResult(stats={"sortino": 1.0}) + with pytest.raises(ValueError, match="不能包含"): + _optimizer(score).optimize(_cfg(backtest_kwargs={"symbols": ["x"]})) + + +def test_base_params_merged_and_overridden_by_sweep(): + # base_params 提供固定参数, combo 覆盖同名; 记录 service 实际收到的 params + def score(p): + return _FakeResult(stats={"sortino": 1.0}) + opt = _optimizer(score) + opt.optimize(_cfg(base_params={"ma_proximity": 0.99, "other": 7})) + # 每次 run 收到的 params: ma_proximity 被 combo 覆盖, other 保留 + for call in opt.service.calls: + assert call["other"] == 7 + assert call["ma_proximity"] in (0.01, 0.02, 0.03) + + +def test_none_and_error_results_sink_to_bottom(): + # ma_proximity=0.02 的组返回 error, 0.03 的 sortino=None -> 都应排在有效结果之后 + def score(p): + if p["ma_proximity"] == 0.02: + return _FakeResult(stats={}, error="boom") + if p["ma_proximity"] == 0.03: + return _FakeResult(stats={"sortino": None}) + return _FakeResult(stats={"sortino": 5.0}) + out = _optimizer(score).optimize(_cfg()) + assert out["best_params"] == {"ma_proximity": 0.01} + assert out["best_score"] == 5.0 + # 失败/None 组仍在结果里但 rank 靠后 + assert out["n_completed"] == 3 + assert out["results"][0]["params"] == {"ma_proximity": 0.01} + + +def test_cancel_event_stops_remaining(): + ev = threading.Event() + ev.set() # 一开始就取消 + + def score(p): + return _FakeResult(stats={"sortino": 1.0}) + opt = _optimizer(score) + out = opt.optimize(_cfg(), cancel_event=ev) + # 取消后所有组跳过 -> 无有效结果 + assert opt.service.calls == [] + assert out["best_params"] is None + + +def test_progress_callback_reports_done_total(): + seen = [] + + def score(p): + return _FakeResult(stats={"sortino": 1.0}) + + def cb(msg): + seen.append(msg) + _optimizer(score).optimize(_cfg(), progress_cb=cb) + assert len(seen) == 3 + assert seen[-1]["done"] == 3 + assert all(m["total"] == 3 for m in seen) + + +def test_invalid_objective_rejected(): + def score(p): + return _FakeResult(stats={"sortino": 1.0}) + with pytest.raises(ValueError, match="不支持的优化目标"): + _optimizer(score).optimize(_cfg(objective="not_a_metric")) diff --git a/frontend/src/lib/optimizerTask.ts b/frontend/src/lib/optimizerTask.ts new file mode 100644 index 0000000..5c6beae --- /dev/null +++ b/frontend/src/lib/optimizerTask.ts @@ -0,0 +1,274 @@ +import { useSyncExternalStore } from 'react' + +/** + * 参数优化任务管理 (SSE 模式 + 重连)。镜像 backtestTask, 结果为排名 dict。 + */ + +export interface OptimizeProgress { + type: string + done: number + total: number + best_score: number | null +} + +export interface OptimizeResultRow { + params: Record + objective_raw: number | null + stats?: Record + rank: number + error?: string +} + +export interface OptimizeResult { + objective: string + direction: string + n_combinations: number + n_completed: number + best_params: Record | null + best_score: number | null + results: OptimizeResultRow[] + elapsed_ms: number +} + +export interface OptimizerTask { + id: number + isPending: boolean + result: OptimizeResult | null + progress: OptimizeProgress | null + error: string | null +} + +export interface StartOptimizeParams { + strategy_id: string + param_grid: Record + objective: string + direction?: string + max_workers?: number + params?: Record | null // 未扫描参数固定为用户当前值 + overrides?: Record | null // 策略当前的 basic_filter/信号/风控覆盖 + symbols?: string[] | null + start?: string | null + end?: string | null + matching?: string + fees_pct?: number + commission_pct?: number + stamp_tax_pct?: number + slippage_bps?: number + max_positions?: number + max_exposure_pct?: number + initial_capital?: number + position_sizing?: string + mode?: 'position' | 'full' + holding_days?: number +} + +let current: OptimizerTask | null = null +const listeners = new Set<() => void>() +let taskSeq = 0 +let eventSource: EventSource | null = null +let currentJobKey: string | null = null +let cancelRequested = false // stop 在拿到 job_key 前被点 -> 收到 job 事件立即补发 cancel +let reconnectAttempts = 0 // 无 data 断线的连续重连计数, 超上限放弃 +const MAX_RECONNECT = 5 + +const RECONNECT_KEY = 'optimizer_reconnect' +const JOB_KEY_KEY = 'optimizer_job_key' + +function emit() { + listeners.forEach(fn => fn()) +} + +function subscribe(fn: () => void) { + listeners.add(fn) + return () => listeners.delete(fn) +} + +function buildQuery(params: Record): string { + const sp = new URLSearchParams() + for (const [k, v] of Object.entries(params)) { + if (v != null && v !== '') sp.set(k, String(v)) + } + return sp.toString() +} + +function connectSSE(url: string): void { + const id = current?.id ?? ++taskSeq + + if (eventSource) { + eventSource.close() + eventSource = null + } + + const es = new EventSource(url) + eventSource = es + + // 首事件: 后端回吐 job_key, 存下供 cancel 直接引用 (无需前端重算) + es.addEventListener('job', (e: MessageEvent) => { + reconnectAttempts = 0 + try { + const key = JSON.parse(e.data)?.key + if (key) { + currentJobKey = key + localStorage.setItem(JOB_KEY_KEY, key) + // 竞态修复: stop 在拿到 key 前被点过 -> 现在补发 cancel 真正停后端任务, 再收尾关闭。 + if (cancelRequested) { + postCancel(key) + es.close() + eventSource = null + currentJobKey = null + localStorage.removeItem(RECONNECT_KEY) + localStorage.removeItem(JOB_KEY_KEY) + } + } + } catch { /* ignore */ } + }) + + es.addEventListener('progress', (e: MessageEvent) => { + if (current?.id !== id) return + reconnectAttempts = 0 + try { + const prog = JSON.parse(e.data) as OptimizeProgress + current = { ...current, progress: prog } + emit() + } catch { /* ignore */ } + }) + + es.addEventListener('done', (e: MessageEvent) => { + if (current?.id !== id) return + try { + const result = JSON.parse(e.data) as OptimizeResult + current = { ...current, isPending: false, result, error: null } + emit() + } catch { + current = { ...current, isPending: false, error: '结果解析失败' } + emit() + } + es.close() + eventSource = null + currentJobKey = null + localStorage.removeItem(RECONNECT_KEY) + localStorage.removeItem(JOB_KEY_KEY) + }) + + es.addEventListener('error', (e: MessageEvent) => { + if (current?.id !== id) return + if (e.data) { + try { + const msg = JSON.parse(e.data)?.message ?? '优化出错' + current = { ...current, isPending: false, error: msg } + emit() + } catch { + current = { ...current, isPending: false, error: '优化出错' } + emit() + } + es.close() + eventSource = null + currentJobKey = null + localStorage.removeItem(RECONNECT_KEY) + localStorage.removeItem(JOB_KEY_KEY) + return + } + // 无 data: 连接异常断开。EventSource 会自动重连, 但设上限避免网络长断时无限 pending。 + if (current?.id === id) { + reconnectAttempts += 1 + if (reconnectAttempts > MAX_RECONNECT) { + es.close() + eventSource = null + current = { ...current, isPending: false, error: '连接中断, 重连多次失败' } + emit() + } + } + }) +} + +/** 调后端 cancel (按回吐的 job_key)。 */ +function postCancel(jobKey: string): void { + fetch('/api/backtest/optimize/cancel', { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ job_key: jobKey }), + }).catch(() => {}) +} + +export function startOptimize(params: StartOptimizeParams): void { + if (eventSource) { + eventSource.close() + eventSource = null + } + + cancelRequested = false + currentJobKey = null + reconnectAttempts = 0 + const id = ++taskSeq + current = { id, isPending: true, result: null, progress: null, error: null } + emit() + + const qs = buildQuery({ + strategy_id: params.strategy_id, + param_grid: JSON.stringify(params.param_grid), + objective: params.objective, + direction: params.direction, + max_workers: params.max_workers, + params: params.params ? JSON.stringify(params.params) : undefined, + overrides: params.overrides ? JSON.stringify(params.overrides) : undefined, + symbols: params.symbols?.join(','), + start: params.start ?? undefined, + end: params.end ?? undefined, + matching: params.matching, + fees_pct: params.fees_pct, + commission_pct: params.commission_pct, + stamp_tax_pct: params.stamp_tax_pct, + slippage_bps: params.slippage_bps, + max_positions: params.max_positions, + max_exposure_pct: params.max_exposure_pct, + initial_capital: params.initial_capital, + position_sizing: params.position_sizing, + mode: params.mode, + holding_days: params.holding_days, + }) + + localStorage.setItem(RECONNECT_KEY, qs) + connectSSE(`/api/backtest/optimize/stream?${qs}`) +} + +export function stopOptimize(): void { + // 竞态: 若刚点开始还没收到 job 事件, job_key 尚未到手。标记 cancelRequested —— + // 有 key 则立即取消并关闭; 无 key 则保持 SSE 打开, 等 job 事件到达时补发 cancel 再关 + // (关闭 SSE 不会停后端 daemon 线程, 必须真正 POST cancel)。5s 兜底防 job 事件永不来。 + cancelRequested = true + const jobKey = currentJobKey ?? localStorage.getItem(JOB_KEY_KEY) + if (jobKey) { + postCancel(jobKey) + if (eventSource) { eventSource.close(); eventSource = null } + currentJobKey = null + localStorage.removeItem(RECONNECT_KEY) + localStorage.removeItem(JOB_KEY_KEY) + } else if (eventSource) { + // 保持连接等 job 事件; 兜底: 5s 后仍没 key 就强关 + const es = eventSource + setTimeout(() => { if (es === eventSource) { es.close(); eventSource = null } }, 5000) + } + if (current?.isPending) { + current = { ...current, isPending: false, error: '已取消' } + emit() + } +} + +export function clearOptimize(): void { + current = null + emit() +} + +export function tryReconnectOptimize(): boolean { + const qs = localStorage.getItem(RECONNECT_KEY) + if (!qs) return false + const id = ++taskSeq + current = { id, isPending: true, result: null, progress: null, error: null } + emit() + connectSSE(`/api/backtest/optimize/stream?${qs}`) + return true +} + +export function useOptimizerTask(): OptimizerTask | null { + return useSyncExternalStore(subscribe, () => current, () => null) +} diff --git a/frontend/src/lib/strategyOverrides.ts b/frontend/src/lib/strategyOverrides.ts new file mode 100644 index 0000000..23106fd --- /dev/null +++ b/frontend/src/lib/strategyOverrides.ts @@ -0,0 +1,24 @@ +import type { StrategyDetail } from './api' + +/** 信号 id 归一 (与策略回测页一致): 裸名补 signal_ 前缀。 */ +export const toSignalId = (sig: string) => + sig.startsWith('signal_') || sig.startsWith('csg_') ? sig : `signal_${sig}` + +/** 从策略详情构建默认 overrides (basic_filter / 信号 / 风控)。 + * 优化器与策略回测页共用, 保证优化的就是用户当前配置的策略。 */ +export function buildDefaultOverrides(detail: StrategyDetail): Record { + return { + basic_filter: { ...detail.basic_filter }, + entry_signals: detail.entry_signals.map(toSignalId), + exit_signals: detail.exit_signals.map(toSignalId), + scoring: { ...detail.scoring }, + stop_loss: detail.stop_loss, + take_profit: detail.take_profit, + trailing_stop: detail.trailing_stop, + trailing_take_profit_activate: detail.trailing_take_profit_activate, + trailing_take_profit_drawdown: detail.trailing_take_profit_drawdown, + score_min: null, + score_max: null, + max_hold_days: detail.max_hold_days, + } +} diff --git a/frontend/src/pages/Backtest.tsx b/frontend/src/pages/Backtest.tsx index 08ab0d7..c496e22 100644 --- a/frontend/src/pages/Backtest.tsx +++ b/frontend/src/pages/Backtest.tsx @@ -2,9 +2,10 @@ import { useState } from 'react' import { PageHeader } from '@/components/PageHeader' import { FactorBacktest } from './backtest/FactorBacktest' import { StrategyBacktest } from './backtest/StrategyBacktest' -import { BarChart3, FlaskConical } from 'lucide-react' +import { StrategyOptimizer } from './backtest/StrategyOptimizer' +import { BarChart3, FlaskConical, SlidersHorizontal } from 'lucide-react' -type Tab = 'factor' | 'strategy' +type Tab = 'factor' | 'strategy' | 'optimizer' const MODES: Record = { factor: { @@ -17,6 +18,17 @@ const MODES: Record = { subtitle: '验证完整选股和交易规则', hint: '看净值曲线、回撤、胜率和交易明细,适合判断策略是否可执行。', }, + optimizer: { + title: '参数优化', + subtitle: '网格搜索最优参数组合', + hint: '并行回测所有参数组合,按夏普/索提诺等目标排序,找到最优参数。', + }, +} + +const TAB_ICONS: Record = { + factor: BarChart3, + strategy: FlaskConical, + optimizer: SlidersHorizontal, } export function Backtest() { @@ -24,8 +36,8 @@ export function Backtest() { const modeSwitch = (
- {(['factor', 'strategy'] as const).map(tab => { - const Icon = tab === 'factor' ? BarChart3 : FlaskConical + {(['factor', 'strategy', 'optimizer'] as const).map(tab => { + const Icon = TAB_ICONS[tab] const active = activeTab === tab return (
) diff --git a/frontend/src/pages/backtest/StrategyOptimizer.tsx b/frontend/src/pages/backtest/StrategyOptimizer.tsx new file mode 100644 index 0000000..c67c514 --- /dev/null +++ b/frontend/src/pages/backtest/StrategyOptimizer.tsx @@ -0,0 +1,348 @@ +import { useEffect, useMemo, useState } from 'react' +import { useQuery } from '@tanstack/react-query' +import { Play, Square, Trophy } from 'lucide-react' +import { api, type StrategyDetail, type StrategyParamDef } from '@/lib/api' +import { fmtPct } from '@/lib/format' +import { EmptyState } from '@/components/EmptyState' +import { DatePicker } from '@/components/DatePicker' +import { + startOptimize, + stopOptimize, + clearOptimize, + tryReconnectOptimize, + useOptimizerTask, +} from '@/lib/optimizerTask' +import { buildDefaultOverrides } from '@/lib/strategyOverrides' + +const INPUT_CLS = 'w-full px-2.5 py-1.5 rounded-input bg-surface border border-border text-xs focus:outline-none focus:border-accent' + +// 可选优化目标 (对齐后端 VALID_OBJECTIVES) + 中文标签 + 是否越小越好 +const OBJECTIVES: { id: string; label: string; min?: boolean }[] = [ + { id: 'sortino', label: '索提诺比率' }, + { id: 'sharpe', label: '夏普比率' }, + { id: 'calmar', label: 'Calmar 比率' }, + { id: 'total_return', label: '总收益' }, + { id: 'annual_return', label: '年化收益' }, + { id: 'win_rate', label: '胜率' }, + { id: 'profit_factor', label: '盈亏比' }, + { id: 'max_drawdown', label: '最大回撤(越小越好)' }, + { id: 'mc_maxdd_p95', label: '蒙卡回撤P95(越小越好)' }, + { id: 'avg_holding_days', label: '平均持仓天数', min: true }, +] + +// 单个可扫参数的网格配置 +interface Sweep { + enabled: boolean + min: string + max: string + step: string +} + +function defaultSweep(p: StrategyParamDef): Sweep { + return { + enabled: false, + min: String(p.min ?? p.default ?? 0), + max: String(p.max ?? p.default ?? 1), + step: String(p.step ?? (p.type === 'int' ? 1 : 0.01)), + } +} + +/** 从 sweep 配置估算某参数候选值个数 (与后端整数计数一致) */ +function candidateCount(p: StrategyParamDef, s: Sweep): number { + if (p.type === 'bool') return 2 + if (p.type === 'select') return p.options?.length ?? 1 + const lo = Number(s.min), hi = Number(s.max), step = Number(s.step) + if (!(step > 0) || hi < lo) return 0 + return Math.round((hi - lo) / step) + 1 +} + +/** 校验某数值参数的 sweep 是否会被后端拒绝 (与后端 _candidates_for 同口径)。 + * 后端按 lo+i*step 生成 (i=0..round((hi-lo)/step)), 任一值超出 [min,max] 即报错。 */ +function sweepError(p: StrategyParamDef, s: Sweep): string | null { + if (p.type === 'bool' || p.type === 'select') return null + const lo = Number(s.min), hi = Number(s.max), step = Number(s.step) + if (Number.isNaN(lo) || Number.isNaN(hi) || Number.isNaN(step)) return `${p.label}: 范围/步长非法` + if (!(step > 0)) return `${p.label}: 步长必须为正` + if (hi < lo) return `${p.label}: max < min` + if (p.min != null && lo < p.min - 1e-9) return `${p.label}: min 小于允许下限 ${p.min}` + if (p.max != null && hi > p.max + 1e-9) return `${p.label}: max 超出允许上限 ${p.max}` + // 后端生成的末值 lo + round((hi-lo)/step)*step 若 > max, 会被拒 + const nSteps = Math.round((hi - lo) / step) + const last = lo + nSteps * step + if (last > hi + 1e-9) return `${p.label}: 步长 ${step} 不整除区间, 末值 ${last.toFixed(4)} 超出 max ${hi}` + return null +} + +const TODAY = new Date().toISOString().slice(0, 10) +const ONE_YEAR_AGO = new Date(Date.now() - 365 * 864e5).toISOString().slice(0, 10) + +export function StrategyOptimizer() { + const task = useOptimizerTask() + const { data: stratData } = useQuery({ queryKey: ['strategies'], queryFn: api.strategyList }) + const strategies: StrategyDetail[] = stratData?.strategies ?? [] + + const [strategyId, setStrategyId] = useState('') + const [objective, setObjective] = useState('sortino') + const [start, setStart] = useState(ONE_YEAR_AGO) + const [end, setEnd] = useState(TODAY) + const [mode, setMode] = useState<'position' | 'full'>('position') + const [sweeps, setSweeps] = useState>({}) + + const selected = strategies.find(s => s.id === strategyId) + const params = selected?.params ?? [] + + // 刷新/切页后: 恢复未完成的优化任务 + useEffect(() => { + tryReconnectOptimize() + // eslint-disable-next-line react-hooks/exhaustive-deps + }, []) + + // 切策略: 若有任务在跑, 先真正取消 (关 SSE + 后端 cancel + 清 localStorage), 不能静默丢。 + const onSelectStrategy = (id: string) => { + if (task?.isPending) stopOptimize() + else clearOptimize() + setStrategyId(id) + const s = strategies.find(x => x.id === id) + const init: Record = {} + for (const p of s?.params ?? []) init[p.id] = defaultSweep(p) + setSweeps(init) + } + + const updateSweep = (pid: string, patch: Partial) => + setSweeps(prev => ({ ...prev, [pid]: { ...prev[pid], ...patch } })) + + // 组合数预估 + const combos = useMemo(() => { + const enabled = params.filter(p => sweeps[p.id]?.enabled) + if (!enabled.length) return 0 + return enabled.reduce((acc, p) => acc * candidateCount(p, sweeps[p.id]), 1) + }, [params, sweeps]) + + // 网格合法性 (与后端展开同口径): 步长不整除/越界会被后端拒, 前端提前拦。 + const gridError = useMemo(() => { + for (const p of params) { + if (!sweeps[p.id]?.enabled) continue + const err = sweepError(p, sweeps[p.id]) + if (err) return err + } + return null + }, [params, sweeps]) + + const buildGrid = (): Record => { + const grid: Record = {} + for (const p of params) { + const s = sweeps[p.id] + if (!s?.enabled) continue + if (p.type === 'bool') grid[p.id] = [true, false] + else if (p.type === 'select') grid[p.id] = p.options ?? [] + else grid[p.id] = { min: Number(s.min), max: Number(s.max), step: Number(s.step) } + } + return grid + } + + const canRun = strategyId && combos > 0 && combos <= 2000 && !gridError && !task?.isPending + + const onRun = () => { + if (!canRun) return + clearOptimize() + startOptimize({ + strategy_id: strategyId, + param_grid: buildGrid(), + objective, + // 未扫描参数固定为策略当前默认值; overrides 让 basic_filter/信号/风控按当前策略参与, + // 保证优化的就是用户实际回测的策略 (而非被剥离配置的裸策略)。 + params: selected?.params_defaults, + overrides: selected ? buildDefaultOverrides(selected) : undefined, + start, + end, + mode, + }) + } + + const result = task?.result + const progress = task?.progress + + return ( +
+ {/* ── 配置面板 ── */} +
+
+ + +
+ +
+ + +
+ +
+
+ + +
+
+ + +
+
+ +
+ + +
+ + {/* 可扫参数 */} + {params.length > 0 && ( +
+
扫描参数 (勾选后设范围)
+
+ {params.map(p => { + const s = sweeps[p.id] ?? defaultSweep(p) + const numeric = p.type === 'float' || p.type === 'int' + return ( +
+ + {s.enabled && numeric && ( +
+ updateSweep(p.id, { min: e.target.value })} placeholder="min" className={INPUT_CLS} /> + updateSweep(p.id, { max: e.target.value })} placeholder="max" className={INPUT_CLS} /> + updateSweep(p.id, { step: e.target.value })} placeholder="step" className={INPUT_CLS} /> +
+ )} + {s.enabled && !numeric && ( +
+ {p.type === 'bool' ? '扫描 [是 / 否]' : `扫描全部选项 (${p.options?.length ?? 0})`} +
+ )} +
+ ) + })} +
+
+ )} + + {/* 组合数 / 校验提示 */} + {strategyId && ( +
2000 || gridError) ? 'text-red-400' : 'text-secondary'}`}> + {gridError + ? gridError + : combos === 0 + ? '请至少勾选一个参数' + : `共 ${combos} 组参数组合${combos > 2000 ? ' — 超过上限 2000, 请增大 step 或缩小范围' : ''}`} +
+ )} + + {task?.isPending ? ( + + ) : ( + + )} +
+ + {/* ── 结果面板 ── */} +
+ {task?.error && ( +
{task.error}
+ )} + + {task?.isPending && progress && ( +
+
+ 进度 {progress.done}/{progress.total} + 当前最优: {progress.best_score != null ? progress.best_score.toFixed(3) : '—'} +
+
+
+
+
+ )} + + {!result && !task?.isPending && ( + + )} + + {result && ( +
+ {/* 最优参数 */} + {result.best_params && ( +
+
+ 最优参数 · {result.objective} = {result.best_score} +
+
+ {Object.entries(result.best_params).map(([k, v]) => ( + {k}: {String(v)} + ))} +
+
+ )} + +
+ {result.n_completed}/{result.n_combinations} 组完成 · 耗时 {(result.elapsed_ms / 1000).toFixed(1)}s +
+ + {/* 排名表 */} +
+ + + + + + + + + + + + + + + + {result.results.slice(0, 50).map(r => ( + + + + + + + + + + + + ))} + +
#参数{result.objective}夏普索提诺总收益最大回撤胜率交易数
{r.rank} + {r.error + ? 失败: {r.error.slice(0, 40)} + : {Object.entries(r.params).map(([k, v]) => `${k}=${v}`).join(', ')}} + {r.objective_raw != null ? r.objective_raw.toFixed(3) : '—'}{r.stats?.sharpe ?? '—'}{r.stats?.sortino ?? '—'}{r.stats?.total_return != null ? fmtPct(r.stats.total_return) : '—'}{r.stats?.max_drawdown != null ? fmtPct(r.stats.max_drawdown) : '—'}{r.stats?.win_rate != null ? fmtPct(r.stats.win_rate) : '—'}{r.stats?.n_trades ?? '—'}
+ {result.results.length > 50 && ( +
+ 仅显示前 50 组 · 共 {result.results.length} 组 +
+ )} +
+
+ )} +
+
+ ) +}