""" 回测引擎模块 — 事件驱动回测,模拟交易执行与组合管理 """ 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