diff --git a/CHANGELOG.md b/CHANGELOG.md index f6d13a1..119d2a4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,15 @@ ## [未发布] +### 性能 + +- **回测引擎信号管线提速 ×12.6(ma_cross 800 根全流程 93.6ms → 7.4ms,Windows/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 与真实客户端的契约。 diff --git a/scripts/bench_engine.py b/scripts/bench_engine.py new file mode 100644 index 0000000..753d40a --- /dev/null +++ b/scripts/bench_engine.py @@ -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()) diff --git a/src/easy_tdx/backtest/engine.py b/src/easy_tdx/backtest/engine.py index 2d5847c..aaedb46 100644 --- a/src/easy_tdx/backtest/engine.py +++ b/src/easy_tdx/backtest/engine.py @@ -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], diff --git a/src/easy_tdx/backtest/orders.py b/src/easy_tdx/backtest/orders.py index 05c45b8..393a18a 100644 --- a/src/easy_tdx/backtest/orders.py +++ b/src/easy_tdx/backtest/orders.py @@ -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 线索引。 diff --git a/src/easy_tdx/backtest/strategies/builtin.py b/src/easy_tdx/backtest/strategies/builtin.py index 5ea2609..3a7f750 100644 --- a/src/easy_tdx/backtest/strategies/builtin.py +++ b/src/easy_tdx/backtest/strategies/builtin.py @@ -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 diff --git a/src/easy_tdx/backtest/strategy.py b/src/easy_tdx/backtest/strategy.py index dbb36d0..16f3008 100644 --- a/src/easy_tdx/backtest/strategy.py +++ b/src/easy_tdx/backtest/strategy.py @@ -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、object(Timestamp)和数值类型。 + 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) diff --git a/tests/unit/test_backtest_engine_vector.py b/tests/unit/test_backtest_engine_vector.py new file mode 100644 index 0000000..92766fd --- /dev/null +++ b/tests/unit/test_backtest_engine_vector.py @@ -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_masks,auto 将永远走逐 bar" + )