feat(backtest): integrate SlippageModel into OrderSimulator

This commit is contained in:
GitHub
2026-06-12 20:50:42 +08:00
parent d081eeb265
commit 6414c2cc11
2 changed files with 114 additions and 3 deletions
+40 -3
View File
@@ -6,11 +6,16 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING
import numpy as np
import pandas as pd
from easy_tdx.backtest.types import Signal, Trade
if TYPE_CHECKING:
from easy_tdx.backtest.slippage import SlippageModel
@dataclass
class OrderSimulator:
@@ -39,6 +44,7 @@ class OrderSimulator:
min_commission: float = 5.0
stamp_tax: float = 0.001
slippage: float = 0.0
slippage_model: SlippageModel | None = None
future_leak_warning: bool = False
def simulate(
@@ -288,6 +294,37 @@ class OrderSimulator:
return commission
def _compute_slippage(self, size: float, price: float, is_sell: bool) -> float:
"""计算滑点成本。"""
if self.slippage_model is not None:
volume = self._get_current_volume()
volatility = self._estimate_volatility()
return self.slippage_model.compute(
price=price,
size=size,
volume=volume,
volatility=volatility,
direction="SELL" if is_sell else "BUY",
)
return size * self.slippage
def _get_current_volume(self) -> float:
"""获取最后一根K线的成交量。"""
if "volume" in self.df.columns and len(self.df) > 0:
return float(self.df["volume"].iloc[-1])
return 0.0
def _estimate_volatility(self) -> float:
"""从收盘价估计年化波动率。"""
if "close" not in self.df.columns or len(self.df) < 2:
return 0.0
close = self.df["close"].to_numpy()
returns = np.diff(close) / close[:-1]
if len(returns) < 2:
return 0.0
daily_vol = float(np.std(returns))
return daily_vol * np.sqrt(252)
def _execute_buy(
self,
signal: Signal,
@@ -338,7 +375,7 @@ class OrderSimulator:
# 计算费用
commission = self._calculate_commission(size, price, is_sell=False)
slippage = size * self.slippage
slippage = self._compute_slippage(size, price, is_sell=False)
# 检查资金是否足够
total_cost = size * price + commission + slippage
@@ -361,7 +398,7 @@ class OrderSimulator:
reduced_size = int(available_cash / price / 100) * 100
if reduced_size > 0:
commission = self._calculate_commission(reduced_size, price, is_sell=False)
slippage = reduced_size * self.slippage
slippage = self._compute_slippage(reduced_size, price, is_sell=False)
return Trade(
datetime=self.df.iloc[exec_idx]["datetime"],
direction="BUY",
@@ -446,7 +483,7 @@ class OrderSimulator:
# 计算费用
commission = self._calculate_commission(size, price, is_sell=True)
slippage = size * self.slippage
slippage = self._compute_slippage(size, price, is_sell=True)
return Trade(
datetime=self.df.iloc[exec_idx]["datetime"],
+74
View File
@@ -3,8 +3,10 @@
from __future__ import annotations
import pandas as pd
import pytest
from easy_tdx.backtest.orders import OrderSimulator
from easy_tdx.backtest.slippage import FixedSlippage, PercentSlippage
from easy_tdx.backtest.types import Signal
# ── Test Fixtures ─────────────────────────────────────────────────────────────
@@ -390,3 +392,75 @@ class TestEdgeCases:
# 简化:只验证成交记录
assert len(trades) == 1
assert trades[0].size == 100
# ── Test SlippageModel Integration ─────────────────────────────────────────────
class TestSlippageModelIntegration:
"""测试 OrderSimulator 与 SlippageModel 集成。"""
def test_fixed_slippage_model(self) -> None:
"""FixedSlippage 与旧 slippage 参数等价。"""
df = _make_df(10)
sim = OrderSimulator(
df,
execution="next_open",
slippage_model=FixedSlippage(per_share=0.01),
)
signals = [_buy_signal(0, size=100)]
trades = sim.simulate(signals, cash=20000, position=0)
assert len(trades) == 1
assert trades[0].slippage == pytest.approx(1.0)
def test_percent_slippage_model(self) -> None:
"""PercentSlippage 计算。"""
df = _make_df(10)
sim = OrderSimulator(
df,
execution="next_open",
slippage_model=PercentSlippage(rate=0.001),
)
signals = [_buy_signal(0, size=100)]
trades = sim.simulate(signals, cash=20000, position=0)
assert len(trades) == 1
# price=101 (next_open), 101 × 100 × 0.001 = 10.1
assert trades[0].slippage == pytest.approx(10.1)
def test_slippage_model_overrides_slippage_param(self) -> None:
"""slippage_model 优先于 slippage 参数。"""
df = _make_df(10)
sim = OrderSimulator(
df,
execution="next_open",
position_mode="fixed",
slippage=999.0,
slippage_model=FixedSlippage(per_share=0.01),
)
signals = [_buy_signal(0, size=100)]
trades = sim.simulate(signals, cash=20000, position=0)
assert len(trades) == 1
assert trades[0].slippage == pytest.approx(1.0)
def test_sell_with_slippage_model(self) -> None:
"""卖出时也使用滑点模型。"""
df = _make_df(10)
sim = OrderSimulator(
df,
execution="next_open",
slippage_model=FixedSlippage(per_share=0.02),
)
signals = [_sell_signal(0, size=100)]
trades = sim.simulate(signals, cash=0, position=100)
assert len(trades) == 1
# position_mode=full, size=0 → sell all position=100
assert trades[0].slippage == pytest.approx(2.0)
def test_no_slippage_model_uses_old_param(self) -> None:
"""不提供 model 时使用旧 slippage 参数。"""
df = _make_df(10)
sim = OrderSimulator(df, execution="next_open", slippage=0.05)
signals = [_buy_signal(0, size=100)]
trades = sim.simulate(signals, cash=20000, position=0)
assert len(trades) == 1
assert trades[0].slippage == pytest.approx(5.0)