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 刻度不匹配(策略恒不交易,既有问题未改)
This commit is contained in:
GitHub
2026-09-01 23:47:06 +08:00
parent 1aacce9b7d
commit 36ea8497ae
7 changed files with 751 additions and 24 deletions
+9
View File
@@ -4,6 +4,15 @@
## [未发布]
### 性能
- **回测引擎信号管线提速 ×12.6ma_cross 800 根全流程 93.6ms → 7.4msWindows/Py3.12 实测)**——升级计划 P4 遗留项。先写基准(`scripts/bench_engine.py`perf_counter 中位数)再优化,profile 归因打破预期:v1.25 以为瓶颈是 `_generate_signals` 的逐 bar 循环,实测 **85% 墙钟在 `OrderSimulator._find_bar_index`**(每个信号都把整列 datetime `strftime("%Y%m%d")` 一遍,46 信号×800 根 ≈ 0.78s),另 ~10% 在 `_bind_data``_datetime_to_int`(同样 strftime 全列)。三处优化(全部保持行为逐位一致):
1. **向量化信号生成快速路径**`Strategy` 新增 `entry_exit_masks()` 钩子(显式声明 entry/exit 布尔掩码,语义约定与约束写在 docstring),19 个内置策略全部实现;引擎按掩码 + 候选事件 bar 状态机一次产出信号(逐 bar 循环退化为逐事件循环),持仓估算逐行复刻 `_update_strategy_position`(含买不足 1 手的退化路径)。约束检测 `_vectorize_eligibility` 显式可测(未实现钩子 / 缠论注入 → 回退逐 bar);`signal_path="auto|vector|loop"` 可强制指定。信号层加速 ×1.39~1.72(800 根,四策略实测);
2. **OrderSimulator 日期查找表**`_build_dt_lookup()` 每个 `simulate()` 只转一次 datetime→行号(O(信号数×bar 数) → O(bar 数)),重复日期取首个、未命中 None、object 列恒不匹配等语义与原全列扫描逐条对齐;
3. **`_datetime_to_int` 去 strftime**datetime64 走 `year*10000+month*100+day` 整数算术(输出、NaT→NaN 行为与 strftime 一致),`_bind_data` 与查找表共用。
效果:单标的 800 根 ma_cross 全流程 **×12.6**macd ×11.3、boll_breakout ×11.2、rsi_reversal ×11.1);32 点网格寻优 3.0s(按基线折算)→ 231ms。对拍单测 `test_backtest_engine_vector.py`(39 例):19 策略 × 默认/非默认参数 × warmup/极低资金/非默认费率/指标缓存,performance/trades/equity_curve/positions 逐位一致。
- 顺带发现(未改行为,仅记录):`wr_reversal` 的默认阈值 -80/-20 是通达信 -100~0 惯例,而 MyTT 的 WR 为 0~100 刻度——默认参数(及边界内任意合法参数)下 entry 恒 False,策略实际不产生交易;对拍用 `skip_bounds` 参数覆盖其掩码路径,语义修正另行排期。
### 新增
- **Playwright E2E 前端测试基建**(升级计划 P4-1)——web-ui 引入 `@playwright/test``e2e/` + `playwright.config.ts``npm run test:e2e`)。**mock 方案选后端合成数据而非 page.route 拦截**`EASY_TDX_E2E_MOCK=1` 时 serve 的 lifespan 把 TDX/MAC 客户端替换为合成数据客户端(`web/e2e_mock.py`,按 (market, code) CRC32 播种的确定性随机游走,分页语义与真实 /bars 一致),回测/WF/一条龙评估/自选/策略库继续走**真实后端代码**(它们本就不依赖行情连接),SSE 由 QuoteStreamer 真轮询合成数据全链路覆盖(mock 模式下轮询降到 2s 一拍,不受交易时段限制)。用例覆盖:看板五大指数区块+SSE 价格渲染、自选增删、回测全流程(净值图/绩效表/成交记录)、「附加分析」开关(WF 逐窗柱状图+一条龙评估卡)、策略库保存;`EASY_TDX_CONFIG_DIR` 指向每轮独立临时目录(断言可写死、不污染真实 `~/.easy_tdx`)。CI frontend job 追加 E2E 步骤;`verify_ci.sh``--no-frontend` 与前端 typecheck+build+E2E 段。新增 `tests/unit/test_e2e_mock.py`(11 例)守护 mock 与真实客户端的契约。
+129
View File
@@ -0,0 +1,129 @@
"""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())
+119 -1
View File
@@ -1,6 +1,17 @@
"""BacktestEngine — orchestrate vectorized execution pipeline.
Coordinates Strategy → OrderSimulator → PortfolioTracker → PerformanceAnalyzer.
信号生成双路径(v1.28):
- **向量化快速路径**:策略实现了
:meth:`~easy_tdx.backtest.strategy.Strategy.entry_exit_masks` 且无缠论注入时,
预计算指标列后用 numpy 掩码 + 事件 bar 状态机一次性产出信号(逐 bar 的
Python 循环是网格寻优的实测瓶颈,见 CHANGELOG v1.25 的诚实数据);
- **逐 bar 路径**:不满足约束时自动回退,行为与历史版本完全一致。
两条路径在同一 df + 同参数下 performance 逐位一致(对拍单测守护,
``tests/unit/test_backtest_engine_vector.py``)。
"""
from __future__ import annotations
@@ -8,6 +19,7 @@ from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
import numpy as np
import pandas as pd
from easy_tdx.backtest.orders import OrderSimulator
@@ -67,6 +79,7 @@ class BacktestEngine:
symbol: str | None = None,
auto_fees: bool = False,
indicator_cache: Any | None = None,
signal_path: str = "auto",
):
"""Initialize engine.
@@ -97,9 +110,16 @@ class BacktestEngine:
indicator_cache: 指标计算缓存(网格寻优跨点复用,见
:class:`~easy_tdx.backtest.indicator_cache.IndicatorCache`)。
None = 每次直接计算(默认,向后兼容)。
signal_path: 信号生成路径。``"auto"``(默认)= 策略满足向量化约束
时走快速路径、否则回退逐 bar;``"vector"`` = 强制向量化
(不满足约束时抛 ValueError,测试/基准用);``"loop"`` = 强制
逐 bar(对拍基准用)。两条路径结果逐位一致(对拍单测守护)。
.. versionchanged:: 1.24
新增 ``symbol`` / ``auto_fees`` 品种感知费率参数。
.. versionchanged:: 1.28
新增 ``signal_path`` 与向量化快速路径(约束见
:meth:`~easy_tdx.backtest.strategy.Strategy.entry_exit_masks`)。
"""
self._strategy_cls = strategy if isinstance(strategy, type) else type(strategy)
self._strategy_instance = strategy if isinstance(strategy, Strategy) else None
@@ -131,6 +151,9 @@ class BacktestEngine:
self._execution_model = execution_model
self._warmup_bars = max(int(warmup_bars), 0)
self._indicator_cache = indicator_cache
if signal_path not in ("auto", "vector", "loop"):
raise ValueError(f"signal_path 取值应为 auto/vector/loop,得到 '{signal_path}'")
self._signal_path = signal_path
def run(self, df: pd.DataFrame, chanlun_result: Any | None = None) -> BacktestResult:
"""Run backtest.
@@ -310,7 +333,19 @@ class BacktestEngine:
# Call init
strat._call_init()
# Generate signals bar by bar
# ── 向量化快速路径(v1.28):满足约束时跳过逐 bar 循环 ──────────────
if self._signal_path in ("auto", "vector"):
eligible, reason = self._vectorize_eligibility(strat, chanlun_result)
if eligible:
signals = self._generate_signals_vector(df, strat)
if signals is not None:
return signals
reason = "entry_exit_masks 形状与 K 线长度不符"
if self._signal_path == "vector":
raise ValueError(f"策略不满足向量化约束:{reason}")
# auto:回退逐 bar 路径(行为与历史版本一致)
# ── 逐 bar 路径(历史行为,逐字未动) ───────────────────────────────
all_signals: list[Signal] = []
active_stops: list[_StopCondition] = []
@@ -358,6 +393,89 @@ class BacktestEngine:
return all_signals
# ── 向量化快速路径(v1.28) ────────────────────────────────────────────────
def _vectorize_eligibility(
self, strat: Strategy, chanlun_result: Any | None
) -> tuple[bool, str]:
"""显式约束检测:策略 + 引擎配置是否可走向量化路径。
约束(可测试,见 test_backtest_engine_vector.py):
1. 策略覆写了 ``entry_exit_masks``(显式声明语义等价的掩码);
2. 无缠论注入(``chanlun_result`` / ``chanlun_level`` 都未设置——
缠论策略的 next() 依赖缠论结果,掩码无法表达)。
Returns:
(是否可向量化, 不可用原因描述)。
"""
if type(strat).entry_exit_masks is Strategy.entry_exit_masks:
return False, "策略未实现 entry_exit_masks"
if chanlun_result is not None or self._chanlun_level is not None:
return False, "缠论注入(chanlun_result/chanlun_level)不支持向量化"
return True, ""
def _generate_signals_vector(self, df: pd.DataFrame, strat: Strategy) -> list[Signal] | None:
"""用 entry/exit 掩码 + 事件 bar 状态机一次性产出信号。
状态机语义与逐 bar 路径严格对齐(对拍单测守护逐位一致):
- 空仓 + ``entry[i]`` → 全仓 BUY(与 ``buy()`` 同参:size=0、无限价、
无止损止盈);
- 持仓 + ``exit[i]`` → 全仓 SELL
- 持仓估算复刻 :meth:`_update_strategy_position`:按收盘价全仓、
100 股整手、资金不足 1 手时仓位状态不变(后续 entry 仍可触发);
- warmup 期不产生信号。
Returns:
信号列表;掩码形状与 K 线不符(策略实现错误)返回 None 供上层回退。
"""
masks = strat.entry_exit_masks()
if masks is None:
return None
entry_raw, exit_raw = masks
entry = np.asarray(entry_raw, dtype=bool)
exit_ = np.asarray(exit_raw, dtype=bool)
n = len(df)
if entry.ndim != 1 or exit_.ndim != 1 or len(entry) != n or len(exit_) != n:
return None
dt_arr = strat._datetime_array
close_arr = df["close"].to_numpy(dtype=np.float64)
if dt_arr is None:
return None
# 候选事件 bar(任一掩码触发),远少于总 bar 数——逐 bar 循环退化为
# 逐事件循环,这是向量化收益的来源
candidates = np.flatnonzero(entry | exit_)
candidates = candidates[candidates >= self._warmup_bars]
signals: list[Signal] = []
cash = self._cash
position = 0.0
for i in candidates:
i_int = int(i)
if position == 0.0 and entry[i_int]:
signals.append(
Signal(datetime=int(dt_arr[i_int]), direction="BUY", size=0, price=None)
)
price = close_arr[i_int]
shares = int(cash / (price * (1 + self._commission)) / 100) * 100
if shares > 0:
position += shares
cash -= shares * price
elif position > 0.0 and exit_[i_int]:
signals.append(
Signal(datetime=int(dt_arr[i_int]), direction="SELL", size=0, price=None)
)
cash += position * close_arr[i_int]
position = 0.0
# 保持策略内部状态与逐 bar 路径一致(后续读取 strat._cash 的代码同构)
strat._cash = cash
strat._position_size = position
return signals
def _check_stop_conditions(
self,
active_stops: list[_StopCondition],
+45 -19
View File
@@ -45,6 +45,12 @@ class OrderSimulator:
slippage: float = 0.0
slippage_model: SlippageModel | None = None
def __post_init__(self) -> None:
# datetime → 行号查找表(惰性构建一次)。v1.28 之前每个信号都对整列
# datetime 做 strftime 扫描(O(信号数×bar 数),占全流程 ~85% 墙钟),
# 网格寻优每点都跑一遍 simulate,是实测最大瓶颈。
self._dt_lookup: dict[int, int] | None = None
def simulate(
self,
signals: list[Signal],
@@ -165,30 +171,50 @@ class OrderSimulator:
Returns:
K 线索引,未找到返回 None
性能(v1.28):查找表在首次调用时构建一次(O(bar 数)),此后每个
信号 O(1) 查询;语义与逐信号全列扫描完全一致——重复日期取首个匹配
(等价于原来的 ``mask.argmax()``),未命中返回 None。
"""
# 检查 df 中的 datetime 列类型
if self._dt_lookup is None:
self._dt_lookup = self._build_dt_lookup()
return self._dt_lookup.get(int(datetime_val))
def _build_dt_lookup(self) -> dict[int, int]:
"""构建 datetime → 行号查找表(重复日期保留首个出现)。
与历史行为的对应关系:
- datetime64 列:向量化转 YYYYMMDD int 后建表(原来每个信号都
``strftime`` 全列扫描一遍,是 O(信号数×bar 数) 的热点);
- 数值列(int/float):整数值等价于原 ``dt_col == datetime_val`` 的
直接比较(非整数 float 不会命中,与原来一致);
- 其他列(object 等):原来直接比较恒不命中、返回 None——这里同样
产出空表(保持行为不变,包括 object-Timestamp 列不匹配的既有行为)。
"""
from easy_tdx.backtest.strategy import _datetime_to_int
dt_col = self.df["datetime"]
lookup: dict[int, int] = {}
# 尝试直接比较(如果是 int 类型)
# 注意:用 to_numpy().argmax() 取位置索引,而非 idxmax()(返回 label),
# 因为后续 self.df.iloc[...] 按位置取行;若 df.index 非默认 RangeIndex
# label != position 会导致撮合取错 bar。
try:
mask = (dt_col == datetime_val).to_numpy()
if mask.any():
return int(mask.argmax())
except (TypeError, ValueError):
pass
# 如果是 datetime 对象,转为 int 比较
if pd.api.types.is_datetime64_any_dtype(dt_col):
dt_ints = dt_col.dt.strftime("%Y%m%d").astype(int)
mask_arr = (dt_ints == datetime_val).to_numpy()
if mask_arr.any():
return int(mask_arr.argmax())
return None
ints = _datetime_to_int(dt_col.to_numpy())
for i, v in enumerate(ints):
if v == v: # NaT → NaN,跳过(原 strftime 同样不产出该行)
lookup.setdefault(int(v), i)
return lookup
return None
arr = dt_col.to_numpy()
if arr.dtype.kind in "iuf":
for i, v in enumerate(arr):
# 浮点列只有整数值(20240104.0)才可能与 int 信号相等,
# 与原 == 比较语义一致
if float(v).is_integer():
lookup.setdefault(int(v), i)
return lookup
# object / 字符串等:原实现 (dt_col == int) 恒为 False → 恒 None
return lookup
def _resolve_exec_index(self, bar_idx: int) -> int | None:
"""根据执行模式确定成交的 K 线索引。
@@ -9,6 +9,10 @@
from __future__ import annotations
from typing import Any
import numpy as np
from easy_tdx.backtest.strategies.registry import (
Param,
ParametrizedStrategy,
@@ -69,6 +73,10 @@ class MaCrossStrategy(ParametrizedStrategy):
elif self.dead[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""与 next() 同源(gold/dead 即 next() 判定用的同一组掩码数组)。"""
return self.gold, self.dead
# ── MACD 金叉 ──────────────────────────────────────────────────────────────────
@@ -106,6 +114,10 @@ class MacdStrategy(ParametrizedStrategy):
elif self.dead[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""与 next() 同源(gold/dead 即 next() 判定用的同一组掩码数组)。"""
return self.gold, self.dead
# ── 布林带突破 ─────────────────────────────────────────────────────────────────
@@ -135,6 +147,11 @@ class BollBreakoutStrategy(ParametrizedStrategy):
elif close >= self.upper[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""触及下轨进 / 触及上轨出(NaN 轨道期比较为 False,与 next() 一致)。"""
close = self.data.close.raw
return close <= self.lower, close >= self.upper
# ── RSI 超买超卖 ───────────────────────────────────────────────────────────────
@@ -165,6 +182,11 @@ class RsiReversalStrategy(ParametrizedStrategy):
elif rsi >= self.p["overbought"] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""RSI 超卖进 / 超买出(NaN 预热期比较为 False,与 next() 一致)。"""
rsi = np.asarray(self.rsi, dtype=np.float64)
return rsi <= self.p["oversold"], rsi >= self.p["overbought"]
# ── KDJ 金叉 ───────────────────────────────────────────────────────────────────
@@ -199,6 +221,10 @@ class KdjCrossStrategy(ParametrizedStrategy):
elif self.dead[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""与 next() 同源(gold/dead 即 next() 判定用的同一组掩码数组)。"""
return self.gold, self.dead
# ── EMA 双线交叉 ──────────────────────────────────────────────────────────────
@@ -228,6 +254,10 @@ class EmaCrossStrategy(ParametrizedStrategy):
elif self.dead[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""与 next() 同源(gold/dead 即 next() 判定用的同一组掩码数组)。"""
return self.gold, self.dead
# ── 三均线系统 ────────────────────────────────────────────────────────────────
@@ -257,6 +287,13 @@ class TripleMaStrategy(ParametrizedStrategy):
elif self.ma_s[i] < self.ma_m[i] < self.ma_l[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""均线多头排列进 / 空头排列出(链式比较逐元素展开,与 next() 一致)。"""
return (
(self.ma_s > self.ma_m) & (self.ma_m > self.ma_l),
(self.ma_s < self.ma_m) & (self.ma_m < self.ma_l),
)
# ── 唐安奇通道(海龟)────────────────────────────────────────────────────────
@@ -282,6 +319,11 @@ class DonchianStrategy(ParametrizedStrategy):
elif close <= self.lower[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""突破 N 日高进 / 跌破 N 日低出(NaN 预热期比较为 False,与 next() 一致)。"""
close = self.data.close.raw
return close >= self.upper, close <= self.lower
# ── 肯特纳通道 ────────────────────────────────────────────────────────────────
@@ -310,6 +352,11 @@ class KeltnerStrategy(ParametrizedStrategy):
elif close <= self.lower[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""突破上轨进 / 跌破下轨出(NaN 预热期比较为 False,与 next() 一致)。"""
close = self.data.close.raw
return close >= self.upper, close <= self.lower
# ── BBI 多空指标 ──────────────────────────────────────────────────────────────
@@ -340,6 +387,11 @@ class BbiStrategy(ParametrizedStrategy):
elif close < self.bbi[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""上穿 BBI 进 / 下穿 BBI 出(NaN 预热期比较为 False,与 next() 一致)。"""
close = self.data.close.raw
return close > self.bbi, close < self.bbi
# ── CCI 顺势指标 ──────────────────────────────────────────────────────────────
@@ -368,6 +420,11 @@ class CciStrategy(ParametrizedStrategy):
elif cci >= self.p["overbought"] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""CCI 跌破超卖线进 / 涨破超买线出(与 next() 同一比较)。"""
cci = np.asarray(self.cci, dtype=np.float64)
return cci <= self.p["oversold"], cci >= self.p["overbought"]
# ── WR 威廉指标 ───────────────────────────────────────────────────────────────
@@ -396,6 +453,11 @@ class WrReversalStrategy(ParametrizedStrategy):
elif wr >= self.p["overbought"] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""WR 进入超卖区进 / 超买区出(与 next() 同一比较)。"""
wr = np.asarray(self.wr, dtype=np.float64)
return wr <= self.p["oversold"], wr >= self.p["overbought"]
# ── BIAS 乖离率 ───────────────────────────────────────────────────────────────
@@ -423,6 +485,12 @@ class BiasReversalStrategy(ParametrizedStrategy):
elif bias_pct >= threshold and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""负乖离超跌进 / 正乖离超涨出(与 next() 同一比较,阈值放大 100 倍口径)。"""
bias_pct = np.asarray(self.bias, dtype=np.float64) * 100
threshold = self.p["threshold"]
return bias_pct <= -threshold, bias_pct >= threshold
# ── DMI 趋向指标 ──────────────────────────────────────────────────────────────
@@ -452,6 +520,10 @@ class DmiStrategy(ParametrizedStrategy):
elif self.dead[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""与 next() 同源(gold/dead 即 next() 判定用的同一组掩码数组)。"""
return self.gold, self.dead
# ── TRIX 三重平滑 ─────────────────────────────────────────────────────────────
@@ -479,6 +551,10 @@ class TrixStrategy(ParametrizedStrategy):
elif self.dead[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""与 next() 同源(gold/dead 即 next() 判定用的同一组掩码数组)。"""
return self.gold, self.dead
# ── EMV 简易波动 ──────────────────────────────────────────────────────────────
@@ -505,6 +581,11 @@ class EmvStrategy(ParametrizedStrategy):
elif self.emv[i] < 0 and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""EMV 上穿 0 轴进 / 下穿 0 轴出(与 next() 同一比较)。"""
emv = np.asarray(self.emv, dtype=np.float64)
return emv > 0, emv < 0
# ── DPO 区间震荡 ──────────────────────────────────────────────────────────────
@@ -531,6 +612,10 @@ class DpoStrategy(ParametrizedStrategy):
elif self.dead[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""与 next() 同源(gold/dead 即 next() 判定用的同一组掩码数组)。"""
return self.gold, self.dead
# ── ATR 通道突破 ──────────────────────────────────────────────────────────────
@@ -561,6 +646,13 @@ class AtrBreakoutStrategy(ParametrizedStrategy):
elif close <= lower and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""突破 均线+K×ATR 进 / 跌破 均线-K×ATR 出(与 next() 同一比较)。"""
close = self.data.close.raw
upper = self.ma + self.p["k"] * self.atr
lower = self.ma - self.p["k"] * self.atr
return close >= upper, close <= lower
# ── FSL 分水岭指标 ────────────────────────────────────────────────────────────
@@ -595,3 +687,7 @@ class FslStrategy(ParametrizedStrategy):
self.buy()
elif self.dead[i] and self.position["size"] > 0:
self.sell()
def entry_exit_masks(self) -> tuple[Any, Any]:
"""与 next() 同源(gold/dead 即 next() 判定用的同一组掩码数组)。"""
return self.gold, self.dead
+43 -4
View File
@@ -255,6 +255,35 @@ class Strategy(ABC):
每根 bar 调用一次,用户通过 self.buy()/self.sell() 生成信号。
"""
# ── 向量化快速路径钩子(v1.28) ────────────────────────────────────────────
def entry_exit_masks(self) -> tuple[NDArray, NDArray] | None:
"""返回 (entry_mask, exit_mask) 布尔数组以启用向量化信号生成。
默认返回 None(不支持向量化,引擎走逐 bar 路径)。策略满足以下
约束时可覆写本方法,让 :class:`~easy_tdx.backtest.engine.BacktestEngine`
用 numpy 一次性产出信号数组(网格寻优提速的关键路径):
- ``init()`` 只注册指标(副作用仅写入 ``self`` 属性);
- ``next()`` 只读 ``self.data`` / 指标数组的**当根值**(或整段掩码),
仅调用无参 ``buy()`` / ``sell()``size=0 全仓、无限价、无止损止盈);
- 持仓判断只依赖「空仓 / 持仓」两态(``self.position["size"]`` 的
零与非零),无动态仓位。
语义约定(掩码与 ``next()`` 必须等价,由对拍单测守护):
- ``entry_mask[i]`` 在**空仓**时触发买入;``exit_mask[i]`` 在**持仓**
时触发卖出——等价于 ``if entry and 空仓: buy() elif exit and 持仓:
sell()``。金叉/死叉这类交替掩码天然满足;
- 同一根 bar 两个掩码**不应同时为 True**;
- 掩码只依赖当根及更早数据(无未来函数),NaN 参与的比较按 False
处理(与逐 bar 路径一致)。
Returns:
(entry_mask, exit_mask) 与 K 线等长的 bool 数组;None = 不支持。
"""
return None
# ── 指标注册 ───────────────────────────────────────────────────────────────
def I( # noqa: E743
@@ -450,6 +479,11 @@ def _datetime_to_int(arr: NDArray) -> NDArray:
"""将 datetime 数组转为 int (YYYYMMDD)。
向量化实现,自动处理 datetime64、objectTimestamp)和数值类型。
datetime 走 ``year*10000 + month*100 + day`` 整数算术而非
``strftime``——strftime 是逐元素格式化(Python 层循环),在
``StrategyDataProxy`` 每次 ``_bind_data`` 都要转一遍的热路径上实测
慢一个数量级(v1.28 性能工作的一部分;输出与 strftime 零填充完全一致,
NaT → NaN 的行为也一致)。
Args:
arr: datetime 数组(np.datetime64、pd.Timestamp 或数值)
@@ -460,14 +494,19 @@ def _datetime_to_int(arr: NDArray) -> NDArray:
if len(arr) == 0:
return np.array([], dtype=np.float64)
arr = np.asarray(arr)
def _ymd_to_int(series: pd.Series) -> NDArray:
out = (series.dt.year * 10000 + series.dt.month * 100 + series.dt.day).to_numpy(
dtype=np.float64
)
return np.asarray(out, dtype=np.float64)
# datetime64 → 向量化转换
if arr.dtype.kind == "M":
return np.asarray(pd.to_datetime(arr).strftime("%Y%m%d").astype(float), dtype=np.float64)
return _ymd_to_int(pd.to_datetime(pd.Series(arr)))
# object 数组(可能包含 Timestamp
if arr.dtype == object:
if len(arr) > 0 and isinstance(arr[0], pd.Timestamp | np.datetime64):
return np.asarray(
pd.to_datetime(arr).strftime("%Y%m%d").astype(float), dtype=np.float64
)
return _ymd_to_int(pd.to_datetime(pd.Series(arr)))
# 已经是数值类型
return arr.astype(np.float64)
+310
View File
@@ -0,0 +1,310 @@
"""向量化快速路径(v1.28)对拍与约束检测单测。
核心保证:**同一 df + 同参数下,向量化路径与逐 bar 路径的输出逐位一致**
performance / trades / equity_curve / positions 全比对)。
- 对拍覆盖:全部 19 个内置策略(默认参数)+ ma_cross/macd/boll/rsi 的非默认
参数组合 + warmup / 极低资金(买不足 1 手的退化路径)/ 非默认费率与成交价模式;
- 约束检测:``_vectorize_eligibility`` 的显式约束(无掩码 / 缠论注入)与
``signal_path`` 的 auto/vector/loop 语义。
"""
from __future__ import annotations
import math
from typing import Any
import numpy as np
import pandas as pd
import pytest
from easy_tdx.backtest.engine import BacktestEngine
from easy_tdx.backtest.strategies import get_registry
from easy_tdx.backtest.strategy import Strategy
from easy_tdx.MyTT import MA
# ── 合成行情(确定性) ───────────────────────────────────────────────────────
def _synthetic_ohlcv(n: int = 800, seed: int = 42, base: float = 20.0) -> pd.DataFrame:
"""确定性随机游走 OHLCV(对拍两路径用同一份 df)。"""
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 _oscillating_ohlcv(n: int = 800, seed: int = 8, base: float = 20.0) -> pd.DataFrame:
"""周期振荡 OHLCV(正弦 + 噪声,high/low 恰为 max/min(open, close))。
donchian / wr_reversal 这类「突破 N 日高进 / 跌破 N 日低出」的策略,
其通道窗口包含当根(upper ≥ 当根 high),``close >= upper`` 只有在
close == high == 窗口最大(即收盘即创新高)时才成立——high/low 不放大,
正弦行情每个周期顶/底都会双向触发,覆盖完整的开平仓循环。
"""
rng = np.random.default_rng(seed)
t = np.arange(n)
close = base * (1.0 + 0.15 * np.sin(2 * np.pi * t / 40) + rng.normal(0.0, 0.002, n))
open_ = np.concatenate([[base], close[:-1]])
high = np.maximum(open_, close)
low = np.minimum(open_, close)
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 _assert_perf_equal(a: dict[str, Any], b: dict[str, Any]) -> None:
"""performance 字典逐位比对(NaN 视为相等——两条路径应产生完全相同的浮点数)。"""
assert set(a.keys()) == set(b.keys()), f"键集不一致: {set(a) ^ set(b)}"
for key in a:
va, vb = a[key], b[key]
if isinstance(va, float) and isinstance(vb, float) and math.isnan(va) and math.isnan(vb):
continue
assert va == vb, f"performance[{key}] 不一致: loop={va!r} vector={vb!r}"
def _assert_results_identical(loop: Any, vec: Any) -> None:
_assert_perf_equal(loop.performance, vec.performance)
pd.testing.assert_frame_equal(loop.trades, vec.trades)
pd.testing.assert_frame_equal(loop.equity_curve, vec.equity_curve)
pd.testing.assert_frame_equal(loop.positions, vec.positions)
assert loop.config == vec.config
def _run_both(
name: str,
df: pd.DataFrame,
params: dict[str, Any] | None = None,
*,
skip_bounds: bool = False,
**engine_kw: Any,
):
"""同一配置下分别跑逐 bar / 自动画两条路径。"""
entry = get_registry().get(name)
loop = BacktestEngine(
entry.build(params, skip_bounds=skip_bounds), signal_path="loop", **engine_kw
).run(df)
vec = BacktestEngine(
entry.build(params, skip_bounds=skip_bounds), signal_path="auto", **engine_kw
).run(df)
return loop, vec
# ── 对拍:19 个内置策略 × 默认参数 ───────────────────────────────────────────
#: 已知「默认参数下不会交易」的策略:MyTT 的 WR 是 0~100 刻度(100=超卖),
#: 而 wr_reversal 默认阈值为 -80/-20(通达信 -100~0 惯例),entry 恒 False。
#: 这是策略的既有行为(两路径一致地不交易),语义修正不属于向量化改动范围;
#: 对拍改用非默认参数覆盖(见 test_key_strategies_alternate_params)。
_DEAD_DEFAULT_STRATEGIES = {"wr_reversal"}
@pytest.mark.parametrize("name", get_registry().names())
def test_all_builtin_strategies_default_params(name: str) -> None:
"""全部内置策略:向量化与逐 bar 输出逐位一致(对拍核心保证)。"""
df = _synthetic_ohlcv()
loop, vec = _run_both(name, df)
# 确认确实产生了交易(空交易的对拍没有意义):随机游走无交易换种子,
# 再无交易换振荡行情(donchian/wr 等带状策略只在振荡行情双向触发)
if len(loop.trades) == 0:
loop, vec = _run_both(name, _synthetic_ohlcv(seed=7, base=50.0))
if len(loop.trades) == 0:
loop, vec = _run_both(name, _oscillating_ohlcv())
if name not in _DEAD_DEFAULT_STRATEGIES:
assert len(loop.trades) > 0, f"{name} 在三组行情下均无交易,对拍无效"
_assert_results_identical(loop, vec)
# ── 对拍:关键策略的非默认参数 + 引擎配置变化 ───────────────────────────────
@pytest.mark.parametrize(
("name", "params"),
[
("ma_cross", {"fast": 10, "slow": 60}),
("ma_cross", {"fast": 3, "slow": 8}),
("macd", {"short": 6, "long": 13, "signal": 5}),
("boll_breakout", {"n": 10, "p": 2.5}),
("rsi_reversal", {"n": 7, "oversold": 25, "overbought": 78}),
("donchian", {"n": 20}),
],
)
def test_key_strategies_alternate_params(name: str, params: dict[str, Any]) -> None:
"""任务点名的四个策略(含非默认参数)对拍一致。"""
df = _synthetic_ohlcv(seed=99)
loop, vec = _run_both(name, df, params)
if len(loop.trades) == 0:
loop, vec = _run_both(name, _oscillating_ohlcv(), params)
assert len(loop.trades) > 0, f"{name}{params} 两组行情下均无交易,对拍无效"
_assert_results_identical(loop, vec)
def test_wr_reversal_parity_with_skip_bounds() -> None:
"""wr_reversal 的向量化机制对拍(需 skip_bounds 越过死的默认边界)。
其阈值参数边界为负数区间(-100~-40 / -60~0),而 MyTT 的 WR 是 0~100
刻度——任何合法参数都无法触发交易(策略现状如此,两路径行为一致)。
为了让它的掩码/状态机路径也被对拍覆盖,用 skip_bounds 传 0~100 刻度内
的阈值绕过边界(寻优器同款机制)。
"""
df = _oscillating_ohlcv()
loop, vec = _run_both(
"wr_reversal", df, {"n": 14, "oversold": 40, "overbought": 60}, skip_bounds=True
)
assert len(loop.trades) > 0
_assert_results_identical(loop, vec)
@pytest.mark.parametrize("warmup", [0, 20, 100])
def test_warmup_bars_consistency(warmup: int) -> None:
"""warmup 期不产生信号:两路径一致(向量化按候选 bar 过滤)。"""
df = _synthetic_ohlcv(seed=11)
loop, vec = _run_both("ma_cross", df, warmup_bars=warmup)
_assert_results_identical(loop, vec)
def test_degenerate_low_cash_consistency() -> None:
"""极低资金:BUY 信号买不足 1 手(策略仓位状态不变),两路径仍一致。"""
df = _synthetic_ohlcv(seed=5, base=200.0) # 高价股 + 小资金 → 整手买入失败
loop, vec = _run_both("ma_cross", df, cash=1500.0)
_assert_results_identical(loop, vec)
def test_nondefault_fees_and_execution_consistency() -> None:
"""非默认费率/滑点/成交价模式:信号路径无关下游,但全流程仍应一致。"""
df = _synthetic_ohlcv(seed=13)
kw: dict[str, Any] = {
"commission": 0.0005,
"stamp_tax": 0.0005,
"slippage": 0.01,
"execution": "next_close",
}
loop, vec = _run_both("macd", df, **kw)
_assert_results_identical(loop, vec)
def test_indicator_cache_consistency() -> None:
"""挂载指标缓存(寻优场景)时向量化路径照常工作且一致。"""
from easy_tdx.backtest.indicator_cache import IndicatorCache
df = _synthetic_ohlcv(seed=21)
loop = BacktestEngine(
get_registry().get("ma_cross").build(), signal_path="loop", indicator_cache=IndicatorCache()
).run(df)
cache = IndicatorCache()
vec = BacktestEngine(
get_registry().get("ma_cross").build(), signal_path="auto", indicator_cache=cache
).run(df)
_assert_results_identical(loop, vec)
# ── 约束检测(显式、可测试) ─────────────────────────────────────────────────
class _PlainStrategy(Strategy):
"""未实现 entry_exit_masks 的普通策略(应走逐 bar)。"""
def init(self) -> None:
self.ma = self.I(MA, self.data.close, 5)
def next(self) -> None:
if self.ma[self._bar_index] > 0 and self.position["size"] == 0:
self.buy()
def test_eligibility_requires_masks_hook() -> None:
"""未覆写 entry_exit_masks → 不具备资格(原因可读)。"""
engine = BacktestEngine(_PlainStrategy)
eligible, reason = engine._vectorize_eligibility(_PlainStrategy(), None)
assert eligible is False
assert "entry_exit_masks" in reason
def test_eligibility_rejects_chanlun() -> None:
"""缠论注入(result 或 level)→ 不具备资格。"""
strat = get_registry().get("ma_cross").build()
engine = BacktestEngine(strat, chanlun_level="DAILY")
eligible, reason = engine._vectorize_eligibility(strat, None)
assert eligible is False
assert "缠论" in reason
engine2 = BacktestEngine(strat)
eligible2, reason2 = engine2._vectorize_eligibility(strat, {"fake": "chanlun"})
assert eligible2 is False
assert "缠论" in reason2
def test_eligibility_accepts_builtin() -> None:
"""内置策略(无缠论)→ 具备资格。"""
strat = get_registry().get("ma_cross").build()
engine = BacktestEngine(strat)
eligible, _ = engine._vectorize_eligibility(strat, None)
assert eligible is True
def test_signal_path_vector_forces_or_raises() -> None:
"""signal_path='vector':满足约束时正常,不满足时显式抛错。"""
df = _synthetic_ohlcv(200)
# 满足约束:正常运行且与 loop 一致
loop, vec = _run_both("ma_cross", df)
forced = BacktestEngine(get_registry().get("ma_cross").build(), signal_path="vector").run(df)
_assert_results_identical(loop, forced)
# 不满足约束:显式 ValueError(而非静默回退)
with pytest.raises(ValueError, match="向量化约束"):
BacktestEngine(_PlainStrategy, signal_path="vector").run(df)
def test_signal_path_invalid_rejected() -> None:
with pytest.raises(ValueError, match="signal_path"):
BacktestEngine(_PlainStrategy, signal_path="fast")
def test_auto_falls_back_on_mask_shape_mismatch() -> None:
"""掩码形状错误(策略实现 bug):auto 静默回退逐 bar,结果仍一致。"""
class _BadMaskStrategy(_PlainStrategy):
def entry_exit_masks(self) -> tuple[np.ndarray, np.ndarray]:
return np.zeros(3, dtype=bool), np.zeros(3, dtype=bool)
df = _synthetic_ohlcv(200, seed=3)
loop = BacktestEngine(_PlainStrategy, signal_path="loop").run(df)
fallback = BacktestEngine(_BadMaskStrategy, signal_path="auto").run(df)
_assert_results_identical(loop, fallback)
def test_vector_path_actually_used_for_builtins() -> None:
"""默认 signal_path='auto' 下内置策略确实走了向量化(防止回退被掩盖)。"""
from easy_tdx.backtest.strategy import Strategy as Base
for name in get_registry().names():
strat_cls = get_registry().get(name).strategy_cls
assert strat_cls.entry_exit_masks is not Base.entry_exit_masks, (
f"{name} 未实现 entry_exit_masksauto 将永远走逐 bar"
)