43b6a58ba7
ashare_engine override load_data: 5m/15m 直查 dbbardata 转 BarData(MINUTE), 绕过 vnpy Interval enum 限制(原生只认 d/1h/1m)。cta_engine +interval 参数, 5m/15m 映射 MINUTE 过父类校验。api(schema/routes)+orchestrator 透传 interval。 顺带修 cta_engine DB路径污染(yaml NAS路径在Win VPS误解析→改读vt_setting.json)。 测试: 15m 600519一年 3888bars total_return-0.22 sharpe-2.04; 5m 11664bars; 日线未回归(237bars)。
274 lines
11 KiB
Python
274 lines
11 KiB
Python
"""A 股适配层:子类化 vnpy_ctastrategy BacktestingEngine + DailyResult。
|
||
|
||
vnpy 源码零修改,全部覆写在本文件。
|
||
|
||
- AShareDailyResult:A 股费用模型(佣金保底、印花税卖方、过户费沪市)
|
||
- AShareBacktestingEngine:
|
||
- send_order 覆写 → 做空拦截(SHORT+OPEN 拒单)+ 定寸重算 volume
|
||
- update_daily_close 覆写 → 工厂换 AShareDailyResult(父类在 :647 实例化 DailyResult)
|
||
- load_data 覆写 → 支持 5m/15m 周期(vnpy Interval enum 不认 '5m'/'15m')
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
from datetime import datetime
|
||
|
||
from vnpy_ctastrategy.backtesting import BacktestingEngine, DailyResult
|
||
from vnpy.trader.constant import Direction, Offset, Interval, Exchange
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
class AShareDailyResult(DailyResult):
|
||
"""A 股日度盈亏:佣金双边保底 + 印花税卖方 + 过户费沪市。
|
||
|
||
父类 calculate_pnl 用 turnover*rate 算佣金(单边商),A 股实际:
|
||
- 佣金 = max(turnover * commission_rate, min_commission),买卖都收
|
||
- 印花税 = turnover * stamp_duty_rate,卖方 only(direction==SHORT)
|
||
- 过户费 = turnover * transfer_fee_rate,沪市 only
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
date,
|
||
close_price: float,
|
||
commission_rate: float = 0.00025,
|
||
min_commission: float = 5.0,
|
||
stamp_duty_rate: float = 0.0005,
|
||
transfer_fee_rate: float = 0.00001,
|
||
is_sse: bool = False,
|
||
) -> None:
|
||
super().__init__(date, close_price)
|
||
self.commission_rate: float = commission_rate
|
||
self.min_commission: float = min_commission
|
||
self.stamp_duty_rate: float = stamp_duty_rate
|
||
self.transfer_fee_rate: float = transfer_fee_rate
|
||
self.is_sse: bool = is_sse
|
||
# 费用明细(持久化可见——父类 calculate_result 会遍历 __dict__ 落列)
|
||
self.stamp_duty: float = 0.0
|
||
self.transfer_fee: float = 0.0
|
||
|
||
def calculate_pnl(
|
||
self,
|
||
pre_close: float,
|
||
start_pos: float,
|
||
size: float,
|
||
rate: float,
|
||
slippage: float,
|
||
) -> None:
|
||
"""覆写父类:用 A 股费用模型替换 commission = turnover * rate。
|
||
|
||
签名与父类一致(calculate_result 传 self.rate/self.slippage/self.size),
|
||
但 rate 参数被忽略——佣金由 self.commission_rate + self.min_commission 决定。
|
||
"""
|
||
# 首日无 pre_close 时用 1 防除零(与父类逻辑一致)
|
||
self.pre_close = pre_close if pre_close else 1
|
||
|
||
self.start_pos = start_pos
|
||
self.end_pos = start_pos
|
||
|
||
self.holding_pnl = self.start_pos * (self.close_price - self.pre_close) * size
|
||
|
||
self.trade_count = len(self.trades)
|
||
for trade in self.trades:
|
||
if trade.direction == Direction.LONG:
|
||
pos_change = trade.volume
|
||
else:
|
||
pos_change = -trade.volume
|
||
|
||
self.end_pos += pos_change
|
||
|
||
turnover: float = trade.volume * size * trade.price
|
||
self.trading_pnl += pos_change * (self.close_price - trade.price) * size
|
||
self.slippage += trade.volume * size * slippage
|
||
|
||
self.turnover += turnover
|
||
# A 股佣金:双边,最低 min_commission 元
|
||
self.commission += max(turnover * self.commission_rate, self.min_commission)
|
||
|
||
# 印花税:卖方 only(SHORT = 卖出)
|
||
if trade.direction == Direction.SHORT:
|
||
self.stamp_duty += turnover * self.stamp_duty_rate
|
||
|
||
# 过户费:沪市 only
|
||
if self.is_sse:
|
||
self.transfer_fee += turnover * self.transfer_fee_rate
|
||
|
||
# net_pnl 扣除全部费用
|
||
self.total_pnl = self.trading_pnl + self.holding_pnl
|
||
self.net_pnl = (
|
||
self.total_pnl
|
||
- self.commission
|
||
- self.slippage
|
||
- self.stamp_duty
|
||
- self.transfer_fee
|
||
)
|
||
|
||
|
||
class AShareBacktestingEngine(BacktestingEngine):
|
||
"""A 股回测引擎:long-only 拦截 + A 股费用。
|
||
|
||
覆写点:
|
||
1. send_order → 拦截 SHORT+OPEN(A股不可做空)
|
||
2. update_daily_close → 工厂换 AShareDailyResult(父类 :647 实例化 DailyResult)
|
||
|
||
定寸(C1)通过 engine.size = N 实现,不在此处处理:
|
||
cta_engine 在 load_data 后取首根 bar close 算 N = floor(capital*pct/close/100)*100,
|
||
设置 engine.size = N。vnpy 的 turnover/PnL 自动 ×size,策略 volume 保持 1 手=N 股=满仓。
|
||
"""
|
||
|
||
def __init__(self) -> None:
|
||
super().__init__()
|
||
# A 股费用参数(默认值,cta_engine 可覆盖)
|
||
self.commission_rate: float = 0.00025 # 万 2.5
|
||
self.min_commission: float = 5.0 # 最低 5 元
|
||
self.stamp_duty_rate: float = 0.0005 # 卖方 0.05%
|
||
self.transfer_fee_rate: float = 0.00001 # 沪市 0.001%
|
||
self.is_sse: bool = False
|
||
# 5m/15m 适配:vnpy Interval enum 不认 '5m'/'15m',cta_engine 把 engine.interval
|
||
# 映射成 Interval.MINUTE 让父类 set_parameters 校验通过,真实 DB interval 字符串
|
||
# 存此字段供 load_data 自定义路径使用。默认 "d" 走 vnpy 原生日线路径。
|
||
self.raw_interval: str = "d"
|
||
self.sqlite_db_path: str | None = None # 由 cta_engine 注入(_dcfg.data_paths["vnpy_db"])
|
||
|
||
def send_order(
|
||
self,
|
||
strategy,
|
||
direction: Direction,
|
||
offset: Offset,
|
||
price: float,
|
||
volume: float,
|
||
stop: bool,
|
||
lock: bool,
|
||
net: bool,
|
||
) -> list:
|
||
"""覆写父类 send_order:做空拦截(C2)。
|
||
|
||
C2 做空拦截:SSE/SZSE 标的不可做空,SHORT+OPEN 直接拒单。
|
||
SHORT+CLOSE(平多)允许。
|
||
volume 不动——定寸由 engine.size = N 实现(见类文档)。
|
||
"""
|
||
if direction == Direction.SHORT and offset == Offset.OPEN:
|
||
logger.warning(
|
||
"A股不支持做空,拒单: direction=%s offset=%s price=%s volume=%s",
|
||
direction, offset, price, volume,
|
||
)
|
||
return []
|
||
|
||
return super().send_order(
|
||
strategy, direction, offset, price, volume, stop, lock, net
|
||
)
|
||
|
||
def update_daily_close(self, price: float) -> None:
|
||
"""覆写父类工厂方法:用 AShareDailyResult 替换 DailyResult。
|
||
|
||
父类原实现(backtesting.py:639-647):
|
||
daily_result = self.daily_results.get(d)
|
||
if daily_result: daily_result.close_price = price
|
||
else: self.daily_results[d] = DailyResult(d, price)
|
||
"""
|
||
d = self.datetime.date()
|
||
daily_result = self.daily_results.get(d, None)
|
||
if daily_result:
|
||
daily_result.close_price = price
|
||
else:
|
||
self.daily_results[d] = AShareDailyResult(
|
||
d,
|
||
price,
|
||
commission_rate=self.commission_rate,
|
||
min_commission=self.min_commission,
|
||
stamp_duty_rate=self.stamp_duty_rate,
|
||
transfer_fee_rate=self.transfer_fee_rate,
|
||
is_sse=self.is_sse,
|
||
)
|
||
|
||
def load_data(self) -> None:
|
||
"""覆写父类:支持 5m/15m 周期(绕过 vnpy Interval enum 限制)。
|
||
|
||
vnpy Interval enum 只有 MINUTE('1m')/HOUR('1h')/DAILY('d') 等,不认 '5m'/'15m'。
|
||
父类 load_data 调 ``INTERVAL_DELTA_MAP[self.interval]`` 和
|
||
``load_bar_data(..., self.interval, ...)`` 都依赖 enum,传 '15m' 会 ValueError。
|
||
|
||
适配思路:self.interval 已被父类 set_parameters 映射成 Interval.MINUTE(enum
|
||
校验通过),真实 DB interval 存 self.raw_interval。当 raw_interval 是 5m/15m
|
||
时,直接 sqlite 查 dbbardata 表,把每行转 BarData(interval=Interval.MINUTE),
|
||
绕过 vnpy database_manager 的 enum 限制;其余周期(d/1m 等)走 vnpy 原生路径。
|
||
"""
|
||
if self.raw_interval not in ("5m", "15m"):
|
||
super().load_data()
|
||
return
|
||
self._load_intraday_data()
|
||
|
||
def _load_intraday_data(self) -> None:
|
||
"""5m/15m 直查 SQLite → BarData(MINUTE),绕过 vnpy enum 限制。"""
|
||
from vnpy.trader.object import BarData
|
||
import sqlite3
|
||
|
||
self.output(f"开始加载 {self.raw_interval} 历史数据(ashare 适配)")
|
||
|
||
if not self.end:
|
||
self.end = datetime.now()
|
||
if self.start >= self.end:
|
||
self.output("起始日期必须小于结束日期")
|
||
return
|
||
|
||
db_path = self.sqlite_db_path
|
||
if not db_path:
|
||
try:
|
||
from vnpy.trader.setting import SETTINGS
|
||
db_path = SETTINGS.get("database.database")
|
||
except Exception:
|
||
db_path = None
|
||
if not db_path:
|
||
raise RuntimeError(
|
||
"5m/15m 回测需要 sqlite_db_path(或 SETTINGS['database.database']),"
|
||
"cta_engine 应在 load_data 前注入。"
|
||
)
|
||
|
||
symbol, exchange_str = self.vt_symbol.split(".")
|
||
# peewee DateTimeField 存的是 ISO 字符串;start/end 用 datetime 比较即可
|
||
# (SQLite 会把参数转成可比较的字符串形式)。
|
||
conn = sqlite3.connect(db_path)
|
||
try:
|
||
cur = conn.execute(
|
||
"SELECT datetime, volume, turnover, open_interest, "
|
||
"open_price, high_price, low_price, close_price "
|
||
"FROM dbbardata "
|
||
"WHERE symbol=? AND exchange=? AND interval=? "
|
||
"AND datetime>=? AND datetime<=? "
|
||
"ORDER BY datetime",
|
||
(symbol, exchange_str, self.raw_interval, self.start, self.end),
|
||
)
|
||
rows = cur.fetchall()
|
||
finally:
|
||
conn.close()
|
||
|
||
exchange = Exchange(exchange_str)
|
||
bars: list[BarData] = []
|
||
for dt, vol, turnover, oi, o, h, l, c in rows:
|
||
if isinstance(dt, str):
|
||
try:
|
||
dt = datetime.fromisoformat(dt)
|
||
except ValueError:
|
||
continue
|
||
# vnpy_sqlite save 时 convert_tz 改成 UTC,回测时按本地时间跑即可
|
||
# (日线回测也是直接读 DB datetime,行为一致)。
|
||
bars.append(BarData(
|
||
symbol=symbol,
|
||
exchange=exchange,
|
||
datetime=dt,
|
||
interval=Interval.MINUTE, # 5m/15m 不在 enum,统一标 MINUTE
|
||
volume=float(vol or 0),
|
||
turnover=float(turnover or 0),
|
||
open_interest=float(oi or 0),
|
||
open_price=float(o or 0),
|
||
high_price=float(h or 0),
|
||
low_price=float(l or 0),
|
||
close_price=float(c or 0),
|
||
gateway_name="sqlite",
|
||
))
|
||
|
||
self.history_data.clear()
|
||
self.history_data.extend(bars)
|
||
self.output(f"历史数据加载完成,数据量:{len(self.history_data)}")
|