backtest_engine.py
"""
QTrading 回测引擎模块
"""
import pandas as pd
import numpy as np
from typing import List, Dict, Optional
from data.data_manager import DataManager
from strategy.strategy_manager import StrategyManager
from strategy.base_strategy import BaseStrategy
from config.config import config
from utils.logger import get_logger
from utils.common import validate_code, format_date, calculate_drawdown, today_str
logger = get_logger(__name__)
class BacktestEngine:
"""回测引擎"""
def __init__(self, tushare_token: str = None):
"""
初始化回测引擎
Args:
tushare_token: Tushare 接口令牌(可选)
"""
self.data_manager = DataManager(tushare_token)
self.strategy_manager = StrategyManager()
self.logger = logger
def backtest_strategy(
self,
strategy_name: str,
adjust_method: str = "qfq",
stock_list: List[str] = None,
start_date: str = None,
end_date: str = None,
holding_period: int = 20,
rebalance_frequency: int = 20,
initial_capital: float = None,
commission_rate: float = None,
slippage: float = None
) -> Dict:
"""
回测策略
Args:
strategy_name: 策略名称
adjust_method: 复权方式(qfq: 前复权, hfq: 后复权, none: 不复权)
stock_list: 股票池(可选,默认从配置文件读取)
start_date: 开始日期(YYYYMMDD,可选)
end_date: 结束日期(YYYYMMDD,可选)
holding_period: 持有周期(天数,默认20)
rebalance_frequency: 调仓频率(天数,默认20)
initial_capital: 初始资金(可选,默认从配置读取)
commission_rate: 佣金率(可选,默认从配置读取)
slippage: 滑点(可选,默认从配置读取)
Returns:
Dict: 回测结果
"""
if stock_list is None:
stock_list = config.stock_pool.stock_list
if start_date is None:
start_date = config.stock_pool.start_date
if end_date is None:
end_date = today_str()
if initial_capital is None:
initial_capital = config.backtest.initial_capital
if commission_rate is None:
commission_rate = config.backtest.commission_rate
if slippage is None:
slippage = config.backtest.slippage
# 创建策略实例
strategy = self.strategy_manager.create_strategy_instance(strategy_name, adjust_method)
if strategy is None:
self.logger.error(f"策略 {strategy_name} 不存在或创建失败")
return {}
self.logger.info(f"开始回测策略: {strategy_name}")
self.logger.info(f"回测参数:")
self.logger.info(f" - 复权方式: {adjust_method}")
self.logger.info(f" - 股票池: {len(stock_list)} 只")
self.logger.info(f" - 日期范围: {start_date} - {end_date}")
self.logger.info(f" - 持有周期: {holding_period} 天")
self.logger.info(f" - 调仓频率: {rebalance_frequency} 天")
self.logger.info(f" - 初始资金: {initial_capital}")
self.logger.info(f" - 佣金率: {commission_rate:.2%}")
self.logger.info(f" - 滑点: {slippage:.2%}")
# 获取交易日历
all_trade_dates = self._get_trade_dates(start_date, end_date)
if len(all_trade_dates) < holding_period + rebalance_frequency:
self.logger.warning("回测时间窗口过短")
return {}
# 回测主逻辑
portfolio = {
'dates': [],
'nav': [initial_capital],
'positions': [],
'cash': [initial_capital],
'total_value': [initial_capital]
}
# 股票历史数据缓存
stock_data_cache = {}
for code in stock_list:
try:
data = self.data_manager.calculate_adjusted_data(
code, adjust_method, start_date, end_date
)
if data is not None and not data.empty:
stock_data_cache[code] = data.sort_values('trade_date')
except Exception as e:
self.logger.warning(f"获取股票 {code} 数据失败: {e}")
# 开始回测
position_count = 0
current_date = start_date
position_count = len([code for code in stock_list if code in stock_data_cache])
self.logger.info(f"有效股票数据: {position_count} 只")
for i, date in enumerate(all_trade_dates):
if i == 0:
continue
# 调仓逻辑
if i % rebalance_frequency == 0:
self._rebalance_portfolio(
strategy, stock_data_cache, date, portfolio, initial_capital,
commission_rate, slippage
)
# 更新持仓价值
self._update_portfolio_value(
stock_data_cache, date, portfolio
)
# 更新日期记录
portfolio['dates'].append(date)
# 计算回测指标
backtest_result = self._calculate_backtest_metrics(portfolio)
self.logger.info("回测完成")
return backtest_result
def _get_trade_dates(self, start_date: str, end_date: str) -> List[str]:
"""获取交易日历"""
trade_cal = self.data_manager.tushare.get_trade_cal(start_date, end_date)
if trade_cal is None or trade_cal.empty:
self.logger.warning("未找到交易日历,使用连续日期")
return []
trade_cal = trade_cal[trade_cal['is_open'] == 1]
return trade_cal['cal_date'].tolist()
def _rebalance_portfolio(
self, strategy: BaseStrategy, stock_data_cache: Dict, date: str,
portfolio: Dict, initial_capital: float, commission_rate: float, slippage: float
) -> None:
"""调仓策略"""
# 计算选股分数
stock_scores = []
for code in stock_data_cache:
try:
data = stock_data_cache[code]
data = data[data['trade_date'] <= date]
if len(data) < 20: # 需要至少20天数据
continue
score = strategy.score(data)
stock_scores.append({
'code': code,
'score': score
})
except Exception as e:
self.logger.warning(f"计算股票 {code} 分数失败: {e}")
continue
# 按分数降序排序,选择前 N 只
stock_scores.sort(key=lambda x: x['score'], reverse=True)
selected_codes = [item['code'] for item in stock_scores[:10]]
# 计算每只股票的持仓比例
position_size = 1.0 / len(selected_codes) if selected_codes else 0
# 记录当前持仓
portfolio['positions'].append(selected_codes)
# 计算需要购买的股票
current_positions = portfolio['positions'][-1] if portfolio['positions'] else []
cash = portfolio['cash'][-1]
# 简单模拟交易逻辑(不考虑持仓比例)
for code in selected_codes:
if code not in current_positions:
# 模拟买入
if code in stock_data_cache:
data = stock_data_cache[code]
data = data[data['trade_date'] <= date]
if not data.empty:
price = data['close'].iloc[-1]
shares = (cash * position_size) / (price * (1 + commission_rate + slippage))
# 简化处理,直接减去成本
cost = shares * price * (1 + commission_rate + slippage)
cash -= cost
# 记录现金
portfolio['cash'].append(cash)
def _update_portfolio_value(
self, stock_data_cache: Dict, date: str, portfolio: Dict
) -> None:
"""更新持仓价值"""
current_positions = portfolio['positions'][-1] if portfolio['positions'] else []
cash = portfolio['cash'][-1]
position_value = 0.0
for code in current_positions:
if code in stock_data_cache:
data = stock_data_cache[code]
data = data[data['trade_date'] <= date]
if not data.empty:
# 简化处理,直接取最新价格计算市值
price = data['close'].iloc[-1]
# 简单模拟持仓成本和数量(实际需要详细记录)
shares = 1000 # 简化假设
position_value += shares * price
total_value = cash + position_value
portfolio['total_value'].append(total_value)
# 计算净值(NAV)
nav = total_value / portfolio['total_value'][0]
portfolio['nav'].append(nav)
def _calculate_backtest_metrics(self, portfolio: Dict) -> Dict:
"""计算回测指标"""
if not portfolio['total_value'] or len(portfolio['total_value']) < 2:
return {}
# 计算收益相关指标
nav_series = pd.Series(portfolio['nav'])
returns = nav_series.pct_change().dropna()
total_return = (nav_series.iloc[-1] - 1) * 100
# 年化收益率
num_years = len(portfolio['dates']) / 252
annual_return = (nav_series.iloc[-1] ** (1 / num_years) - 1) * 100
# 最大回撤
max_drawdown = calculate_drawdown(nav_series.values)
# 波动率
annual_volatility = returns.std() * np.sqrt(252) * 100
# 夏普比率(假设无风险利率为3%)
risk_free_rate = 0.03
excess_returns = returns - (risk_free_rate / 252)
sharpe_ratio = (excess_returns.mean() / excess_returns.std()) * np.sqrt(252)
# 胜率
positive_returns = returns[returns > 0]
win_rate = len(positive_returns) / len(returns) * 100
# 盈亏比
average_win = positive_returns.mean()
negative_returns = returns[returns < 0]
average_loss = abs(negative_returns.mean())
profit_loss_ratio = average_win / average_loss if average_loss != 0 else 0
return {
'total_return': round(total_return, 2),
'annual_return': round(annual_return, 2),
'max_drawdown': round(max_drawdown, 2),
'annual_volatility': round(annual_volatility, 2),
'sharpe_ratio': round(sharpe_ratio, 2),
'win_rate': round(win_rate, 2),
'profit_loss_ratio': round(profit_loss_ratio, 2),
'total_trading_days': len(portfolio['dates']),
'nav_series': nav_series.tolist(),
'date_series': portfolio['dates']
}
def compare_strategies(
self, strategy_names: List[str], adjust_method: str = "qfq",
**kwargs
) -> List[Dict]:
"""比较多个策略的回测结果"""
comparison_results = []
for strategy_name in strategy_names:
result = self.backtest_strategy(strategy_name, adjust_method, **kwargs)
if result:
comparison_results.append({
'strategy_name': strategy_name,
'description': self.strategy_manager.get_strategy_description(strategy_name),
'backtest_result': result
})
return comparison_results
def run_parameter_sweep(
self, strategy_name: str, parameter_name: str, parameter_values: List,
adjust_method: str = "qfq", **kwargs
) -> Dict:
"""参数扫描回测"""
sweep_results = []
for value in parameter_values:
self.logger.info(f"正在测试参数 {parameter_name} = {value}")
result = self.backtest_strategy(strategy_name, adjust_method, **kwargs)
if result:
sweep_results.append({
parameter_name: value,
'backtest_result': result
})
# 找到最优结果
if sweep_results:
best_result = max(
sweep_results,
key=lambda x: x['backtest_result']['sharpe_ratio']
)
return {
'parameter_name': parameter_name,
'parameter_values': parameter_values,
'sweep_results': sweep_results,
'best_result': best_result
}
return {}