feat:添加alpha模块
This commit is contained in:
@@ -0,0 +1,401 @@
|
||||
"""
|
||||
回测引擎模块 — 事件驱动回测,模拟交易执行与组合管理
|
||||
"""
|
||||
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
|
||||
Reference in New Issue
Block a user