Files
GitHub 36ea8497ae perf(engine): 信号管线提速 ×12.6 — 向量化快速路径 + 日期查找表 + 去 strftime
基准先行(scripts/bench_engine.py):ma_cross 800 根全流程基线 93.6ms。
profile 打破预期:85% 墙钟在 OrderSimulator._find_bar_index(每信号全列
strftime),而非计划认为的逐 bar 信号循环。三处优化(行为逐位一致):

- Strategy 新增 entry_exit_masks() 显式钩子,19 个内置策略全实现;引擎按
  掩码 + 候选事件 bar 状态机一次产出信号,持仓估算逐行复刻
  _update_strategy_position(含买不足 1 手退化路径);约束检测
  _vectorize_eligibility 显式可测,signal_path=auto/vector/loop 可指定;
  不满足约束(无钩子/缠论注入/掩码形状错)自动回退逐 bar
- OrderSimulator._build_dt_lookup:O(信号×bar) 全列扫描 → 每 simulate 一次
  O(bar) 查找表(重复日期取首个、未命中 None、object 恒不匹配语义对齐)
- _datetime_to_int 去 strftime:year*10000+month*100+day 整数算术
  (NaT→NaN 行为一致),_bind_data 与查找表共用

实测:ma_cross 全流程 93.6→7.4ms(×12.6),信号层 ×1.39~1.72(四策略);
32 点网格寻优 ~3.0s→231ms。对拍 39 例:19 策略×参数变体×warmup/低资金/
费率/缓存,performance/trades/equity/positions 逐位一致。
顺带记录:wr_reversal 默认阈值与 MyTT WR 刻度不匹配(策略恒不交易,既有问题未改)
2026-09-01 23:47:06 +08:00

130 lines
4.9 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""bench_engine — 回测引擎信号生成路径基准(v1.28)。
测「单标的 N 根 × 指定内置策略 全流程」(engine.run 端到端)在两条信号
路径下的耗时,并附信号生成阶段(_generate_signals)的归因分解——
v1.25 实测指标缓存墙钟仅 ~1.01x,瓶颈在逐 bar Python 循环,本脚本即为其
优化前后的量化依据(数字进 CHANGELOG,报告用 --json 机器可读)。
用法::
.venv/Scripts/python.exe scripts/bench_engine.py # 800 根 ma_cross
.venv/Scripts/python.exe scripts/bench_engine.py --bars 300 --runs 30 --strategy macd
.venv/Scripts/python.exe scripts/bench_engine.py --json # CI/脚本消费
"""
from __future__ import annotations
import argparse
import json
import statistics
import time
import numpy as np
import pandas as pd
from easy_tdx.backtest.engine import BacktestEngine
from easy_tdx.backtest.strategies import get_registry
def synthetic_ohlcv(n: int, seed: int = 42, base: float = 20.0) -> pd.DataFrame:
"""确定性随机游走 OHLCV(与对拍单测同款生成方式)。"""
rng = np.random.default_rng(seed)
rets = rng.normal(0.0005, 0.02, n)
close = base * np.cumprod(1.0 + rets)
open_ = np.concatenate([[base], close[:-1]])
high = np.maximum(open_, close) * (1.0 + np.abs(rng.normal(0.0, 0.008, n)))
low = np.minimum(open_, close) * (1.0 - np.abs(rng.normal(0.0, 0.008, n)))
vol = rng.integers(50_000, 5_000_000, n).astype(float)
dates = pd.bdate_range("2022-01-04", periods=n)
return pd.DataFrame(
{
"datetime": dates,
"open": open_,
"high": high,
"low": low,
"close": close,
"vol": vol,
"amount": vol * close,
}
)
def _time_path(engine: BacktestEngine, df: pd.DataFrame, runs: int) -> dict[str, float]:
"""计时一条路径:全流程 run + 仅信号生成阶段(归因用)。"""
engine.run(df) # 预热(numpy/pandas 内部缓存、导入惰性初始化)
full_times: list[float] = []
for _ in range(runs):
t0 = time.perf_counter()
engine.run(df)
full_times.append(time.perf_counter() - t0)
# 信号生成单独计时(引擎其余阶段:OrderSimulator/Portfolio/Performance
sig_times: list[float] = []
for _ in range(runs):
t0 = time.perf_counter()
engine._generate_signals(df, None)
sig_times.append(time.perf_counter() - t0)
return {
"full_median_ms": statistics.median(full_times) * 1000,
"full_min_ms": min(full_times) * 1000,
"signal_median_ms": statistics.median(sig_times) * 1000,
"signal_min_ms": min(sig_times) * 1000,
}
def main() -> int:
parser = argparse.ArgumentParser(description="回测引擎信号路径基准")
parser.add_argument("--bars", type=int, default=800, help="K 线根数(默认 800")
parser.add_argument("--runs", type=int, default=20, help="每条路径计时次数(取中位数)")
parser.add_argument("--strategy", default="ma_cross", help="内置策略名(默认 ma_cross")
parser.add_argument("--seed", type=int, default=42, help="合成行情种子")
parser.add_argument("--json", action="store_true", help="JSON 输出(机器可读)")
args = parser.parse_args()
df = synthetic_ohlcv(args.bars, seed=args.seed)
strat = get_registry().get(args.strategy).build()
loop = BacktestEngine(strat, signal_path="loop")
vec = BacktestEngine(get_registry().get(args.strategy).build(), signal_path="vector")
loop_stats = _time_path(loop, df, args.runs)
vec_stats = _time_path(vec, df, args.runs)
result = {
"strategy": args.strategy,
"bars": args.bars,
"runs": args.runs,
"trades": len(loop.run(df).trades),
"loop": loop_stats,
"vector": vec_stats,
"full_speedup": loop_stats["full_median_ms"] / vec_stats["full_median_ms"],
"signal_speedup": loop_stats["signal_median_ms"] / vec_stats["signal_median_ms"],
}
if args.json:
print(json.dumps(result, ensure_ascii=False, indent=2))
else:
print(
f"策略={result['strategy']} bars={result['bars']} runs={result['runs']}"
f" trades={result['trades']}"
)
print(
f"逐 bar 路径:全流程 {loop_stats['full_median_ms']:.1f}ms"
f"(最快 {loop_stats['full_min_ms']:.1f})|信号生成"
f" {loop_stats['signal_median_ms']:.1f}ms"
)
print(
f"向量化路径:全流程 {vec_stats['full_median_ms']:.1f}ms"
f"(最快 {vec_stats['full_min_ms']:.1f})|信号生成"
f" {vec_stats['signal_median_ms']:.1f}ms"
)
print(f"加速比:全流程 ×{result['full_speedup']:.2f}", end="")
print(f"|信号生成 ×{result['signal_speedup']:.2f}")
return 0
if __name__ == "__main__":
raise SystemExit(main())