mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 16:44:15 +08:00
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:
@@ -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": "任务不存在或已完成"}
|
||||
|
||||
|
||||
@@ -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"))
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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>
|
||||
)
|
||||
}
|
||||
Reference in New Issue
Block a user