Files
quanxiel/alpha/backtest.py
T
2026-07-31 21:23:35 +08:00

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