feat(optimizer): 参数网格搜索优化器 (并行回测 + 目标排序 + SSE + 前端面板) (#82)

* 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 通过。
This commit is contained in:
im47cn
2026-07-10 11:38:30 +08:00
committed by GitHub
parent cfc48ce424
commit ec3309163b
9 changed files with 1550 additions and 4 deletions
+208
View File
@@ -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": "任务不存在或已完成"}
+295
View File
@@ -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),
}
@@ -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)
@@ -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)
@@ -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"))
+274
View File
@@ -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<string, any>
objective_raw: number | null
stats?: Record<string, any>
rank: number
error?: string
}
export interface OptimizeResult {
objective: string
direction: string
n_combinations: number
n_completed: number
best_params: Record<string, any> | 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<string, any>
objective: string
direction?: string
max_workers?: number
params?: Record<string, any> | null // 未扫描参数固定为用户当前值
overrides?: Record<string, any> | 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, string | number | boolean | undefined | null>): 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)
}
+24
View File
@@ -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<string, any> {
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,
}
}
+17 -4
View File
@@ -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<Tab, { title: string; subtitle: string; hint: string }> = {
factor: {
@@ -17,6 +18,17 @@ const MODES: Record<Tab, { title: string; subtitle: string; hint: string }> = {
subtitle: '验证完整选股和交易规则',
hint: '看净值曲线、回撤、胜率和交易明细,适合判断策略是否可执行。',
},
optimizer: {
title: '参数优化',
subtitle: '网格搜索最优参数组合',
hint: '并行回测所有参数组合,按夏普/索提诺等目标排序,找到最优参数。',
},
}
const TAB_ICONS: Record<Tab, typeof BarChart3> = {
factor: BarChart3,
strategy: FlaskConical,
optimizer: SlidersHorizontal,
}
export function Backtest() {
@@ -24,8 +36,8 @@ export function Backtest() {
const modeSwitch = (
<div className="inline-flex rounded-btn border border-border bg-surface/80 p-0.5 shadow-sm">
{(['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 (
<button
@@ -57,6 +69,7 @@ export function Backtest() {
<main className="flex-1 min-h-0 px-3 pb-3 pt-3 lg:px-4 lg:pb-4">
{activeTab === 'factor' && <FactorBacktest />}
{activeTab === 'strategy' && <StrategyBacktest />}
{activeTab === 'optimizer' && <StrategyOptimizer />}
</main>
</div>
)
@@ -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<string>('')
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<Record<string, Sweep>>({})
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<string, Sweep> = {}
for (const p of s?.params ?? []) init[p.id] = defaultSweep(p)
setSweeps(init)
}
const updateSweep = (pid: string, patch: Partial<Sweep>) =>
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<string, any> => {
const grid: Record<string, any> = {}
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 (
<div className="grid grid-cols-1 gap-3 lg:grid-cols-[320px_1fr]">
{/* ── 配置面板 ── */}
<div className="space-y-3 rounded-card border border-border bg-surface p-4">
<div>
<label className="mb-1.5 block text-xs font-medium text-secondary"></label>
<select value={strategyId} onChange={e => onSelectStrategy(e.target.value)} className={INPUT_CLS}>
<option value=""></option>
{strategies.map(s => <option key={s.id} value={s.id}>{s.name}</option>)}
</select>
</div>
<div>
<label className="mb-1.5 block text-xs font-medium text-secondary"></label>
<select value={objective} onChange={e => setObjective(e.target.value)} className={INPUT_CLS}>
{OBJECTIVES.map(o => <option key={o.id} value={o.id}>{o.label}</option>)}
</select>
</div>
<div className="grid grid-cols-2 gap-2">
<div>
<label className="mb-1.5 block text-xs font-medium text-secondary"></label>
<DatePicker value={start} onChange={setStart} />
</div>
<div>
<label className="mb-1.5 block text-xs font-medium text-secondary"></label>
<DatePicker value={end} onChange={setEnd} />
</div>
</div>
<div>
<label className="mb-1.5 block text-xs font-medium text-secondary"></label>
<select value={mode} onChange={e => setMode(e.target.value as any)} className={INPUT_CLS}>
<option value="position"></option>
<option value="full"></option>
</select>
</div>
{/* 可扫参数 */}
{params.length > 0 && (
<div>
<div className="mb-1.5 text-xs font-medium text-secondary"> ()</div>
<div className="space-y-2">
{params.map(p => {
const s = sweeps[p.id] ?? defaultSweep(p)
const numeric = p.type === 'float' || p.type === 'int'
return (
<div key={p.id} className="rounded-input border border-border/60 p-2">
<label className="flex items-center gap-2 text-xs">
<input type="checkbox" checked={s.enabled} onChange={e => updateSweep(p.id, { enabled: e.target.checked })} />
<span className="font-medium text-foreground">{p.label}</span>
<span className="text-secondary">({p.type})</span>
</label>
{s.enabled && numeric && (
<div className="mt-2 grid grid-cols-3 gap-1.5">
<input type="number" value={s.min} onChange={e => updateSweep(p.id, { min: e.target.value })} placeholder="min" className={INPUT_CLS} />
<input type="number" value={s.max} onChange={e => updateSweep(p.id, { max: e.target.value })} placeholder="max" className={INPUT_CLS} />
<input type="number" value={s.step} onChange={e => updateSweep(p.id, { step: e.target.value })} placeholder="step" className={INPUT_CLS} />
</div>
)}
{s.enabled && !numeric && (
<div className="mt-1 text-[11px] text-secondary">
{p.type === 'bool' ? '扫描 [是 / 否]' : `扫描全部选项 (${p.options?.length ?? 0})`}
</div>
)}
</div>
)
})}
</div>
</div>
)}
{/* 组合数 / 校验提示 */}
{strategyId && (
<div className={`text-xs ${(combos > 2000 || gridError) ? 'text-red-400' : 'text-secondary'}`}>
{gridError
? gridError
: combos === 0
? '请至少勾选一个参数'
: `${combos} 组参数组合${combos > 2000 ? ' — 超过上限 2000, 请增大 step 或缩小范围' : ''}`}
</div>
)}
{task?.isPending ? (
<button onClick={stopOptimize} className="inline-flex w-full items-center justify-center gap-1.5 rounded-btn bg-red-500/90 px-3 py-2 text-xs font-medium text-white hover:bg-red-500">
<Square className="h-3.5 w-3.5" />
</button>
) : (
<button onClick={onRun} disabled={!canRun} className="inline-flex w-full items-center justify-center gap-1.5 rounded-btn bg-accent px-3 py-2 text-xs font-medium text-white hover:opacity-90 disabled:opacity-40 disabled:cursor-not-allowed">
<Play className="h-3.5 w-3.5" />
</button>
)}
</div>
{/* ── 结果面板 ── */}
<div className="min-h-[300px] rounded-card border border-border bg-surface p-4">
{task?.error && (
<div className="mb-3 rounded-input border border-red-500/30 bg-red-500/10 px-3 py-2 text-xs text-red-400">{task.error}</div>
)}
{task?.isPending && progress && (
<div className="mb-4">
<div className="mb-1 flex justify-between text-xs text-secondary">
<span> {progress.done}/{progress.total}</span>
<span>: {progress.best_score != null ? progress.best_score.toFixed(3) : '—'}</span>
</div>
<div className="h-1.5 overflow-hidden rounded-full bg-elevated">
<div className="h-full bg-accent transition-all" style={{ width: `${progress.total ? (progress.done / progress.total) * 100 : 0}%` }} />
</div>
</div>
)}
{!result && !task?.isPending && (
<EmptyState title="参数优化" hint="选择策略、勾选要扫描的参数与优化目标,网格搜索会并行回测所有组合并按目标排序。" />
)}
{result && (
<div className="space-y-4">
{/* 最优参数 */}
{result.best_params && (
<div className="rounded-card border border-accent/30 bg-accent/5 p-3">
<div className="mb-1.5 flex items-center gap-1.5 text-xs font-semibold text-accent">
<Trophy className="h-3.5 w-3.5" /> · {result.objective} = {result.best_score}
</div>
<div className="flex flex-wrap gap-1.5">
{Object.entries(result.best_params).map(([k, v]) => (
<span key={k} className="rounded-full border border-border bg-surface px-2 py-0.5 text-[11px]">{k}: {String(v)}</span>
))}
</div>
</div>
)}
<div className="text-xs text-secondary">
{result.n_completed}/{result.n_combinations} · {(result.elapsed_ms / 1000).toFixed(1)}s
</div>
{/* 排名表 */}
<div className="overflow-x-auto">
<table className="w-full text-xs">
<thead>
<tr className="border-b border-border text-secondary">
<th className="px-2 py-1.5 text-left">#</th>
<th className="px-2 py-1.5 text-left"></th>
<th className="px-2 py-1.5 text-right">{result.objective}</th>
<th className="px-2 py-1.5 text-right"></th>
<th className="px-2 py-1.5 text-right"></th>
<th className="px-2 py-1.5 text-right"></th>
<th className="px-2 py-1.5 text-right"></th>
<th className="px-2 py-1.5 text-right"></th>
<th className="px-2 py-1.5 text-right"></th>
</tr>
</thead>
<tbody>
{result.results.slice(0, 50).map(r => (
<tr key={r.rank} className="border-b border-border/40 hover:bg-elevated/50">
<td className="px-2 py-1.5 text-secondary">{r.rank}</td>
<td className="px-2 py-1.5">
{r.error
? <span className="text-red-400">: {r.error.slice(0, 40)}</span>
: <span className="text-foreground">{Object.entries(r.params).map(([k, v]) => `${k}=${v}`).join(', ')}</span>}
</td>
<td className="px-2 py-1.5 text-right font-medium">{r.objective_raw != null ? r.objective_raw.toFixed(3) : '—'}</td>
<td className="px-2 py-1.5 text-right">{r.stats?.sharpe ?? '—'}</td>
<td className="px-2 py-1.5 text-right">{r.stats?.sortino ?? '—'}</td>
<td className="px-2 py-1.5 text-right">{r.stats?.total_return != null ? fmtPct(r.stats.total_return) : '—'}</td>
<td className="px-2 py-1.5 text-right">{r.stats?.max_drawdown != null ? fmtPct(r.stats.max_drawdown) : '—'}</td>
<td className="px-2 py-1.5 text-right">{r.stats?.win_rate != null ? fmtPct(r.stats.win_rate) : '—'}</td>
<td className="px-2 py-1.5 text-right">{r.stats?.n_trades ?? '—'}</td>
</tr>
))}
</tbody>
</table>
{result.results.length > 50 && (
<div className="mt-2 text-center text-[11px] text-secondary">
50 · {result.results.length}
</div>
)}
</div>
</div>
)}
</div>
</div>
)
}