mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 23:44:16 +08:00
* 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 通过。
296 lines
12 KiB
Python
296 lines
12 KiB
Python
"""参数网格搜索优化器。
|
|
|
|
给定策略 + 参数网格, 遍历所有参数组合各跑一次回测, 按目标指标排序, 返回最优参数。
|
|
|
|
- 参数网格校验对齐 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),
|
|
}
|