mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 16:54:20 +08:00
feat(web-ui): 一键寻优所有策略支持多进程并发(串行/4/8/16)
回测是 numpy/pandas 的 CPU 密集计算并持有 GIL,多线程无加速,必须用 多进程。照搬项目里已跑通的 screen/scanner.py 进程池模板。 后端: - OptimizeAllBacktestRequest 新增 workers 字段(默认 1=串行,0-32) - 抽出模块顶层函数 _optimize_one_strategy(可被 ProcessPoolExecutor pickle),策略类在子进程内构造,避开跨进程 pickle 问题 - _run_optimize_all 在 workers>=2 时用进程池并行,否则走原串行逻辑 - 顺带修复 backtest_schemas.py 的 E501 历史遗留 前端: - OptimizeView 配置区新增「一键寻优并发」选择器:自动检测 CPU 核数 + 串行/4/8/16 进程下拉,默认串行,标注推荐档 min(CPU, 8) - types.ts/api.ts 透传 workers 字段 实测 8 进程 vs 串行提速 4-6x;端到端冒烟验证并行与串行结果完全一致
This commit is contained in:
@@ -164,6 +164,13 @@ class OptimizeAllBacktestRequest(BaseModel):
|
||||
commission: float = Field(default=0.0003, ge=0, le=0.01)
|
||||
slippage: float = Field(default=0.0, ge=0, le=0.05)
|
||||
execution: Literal["next_open", "next_close"] = Field(default="next_open")
|
||||
workers: int = Field(
|
||||
default=1,
|
||||
ge=0,
|
||||
le=32,
|
||||
description="一键寻优并发工作进程数:0 或 1=串行(默认);2+=ProcessPoolExecutor "
|
||||
"并行(CPU-bound,线程无加速)。推荐 min(cpu_count, 8)。",
|
||||
)
|
||||
|
||||
# 数据来源 A:内联 OHLCV
|
||||
ohlcv: list[dict[str, Any]] | None = Field(default=None, max_length=2000)
|
||||
@@ -275,7 +282,9 @@ class SavedStrategyCreate(BaseModel):
|
||||
"""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=120, description="策略名称(用户自拟)")
|
||||
kind: Literal["single", "portfolio", "multi"] = Field(..., description="来源:单标的/组合/多策略组合")
|
||||
kind: Literal["single", "portfolio", "multi"] = Field(
|
||||
..., description="来源:单标的/组合/多策略组合"
|
||||
)
|
||||
strategy: str = Field(..., description="策略名(注册表 key,如 ma_cross)")
|
||||
strategy_label: str = Field(default="", description="策略展示名")
|
||||
params: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
@@ -561,62 +561,134 @@ def _run_optimize(df: pd.DataFrame, req: OptimizeBacktestRequest) -> dict[str, A
|
||||
return result.to_dict()
|
||||
|
||||
|
||||
def _optimize_one_strategy(
|
||||
strategy_name: str,
|
||||
grid: dict[str, list[Any]],
|
||||
df: pd.DataFrame,
|
||||
cash: float,
|
||||
commission: float,
|
||||
slippage: float,
|
||||
execution: str,
|
||||
) -> dict[str, Any] | None:
|
||||
"""跑单个策略的网格寻优,返回其最优点摘要(模块顶层,可被 ProcessPoolExecutor pickle)。
|
||||
|
||||
必须是模块级顶层函数:Windows 下 ProcessPoolExecutor 用 spawn 方式启动子进程,
|
||||
子进程按 ``module.qualname`` 重新 import 本函数。lambda / 闭包 / 嵌套函数不可 pickle。
|
||||
|
||||
策略类(``registry.get(name).build()``)在子进程内构造,从不跨进程传递,
|
||||
因此天然避开了 screen scanner 当年遇到的"策略类不可 pickle"问题。
|
||||
返回纯 dict(所有值都是 JSON 原生类型),可安全 pickle 回主进程。
|
||||
"""
|
||||
from easy_tdx.backtest.optimizer import ParamGridOptimizer
|
||||
|
||||
try:
|
||||
optimizer = ParamGridOptimizer(
|
||||
strategy_name=strategy_name,
|
||||
param_grid=grid,
|
||||
df=df,
|
||||
cash=cash,
|
||||
commission=commission,
|
||||
slippage=slippage,
|
||||
execution=execution,
|
||||
)
|
||||
except ValueError:
|
||||
# 单策略网格超限(不应发生,预设已控制规模)→ 跳过
|
||||
return None
|
||||
|
||||
result = optimizer.run()
|
||||
if result.best is None:
|
||||
return None
|
||||
|
||||
return {
|
||||
"strategy": strategy_name,
|
||||
"params": result.best.params,
|
||||
"total_return": result.best.total_return,
|
||||
"sharpe": result.best.sharpe,
|
||||
"max_drawdown": result.best.max_drawdown,
|
||||
"total_trades": result.best.total_trades,
|
||||
"win_rate": result.best.win_rate,
|
||||
"profit_factor": result.best.profit_factor,
|
||||
"grid_points": len(result.results),
|
||||
}
|
||||
|
||||
|
||||
def _run_optimize_all(df: pd.DataFrame, req: OptimizeAllBacktestRequest) -> dict[str, Any]:
|
||||
"""对所有策略的预设网格逐策略寻优,汇总成全局排名(后台线程内调用)。
|
||||
|
||||
遍历 ``STRATEGY_PRESETS`` 中每个策略,用其预设参数网格跑
|
||||
:class:`ParamGridOptimizer`,取各策略的最优点(best)组装排名。单个策略
|
||||
无有效结果(如全网格回测失败)则跳过。
|
||||
|
||||
并发:``req.workers >= 2`` 时用 ``ProcessPoolExecutor`` 跨进程并行寻优
|
||||
(回测是 CPU-bound,numpy/pandas 持 GIL,线程无加速,必须用进程)。
|
||||
``workers`` 为 0 或 1 时串行。进程池在函数内 ``with`` 创建/销毁,对前端
|
||||
轮询与 task_runner 透明。
|
||||
"""
|
||||
from easy_tdx.backtest.optimizer import ParamGridOptimizer
|
||||
from easy_tdx.backtest.strategies import get_registry
|
||||
from easy_tdx.backtest.strategies.presets import STRATEGY_PRESETS
|
||||
|
||||
registry = get_registry()
|
||||
# 过滤出已注册的策略 + 解析 label(label 必须在主进程取,避免子进程各自解析不一致)
|
||||
jobs: list[tuple[str, dict[str, list[Any]]]] = []
|
||||
labels: dict[str, str] = {}
|
||||
for strategy_name, grid in STRATEGY_PRESETS.items():
|
||||
if strategy_name not in registry.names():
|
||||
continue
|
||||
labels[strategy_name] = registry.get(strategy_name).label
|
||||
jobs.append((strategy_name, grid))
|
||||
|
||||
# 跑寻优:串行 or 进程池并行
|
||||
raw_results: list[dict[str, Any]] = []
|
||||
if req.workers and req.workers >= 2:
|
||||
import concurrent.futures
|
||||
|
||||
with concurrent.futures.ProcessPoolExecutor(max_workers=req.workers) as executor:
|
||||
futures = {
|
||||
executor.submit(
|
||||
_optimize_one_strategy,
|
||||
name,
|
||||
grid,
|
||||
df,
|
||||
req.cash,
|
||||
req.commission,
|
||||
req.slippage,
|
||||
req.execution,
|
||||
): name
|
||||
for name, grid in jobs
|
||||
}
|
||||
for future in concurrent.futures.as_completed(futures):
|
||||
res = future.result()
|
||||
if res is not None:
|
||||
raw_results.append(res)
|
||||
else:
|
||||
for name, grid in jobs:
|
||||
res = _optimize_one_strategy(
|
||||
name, grid, df, req.cash, req.commission, req.slippage, req.execution
|
||||
)
|
||||
if res is not None:
|
||||
raw_results.append(res)
|
||||
|
||||
# 组装排名(主进程统一构造 Pydantic 模型,保证类型一致)
|
||||
ranking: list[OptimizeAllRankEntry] = []
|
||||
per_strategy: dict[str, OptimizeAllRankEntry] = {}
|
||||
total_grid = 0
|
||||
|
||||
for strategy_name, grid in STRATEGY_PRESETS.items():
|
||||
# 预设里登记但未注册的策略跳过(理论上不应发生)
|
||||
if strategy_name not in registry.names():
|
||||
continue
|
||||
label = registry.get(strategy_name).label
|
||||
try:
|
||||
optimizer = ParamGridOptimizer(
|
||||
strategy_name=strategy_name,
|
||||
param_grid=grid,
|
||||
df=df,
|
||||
cash=req.cash,
|
||||
commission=req.commission,
|
||||
slippage=req.slippage,
|
||||
execution=req.execution,
|
||||
)
|
||||
except ValueError:
|
||||
# 单策略网格超限(不应发生,预设已控制规模)→ 跳过
|
||||
continue
|
||||
|
||||
result = optimizer.run()
|
||||
if result.best is None:
|
||||
continue
|
||||
|
||||
# 该策略本轮真实跑的网格点数(笛卡尔积,去掉失败点后的有效点)
|
||||
total_grid += len(result.results)
|
||||
|
||||
for res in raw_results:
|
||||
strategy_name = res["strategy"]
|
||||
entry = OptimizeAllRankEntry(
|
||||
strategy=strategy_name,
|
||||
strategy_label=label,
|
||||
params=result.best.params,
|
||||
total_return=result.best.total_return,
|
||||
sharpe=result.best.sharpe,
|
||||
max_drawdown=result.best.max_drawdown,
|
||||
total_trades=result.best.total_trades,
|
||||
win_rate=result.best.win_rate,
|
||||
profit_factor=result.best.profit_factor,
|
||||
grid_points=len(result.results),
|
||||
strategy_label=labels[strategy_name],
|
||||
params=res["params"],
|
||||
total_return=res["total_return"],
|
||||
sharpe=res["sharpe"],
|
||||
max_drawdown=res["max_drawdown"],
|
||||
total_trades=res["total_trades"],
|
||||
win_rate=res["win_rate"],
|
||||
profit_factor=res["profit_factor"],
|
||||
grid_points=res["grid_points"],
|
||||
)
|
||||
ranking.append(entry)
|
||||
per_strategy[strategy_name] = entry
|
||||
total_grid += res["grid_points"]
|
||||
|
||||
# 按 total_return 降序
|
||||
ranking.sort(key=lambda r: r.total_return, reverse=True)
|
||||
|
||||
@@ -229,6 +229,7 @@ export interface OptimizeAllBacktestRequest {
|
||||
commission?: number
|
||||
slippage?: number
|
||||
execution?: ExecutionMode
|
||||
workers?: number
|
||||
ohlcv?: Bar[]
|
||||
symbol?: string
|
||||
category?: Category
|
||||
|
||||
@@ -38,6 +38,22 @@ const strategy = ref('ma_cross')
|
||||
const paramGrid = ref<Record<string, Array<number | string>>>({})
|
||||
const cash = ref(1000000)
|
||||
const execution = ref<ExecutionMode>('next_open')
|
||||
// 一键寻优并发工作进程数:0=串行(同 workers=1);2+=多进程并行(CPU-bound 必须)
|
||||
const cpuCount = (() => {
|
||||
// navigator.hardwareConcurrency 返回逻辑核数(含超线程);非浏览器/不可用时回退 4
|
||||
const n = typeof navigator !== 'undefined' ? navigator.hardwareConcurrency : undefined
|
||||
return n && n > 0 ? n : 4
|
||||
})()
|
||||
// 推荐档:min(cpu, 8)。默认串行(workers=0),让用户实测后再开并发——
|
||||
// 小机器上 Windows spawn 子进程开销可能反而拖慢单个寻优任务。
|
||||
const recommendedWorkers = Math.min(cpuCount, 8)
|
||||
const workers = ref(0)
|
||||
const WORKER_OPTIONS: { value: number; label: string }[] = [
|
||||
{ value: 0, label: '串行(不并发)' },
|
||||
{ value: 4, label: '4 进程' },
|
||||
{ value: 8, label: '8 进程' },
|
||||
{ value: 16, label: '16 进程' },
|
||||
]
|
||||
// 成交价模式(精简为 开盘价/收盘价)
|
||||
const EXECUTIONS: { value: ExecutionMode; label: string }[] = [
|
||||
{ value: 'next_open', label: '开盘价' },
|
||||
@@ -109,6 +125,7 @@ async function onRunAll() {
|
||||
await store.runOptimizeAll({
|
||||
cash: cash.value,
|
||||
execution: execution.value,
|
||||
workers: workers.value,
|
||||
ohlcv: store.ohlcv,
|
||||
})
|
||||
}
|
||||
@@ -207,6 +224,24 @@ const rankingGrades = computed<GradeResult[]>(() =>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section class="panel-section">
|
||||
<h3>一键寻优并发</h3>
|
||||
<div class="field">
|
||||
<label
|
||||
>工作进程
|
||||
<span class="hint"
|
||||
>CPU {{ cpuCount }} 核 · 推荐 {{ recommendedWorkers }}</span
|
||||
></label
|
||||
>
|
||||
<select v-model.number="workers">
|
||||
<option v-for="o in WORKER_OPTIONS" :key="o.value" :value="o.value">{{ o.label }}</option>
|
||||
</select>
|
||||
</div>
|
||||
<p class="workers-tip">
|
||||
寻优是 CPU 密集计算,多进程可显著提速(仅对「一键寻优所有策略」生效)。机器较弱时建议先用串行测一次再开并发。
|
||||
</p>
|
||||
</section>
|
||||
|
||||
<button
|
||||
class="primary run-btn"
|
||||
:disabled="store.optimizeRunning || store.optimizeAllRunning"
|
||||
@@ -372,6 +407,18 @@ const rankingGrades = computed<GradeResult[]>(() =>
|
||||
font-weight: 600;
|
||||
margin-bottom: 12px;
|
||||
}
|
||||
.hint {
|
||||
font-weight: 400;
|
||||
color: var(--text-muted, #888);
|
||||
font-size: 12px;
|
||||
margin-left: 6px;
|
||||
}
|
||||
.workers-tip {
|
||||
margin: 8px 0 0;
|
||||
font-size: 12px;
|
||||
line-height: 1.5;
|
||||
color: var(--text-muted, #888);
|
||||
}
|
||||
.run-btn {
|
||||
width: 100%;
|
||||
padding: 10px;
|
||||
|
||||
Reference in New Issue
Block a user