401 lines
13 KiB
Python
401 lines
13 KiB
Python
"""
|
|
回测引擎模块 — 事件驱动回测,模拟交易执行与组合管理
|
|
"""
|
|
from dataclasses import dataclass, field
|
|
from typing import Any, Callable, Dict, List, Optional, Set, Tuple
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
|
|
from .config import AlphaConfig
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# 交易记录 & 持仓
|
|
# ---------------------------------------------------------------------------
|
|
|
|
@dataclass
|
|
class TradeRecord:
|
|
"""单笔交易记录"""
|
|
date: pd.Timestamp
|
|
stock: str
|
|
side: str # 'buy' / 'sell'
|
|
quantity: int
|
|
price: float
|
|
commission: float = 0.0
|
|
stamp_tax: float = 0.0
|
|
slippage_cost: float = 0.0
|
|
signal: str = ""
|
|
|
|
|
|
@dataclass
|
|
class Position:
|
|
"""单只股票持仓"""
|
|
stock: str
|
|
quantity: int = 0
|
|
avg_cost: float = 0.0
|
|
|
|
@property
|
|
def market_value(self) -> float:
|
|
return self.quantity * self.current_price if hasattr(self, "current_price") else 0.0
|
|
|
|
def update_cost(self, qty: int, price: float):
|
|
"""更新平均成本(买入时)"""
|
|
total_cost = abs(self.quantity) * self.avg_cost + abs(qty) * price
|
|
self.quantity += qty
|
|
if self.quantity != 0:
|
|
self.avg_cost = total_cost / abs(self.quantity)
|
|
else:
|
|
self.avg_cost = 0.0
|
|
|
|
|
|
class Portfolio:
|
|
"""投资组合"""
|
|
|
|
def __init__(self, initial_cash: float = 1_000_000.0):
|
|
self.initial_cash = initial_cash
|
|
self.cash = initial_cash
|
|
self.positions: Dict[str, Position] = {} # stock -> Position
|
|
self.trades: List[TradeRecord] = []
|
|
self.daily_values: List[Dict[str, Any]] = [] # 每日净值记录
|
|
|
|
@property
|
|
def total_equity(self) -> float:
|
|
pos_value = sum(
|
|
p.market_value for p in self.positions.values()
|
|
)
|
|
return self.cash + pos_value
|
|
|
|
@property
|
|
def total_return(self) -> float:
|
|
return (self.total_equity / self.initial_cash) - 1.0
|
|
|
|
def get_position(self, stock: str) -> Position:
|
|
if stock not in self.positions:
|
|
self.positions[stock] = Position(stock=stock)
|
|
return self.positions[stock]
|
|
|
|
def update_market_prices(self, prices: Dict[str, float]):
|
|
"""更新所有持仓的市价"""
|
|
for stock, price in prices.items():
|
|
if stock in self.positions:
|
|
self.positions[stock].current_price = price
|
|
|
|
def record_daily(self, date: pd.Timestamp, prices: Dict[str, float]):
|
|
"""记录每日快照"""
|
|
self.update_market_prices(prices)
|
|
pos_value = sum(p.market_value for p in self.positions.values())
|
|
self.daily_values.append({
|
|
"date": date,
|
|
"cash": self.cash,
|
|
"position_value": pos_value,
|
|
"total_equity": self.cash + pos_value,
|
|
"return": (self.cash + pos_value) / self.initial_cash - 1.0,
|
|
})
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# 券商(模拟交易执行)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class Broker:
|
|
"""模拟券商 — 处理订单执行、交易成本"""
|
|
|
|
def __init__(self, config: AlphaConfig):
|
|
self.commission_rate = config.commission_rate
|
|
self.slippage = config.slippage
|
|
self.stamp_tax = config.stamp_tax
|
|
self.max_position_pct = config.max_position_pct
|
|
self.max_turnover = config.max_turnover
|
|
self.min_holding_period = config.min_holding_period
|
|
|
|
def execute(
|
|
self,
|
|
portfolio: Portfolio,
|
|
target_weights: Dict[str, float],
|
|
prices: Dict[str, float],
|
|
date: pd.Timestamp,
|
|
strategy_name: str = "",
|
|
) -> List[TradeRecord]:
|
|
"""
|
|
执行调仓:比较当前持仓与目标权重,生成订单并执行。
|
|
返回新的交易记录列表。
|
|
"""
|
|
if not target_weights:
|
|
# 清仓信号
|
|
return self._liquidate(portfolio, prices, date, strategy_name)
|
|
|
|
total_equity = 0.0
|
|
# 先更新市价以计算当前权益
|
|
portfolio.update_market_prices(prices)
|
|
total_equity = portfolio.total_equity
|
|
|
|
trades: List[TradeRecord] = []
|
|
|
|
# 目标持仓市值
|
|
target_map: Dict[str, float] = {}
|
|
for stock, w in target_weights.items():
|
|
if stock in prices and prices[stock] > 0:
|
|
target_map[stock] = total_equity * w
|
|
|
|
# 卖出不在目标中的持仓
|
|
for stock in list(portfolio.positions.keys()):
|
|
pos = portfolio.positions[stock]
|
|
if pos.quantity <= 0:
|
|
continue
|
|
if stock not in target_map:
|
|
trades.extend(
|
|
self._sell(portfolio, stock, pos.quantity, prices, date,
|
|
strategy_name)
|
|
)
|
|
|
|
# 调整持仓到目标权重
|
|
for stock, target_value in target_map.items():
|
|
price = prices.get(stock, 0)
|
|
if price <= 0:
|
|
continue
|
|
current_pos = portfolio.get_position(stock)
|
|
current_value = current_pos.quantity * price
|
|
diff_value = target_value - current_value
|
|
if abs(diff_value) < price: # 差价不足一手,忽略
|
|
continue
|
|
|
|
diff_qty = int(diff_value / price / 100) * 100 # 整手
|
|
if diff_qty == 0:
|
|
continue
|
|
|
|
if diff_qty > 0:
|
|
trades.extend(
|
|
self._buy(portfolio, stock, diff_qty, price, date,
|
|
strategy_name)
|
|
)
|
|
else:
|
|
sell_qty = min(-diff_qty, current_pos.quantity)
|
|
trades.extend(
|
|
self._sell(portfolio, stock, sell_qty, price, date,
|
|
strategy_name)
|
|
)
|
|
|
|
return trades
|
|
|
|
def _buy(
|
|
self,
|
|
portfolio: Portfolio,
|
|
stock: str,
|
|
qty: int,
|
|
price: float,
|
|
date: pd.Timestamp,
|
|
strategy: str = "",
|
|
) -> List[TradeRecord]:
|
|
"""买入执行"""
|
|
# 滑点
|
|
exec_price = price * (1 + self.slippage)
|
|
cost = qty * exec_price
|
|
commission = cost * self.commission_rate
|
|
total_cost = cost + commission
|
|
|
|
if portfolio.cash < total_cost:
|
|
# 现金不足,调整数量
|
|
affordable_qty = int(
|
|
(portfolio.cash / (exec_price * (1 + self.commission_rate)))
|
|
/ 100
|
|
) * 100
|
|
if affordable_qty <= 0:
|
|
return []
|
|
qty = affordable_qty
|
|
cost = qty * exec_price
|
|
commission = cost * self.commission_rate
|
|
total_cost = cost + commission
|
|
|
|
portfolio.cash -= total_cost
|
|
pos = portfolio.get_position(stock)
|
|
pos.update_cost(qty, exec_price)
|
|
|
|
trade = TradeRecord(
|
|
date=date,
|
|
stock=stock,
|
|
side="buy",
|
|
quantity=qty,
|
|
price=exec_price,
|
|
commission=commission,
|
|
slippage_cost=qty * (exec_price - price),
|
|
signal=strategy,
|
|
)
|
|
portfolio.trades.append(trade)
|
|
return [trade]
|
|
|
|
def _sell(
|
|
self,
|
|
portfolio: Portfolio,
|
|
stock: str,
|
|
qty: int,
|
|
price: float,
|
|
date: pd.Timestamp,
|
|
strategy: str = "",
|
|
) -> List[TradeRecord]:
|
|
"""卖出执行"""
|
|
exec_price = price * (1 - self.slippage)
|
|
proceeds = qty * exec_price
|
|
commission = proceeds * self.commission_rate
|
|
stamp = proceeds * self.stamp_tax
|
|
net_proceeds = proceeds - commission - stamp
|
|
|
|
pos = portfolio.get_position(stock)
|
|
actual_qty = min(qty, pos.quantity)
|
|
if actual_qty <= 0:
|
|
return []
|
|
|
|
# 更新持仓
|
|
pos.quantity -= actual_qty
|
|
if pos.quantity == 0:
|
|
pos.avg_cost = 0.0
|
|
|
|
portfolio.cash += net_proceeds
|
|
|
|
trade = TradeRecord(
|
|
date=date,
|
|
stock=stock,
|
|
side="sell",
|
|
quantity=actual_qty,
|
|
price=exec_price,
|
|
commission=commission,
|
|
stamp_tax=stamp,
|
|
slippage_cost=actual_qty * (price - exec_price),
|
|
signal=strategy,
|
|
)
|
|
portfolio.trades.append(trade)
|
|
return [trade]
|
|
|
|
def _liquidate(
|
|
self,
|
|
portfolio: Portfolio,
|
|
prices: Dict[str, float],
|
|
date: pd.Timestamp,
|
|
strategy: str = "",
|
|
) -> List[TradeRecord]:
|
|
"""全部平仓"""
|
|
trades = []
|
|
for stock in list(portfolio.positions.keys()):
|
|
pos = portfolio.positions[stock]
|
|
if pos.quantity > 0 and stock in prices:
|
|
trades.extend(
|
|
self._sell(portfolio, stock, pos.quantity, prices, date,
|
|
strategy)
|
|
)
|
|
return trades
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# 回测引擎
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class BacktestEngine:
|
|
"""事件驱动回测引擎"""
|
|
|
|
def __init__(self, config: Optional[AlphaConfig] = None):
|
|
self.config = config or AlphaConfig()
|
|
self.broker = Broker(self.config)
|
|
self.portfolio = Portfolio(self.config.initial_cash)
|
|
|
|
def run(
|
|
self,
|
|
strategy: Any,
|
|
price_data: pd.DataFrame,
|
|
factor_data: Dict[str, pd.DataFrame],
|
|
dates: Optional[List[pd.Timestamp]] = None,
|
|
rebalance_freq: str = "M", # 'D'/'W'/'M'
|
|
progress_callback: Optional[Callable] = None,
|
|
) -> pd.DataFrame:
|
|
"""
|
|
运行回测。
|
|
|
|
参数
|
|
----
|
|
strategy : Strategy 实例或兼容接口
|
|
price_data : DataFrame, index=date, columns=stocks, values=价格
|
|
factor_data : {factor_name: DataFrame(index=date, columns=stocks)}
|
|
dates : 回测日期列表,默认为 price_data 所有日期
|
|
rebalance_freq: 调仓频率 'D'(日), 'W'(周), 'M'(月)
|
|
|
|
返回
|
|
----
|
|
daily_values : DataFrame 每日净值曲线
|
|
"""
|
|
if dates is None:
|
|
all_dates = sorted(price_data.index)
|
|
else:
|
|
all_dates = sorted(dates)
|
|
|
|
# 确定调仓日
|
|
date_series = pd.Series(all_dates, index=all_dates)
|
|
if rebalance_freq == "M":
|
|
rebalance_dates = date_series.resample("M").last().tolist()
|
|
elif rebalance_freq == "W":
|
|
rebalance_dates = date_series.resample("W").last().tolist()
|
|
else:
|
|
rebalance_dates = all_dates
|
|
|
|
rebalance_set = set(pd.to_datetime(rebalance_dates).date)
|
|
|
|
total_dates = len(all_dates)
|
|
for i, date in enumerate(all_dates):
|
|
# 进度回调
|
|
if progress_callback:
|
|
progress_callback(i, total_dates)
|
|
|
|
# 当前截面价格
|
|
current_prices = {}
|
|
if date in price_data.index:
|
|
row = price_data.loc[date]
|
|
current_prices = row.dropna().to_dict()
|
|
|
|
# 调仓日执行交易
|
|
trade_date = pd.Timestamp(date).date()
|
|
if trade_date in rebalance_set:
|
|
target_weights = strategy.run_step(
|
|
date=pd.Timestamp(date),
|
|
factor_data=factor_data,
|
|
prices=pd.DataFrame([current_prices]),
|
|
cash=self.portfolio.cash,
|
|
positions={
|
|
s: p.quantity
|
|
for s, p in self.portfolio.positions.items()
|
|
},
|
|
)
|
|
self.broker.execute(
|
|
self.portfolio,
|
|
target_weights,
|
|
current_prices,
|
|
date=pd.Timestamp(date),
|
|
strategy_name=strategy.name,
|
|
)
|
|
|
|
# 记录每日净值
|
|
self.portfolio.record_daily(pd.Timestamp(date), current_prices)
|
|
|
|
# 最终清算
|
|
if not all_dates:
|
|
return pd.DataFrame()
|
|
|
|
final_date = all_dates[-1]
|
|
if final_date in price_data.index:
|
|
final_prices = price_data.loc[final_date].dropna().to_dict()
|
|
else:
|
|
final_prices = {}
|
|
|
|
self.broker._liquidate(
|
|
self.portfolio, final_prices, pd.Timestamp(final_date),
|
|
strategy_name=strategy.name,
|
|
)
|
|
self.portfolio.record_daily(pd.Timestamp(final_date), final_prices)
|
|
|
|
return self._to_equity_curve()
|
|
|
|
def _to_equity_curve(self) -> pd.DataFrame:
|
|
"""输出净值曲线 DataFrame"""
|
|
df = pd.DataFrame(self.portfolio.daily_values)
|
|
if df.empty:
|
|
return pd.DataFrame(columns=["date", "total_equity", "return"])
|
|
df = df.set_index("date").sort_index()
|
|
df["nav"] = df["total_equity"] / self.config.initial_cash
|
|
return df |