mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 20:24:19 +08:00
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:
@@ -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,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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user