feat(backtest): add TWAPExecution + VWAPExecution

This commit is contained in:
GitHub
2026-06-12 20:56:57 +08:00
parent 0772666be3
commit fe68d9da95
2 changed files with 489 additions and 1 deletions
+153 -1
View File
@@ -5,7 +5,12 @@ from __future__ import annotations
import pandas as pd
import pytest
from easy_tdx.backtest.execution import ExecutionModel, ImmediateExecution
from easy_tdx.backtest.execution import (
ExecutionModel,
ImmediateExecution,
TWAPExecution,
VWAPExecution,
)
from easy_tdx.backtest.types import Signal
@@ -167,3 +172,150 @@ class TestImmediateExecution:
)
assert len(trades) == 1
assert trades[0].size == 100
class TestTWAPExecution:
"""时间加权平均价格执行。"""
def test_split_buy_into_3_bars(self) -> None:
df = _make_df(20)
model = TWAPExecution(n_bars=3)
signal = Signal(datetime=20240101, direction="BUY", size=300)
trades = model.execute(
signal=signal,
df=df,
bar_idx=0,
cash=100000,
position=0,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) == 3
total_size = sum(t.size for t in trades)
assert total_size <= 300
prices = [t.price for t in trades]
assert prices[0] != prices[1]
def test_split_sell_into_2_bars(self) -> None:
df = _make_df(20)
model = TWAPExecution(n_bars=2)
signal = Signal(datetime=20240101, direction="SELL", size=200)
trades = model.execute(
signal=signal,
df=df,
bar_idx=0,
cash=0,
position=500,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) == 2
assert sum(t.size for t in trades) == 200.0
def test_truncates_at_data_end(self) -> None:
df = _make_df(5)
model = TWAPExecution(n_bars=10)
signal = Signal(datetime=20240101, direction="BUY", size=1000)
trades = model.execute(
signal=signal,
df=df,
bar_idx=0,
cash=100000,
position=0,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) <= 4
def test_full_position_mode(self) -> None:
df = _make_df(20)
model = TWAPExecution(n_bars=3)
signal = Signal(datetime=20240101, direction="BUY", size=0)
trades = model.execute(
signal=signal,
df=df,
bar_idx=0,
cash=60000,
position=0,
position_mode="full",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) == 3
assert all(t.size > 0 for t in trades)
class TestVWAPExecution:
"""成交量加权平均价格执行。"""
def test_basic_buy(self) -> None:
df = _make_df(20)
model = VWAPExecution(n_bars=3, volume_lookback=10)
signal = Signal(datetime=20240101, direction="BUY", size=300)
trades = model.execute(
signal=signal,
df=df,
bar_idx=5,
cash=100000,
position=0,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) == 3
total_size = sum(t.size for t in trades)
assert total_size <= 300
def test_volume_weighted_split(self) -> None:
df = _make_df(20)
df.loc[6, "volume"] = 50000
df.loc[7, "volume"] = 50000
df.loc[8, "volume"] = 50000
model = VWAPExecution(n_bars=3, volume_lookback=5)
signal = Signal(datetime=20240105, direction="BUY", size=300)
trades = model.execute(
signal=signal,
df=df,
bar_idx=5,
cash=100000,
position=0,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) == 3
sizes = [t.size for t in trades]
assert sum(sizes) <= 300
def test_truncates_at_data_end(self) -> None:
df = _make_df(5)
model = VWAPExecution(n_bars=10, volume_lookback=3)
signal = Signal(datetime=20240101, direction="BUY", size=1000)
trades = model.execute(
signal=signal,
df=df,
bar_idx=0,
cash=100000,
position=0,
position_mode="fixed",
commission=0.0003,
min_commission=5.0,
stamp_tax=0.001,
slippage_model=None,
)
assert len(trades) <= 4