backtester.py
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
全样本回测模块
严格模拟真实交易,对历史每一个交易日做全样本回测
"""
import os
import logging
import pandas as pd
import numpy as np
import json
from pathlib import Path
from datetime import datetime, timedelta
from dotenv import load_dotenv
# 加载环境变量
load_dotenv()
# 配置日志
logger = logging.getLogger(__name__)
class Backtester:
"""全样本回测类"""
def __init__(self):
"""初始化回测对象"""
self.start_date = os.getenv("BACKTEST_START_DATE", "20260101")
self.end_date = os.getenv("BACKTEST_END_DATE", "20260204")
# 回测结果保存目录
self.result_dir = "results/backtests"
Path(self.result_dir).mkdir(parents=True, exist_ok=True)
# 初始化数据管理模块
from src.data_manager import DataManager
self.data_manager = DataManager()
# 初始化大模型模块
from src.model_api import ModelAPI
self.model_api = ModelAPI()
# 加载最优参数
self.best_params = self._load_best_parameters()
def _load_best_parameters(self):
"""加载最优参数"""
params_file = "results/optimization/best_parameters.json"
if Path(params_file).exists():
try:
with open(params_file, 'r', encoding='utf-8') as f:
return json.load(f)
except Exception as e:
logger.error(f"加载最优参数失败:{e}")
# 默认参数
return {
"learning_window": "90日",
"profit_threshold": "3%",
"stock_count": "10只",
"volume_ratio_threshold": "1.5倍"
}
def run(self, params=None, start_date=None, end_date=None):
"""执行全样本回测"""
try:
logger.info("开始全样本回测...")
# 使用传入的参数或默认参数
if not params:
params = self.best_params
if not start_date:
start_date = self.start_date
if not end_date:
end_date = self.end_date
logger.info(f"回测参数:{params}")
logger.info(f"回测日期范围:{start_date} - {end_date}")
# 获取回测所需的所有数据
logger.info("获取回测数据...")
all_stock_data = self.data_manager.get_all_stock_data_for_backtest(start_date, end_date)
if not all_stock_data:
logger.warning("无回测数据,回测终止")
return self._get_empty_report()
logger.info(f"回测股票数量:{len(all_stock_data)}")
# 获取所有交易日
trade_dates = self._get_trading_dates(start_date, end_date)
# 初始化回测结果
backtest_results = []
# 对每个交易日进行回测
for date_idx, trade_date in enumerate(trade_dates):
logger.info(f"正在回测 {trade_date} ({date_idx+1}/{len(trade_dates)})")
# 获取当日数据
daily_data = self._get_daily_data(all_stock_data, trade_date)
# 筛选有效股票
valid_stocks = self._filter_valid_stocks(daily_data)
# 解析参数值,去掉单位(如'只'、'日'、'%'等)
def parse_param(value):
if isinstance(value, str):
# 去掉数字后面的单位
import re
match = re.search(r'\d+', value)
if match:
return int(match.group())
else:
return 0
return value
if len(valid_stocks) < parse_param(params['stock_count']):
logger.warning(f"{trade_date} 有效股票数量不足,跳过回测")
continue
# 获取前一日数据(用于计算收益率)
prev_date = self._get_prev_trading_date(trade_date, trade_dates)
prev_data = self._get_daily_data(all_stock_data, prev_date)
# 获取后一日数据(用于计算T+1收益率)
next_date = self._get_next_trading_date(trade_date, trade_dates)
next_data = self._get_daily_data(all_stock_data, next_date)
# 计算收益率
returns = self._calculate_returns(prev_data, daily_data, next_data)
# 选股(模拟大模型评分)
selected_stocks = self._simulate_stock_selection(valid_stocks, params)
# 计算选股收益率
stock_returns = []
for stock_code, stock in selected_stocks.items():
if stock_code in returns:
stock_returns.append(returns[stock_code])
# 记录回测结果
backtest_results.append({
"trade_date": trade_date,
"selected_stocks": list(selected_stocks.keys()),
"stock_count": len(selected_stocks),
"avg_return": np.mean(stock_returns) if stock_returns else 0,
"max_return": np.max(stock_returns) if stock_returns else 0,
"min_return": np.min(stock_returns) if stock_returns else 0,
"win_rate": self._calculate_win_rate(stock_returns),
"profit_loss_ratio": self._calculate_profit_loss_ratio(stock_returns)
})
# 生成回测报告
report = self._generate_backtest_report(backtest_results)
logger.info("回测完成")
return report
except Exception as e:
logger.error(f"回测失败:{e}")
return self._get_empty_report()
def _get_trading_dates(self, start_date, end_date):
"""获取交易日列表"""
try:
# 使用数据管理模块的方法获取交易日列表
trade_dates = self.data_manager._get_trading_dates(start_date, end_date)
logger.info(f"交易日数量:{len(trade_dates)}")
return trade_dates
except Exception as e:
logger.error(f"交易日获取失败:{e}")
return []
def _get_daily_data(self, all_stock_data, trade_date):
"""获取指定日期的所有股票数据"""
daily_data = {}
for ts_code, df_stock in all_stock_data.items():
date_mask = df_stock['trade_date'] == trade_date
if date_mask.any():
daily_data[ts_code] = df_stock[date_mask].iloc[0]
return daily_data
def _filter_valid_stocks(self, daily_data):
"""筛选有效股票"""
valid_stocks = {}
for ts_code, data in daily_data.items():
# 过滤ST股票和停牌股票
# 从股票基本信息表中获取is_st字段
stock_basic = self.data_manager.get_valid_stocks()
if ts_code in stock_basic['ts_code'].values:
continue
# 过滤换手率过低的股票
if data.get('turnover_rate', 0) < float(os.getenv("MIN_TURNOVER_RATE", 1.0)):
continue
# 过滤成交量过低的股票
if data.get('volume', 0) < 1000000:
continue
valid_stocks[ts_code] = data
logger.info(f"有效股票数量:{len(valid_stocks)}")
return valid_stocks
def _get_prev_trading_date(self, trade_date, trade_dates):
"""获取前一个交易日"""
idx = trade_dates.index(trade_date)
if idx > 0:
return trade_dates[idx - 1]
return None
def _get_next_trading_date(self, trade_date, trade_dates):
"""获取后一个交易日"""
idx = trade_dates.index(trade_date)
if idx < len(trade_dates) - 1:
return trade_dates[idx + 1]
return None
def _calculate_returns(self, prev_data, daily_data, next_data):
"""计算股票收益率"""
returns = {}
for ts_code, daily in daily_data.items():
if ts_code in next_data:
# 计算T+1收益率(收盘买入,次日任意时点卖出)
# 简单模拟:T+1日收盘价/当日收盘价 - 1
returns[ts_code] = (next_data[ts_code]['close'] / daily['close']) - 1
return returns
def _simulate_stock_selection(self, valid_stocks, params):
"""模拟大模型选股(使用历史数据规律)"""
# 简单模拟:基于历史数据和参数筛选股票
# 实际应用中需要调用大模型API进行评分
# 按成交量排名
sorted_stocks = sorted(
valid_stocks.items(),
key=lambda x: x[1].get('volume', 0),
reverse=True
)
# 选择前N只股票
stock_count = int(params['stock_count'])
selected = sorted_stocks[:stock_count]
# 转换为字典
return dict(selected)
def _calculate_win_rate(self, returns):
"""计算胜率"""
if not returns:
return 0
winning_returns = [r for r in returns if r > 0]
return len(winning_returns) / len(returns)
def _calculate_profit_loss_ratio(self, returns):
"""计算盈亏比"""
if not returns:
return 0
winning_returns = [r for r in returns if r > 0]
losing_returns = [abs(r) for r in returns if r < 0]
if not winning_returns or not losing_returns:
return 0
avg_win = np.mean(winning_returns)
avg_loss = np.mean(losing_returns)
return avg_win / avg_loss
def _generate_backtest_report(self, backtest_results):
"""生成回测报告"""
if not backtest_results:
return self._get_empty_report()
report = {
"total_days": len(backtest_results),
"total_trades": sum(result['stock_count'] for result in backtest_results),
"win_rate": 0,
"avg_return": 0,
"total_return": 0,
"max_drawdown": 0,
"profit_loss_ratio": 0,
"avg_daily_return": 0
}
# 计算各项指标
all_returns = []
daily_returns = []
for result in backtest_results:
daily_returns.append(result['avg_return'])
all_returns.extend([
result['avg_return'] / result['stock_count']
for _ in range(result['stock_count'])
])
report['win_rate'] = self._calculate_win_rate(all_returns)
report['avg_return'] = np.mean(all_returns)
report['total_return'] = np.prod([1 + r for r in daily_returns]) - 1
report['avg_daily_return'] = np.mean(daily_returns)
report['profit_loss_ratio'] = self._calculate_profit_loss_ratio(all_returns)
# 计算最大回撤
cumulative_returns = np.cumprod([1 + r for r in daily_returns])
running_max = np.maximum.accumulate(cumulative_returns)
drawdowns = (cumulative_returns - running_max) / running_max
report['max_drawdown'] = np.min(drawdowns)
# 格式化输出
report = self._format_report(report)
logger.info("回测报告生成完成")
return report
def _format_report(self, report):
"""格式化报告"""
return {
"total_days": report['total_days'],
"total_trades": report['total_trades'],
"win_rate": round(report['win_rate'], 4),
"avg_return": round(report['avg_return'], 4),
"total_return": round(report['total_return'], 4),
"max_drawdown": round(report['max_drawdown'], 4),
"profit_loss_ratio": round(report['profit_loss_ratio'], 2),
"avg_daily_return": round(report['avg_daily_return'], 4)
}
def _get_empty_report(self):
"""获取空回测报告"""
return {
"total_days": 0,
"total_trades": 0,
"win_rate": 0,
"avg_return": 0,
"total_return": 0,
"max_drawdown": 0,
"profit_loss_ratio": 0,
"avg_daily_return": 0
}
def save_report(self, report, trade_date=None):
"""保存回测报告"""
try:
if not trade_date:
trade_date = datetime.now().strftime('%Y%m%d')
# 保存CSV文件
csv_path = Path(self.result_dir) / f"{trade_date}_backtest_report.csv"
pd.DataFrame([report]).to_csv(csv_path, index=False, encoding='utf-8-sig')
# 保存文本文件
txt_path = Path(self.result_dir) / f"{trade_date}_backtest_report.txt"
with open(txt_path, 'w', encoding='utf-8') as f:
f.write("=" * 80 + "\n")
f.write(" 超短线量化选股回测报告\n")
f.write("=" * 80 + "\n")
f.write(f"回测日期:{self.start_date} - {self.end_date}\n")
f.write(f"回测股票数量:{len(self.data_manager.get_all_stock_data_for_backtest())}\n")
f.write("-" * 80 + "\n")
f.write(f"总交易天数:{report['total_days']}天\n")
f.write(f"总交易次数:{report['total_trades']}次\n")
f.write(f"胜率:{report['win_rate']:.2%}\n")
f.write(f"平均单只收益率:{report['avg_return']:.2%}\n")
f.write(f"累计收益率:{report['total_return']:.2%}\n")
f.write(f"最大回撤:{report['max_drawdown']:.2%}\n")
f.write(f"盈亏比:{report['profit_loss_ratio']:.2f}\n")
f.write(f"每日平均收益率:{report['avg_daily_return']:.2%}\n")
f.write("=" * 80 + "\n")
logger.info(f"回测报告已保存:")
logger.info(f"CSV文件:{csv_path}")
logger.info(f"文本文件:{txt_path}")
return True
except Exception as e:
logger.error(f"回测报告保存失败:{e}")
return False
def print_report(self, report):
"""打印回测报告"""
logger.info("=" * 80)
logger.info("超短线量化选股回测报告")
logger.info("=" * 80)
print(f"回测日期:{self.start_date} - {self.end_date}")
print(f"回测股票数量:{len(self.data_manager.get_all_stock_data_for_backtest())}")
print("-" * 80)
print(f"总交易天数:{report['total_days']}天")
print(f"总交易次数:{report['total_trades']}次")
print(f"胜率:{report['win_rate']:.2%}")
print(f"平均单只收益率:{report['avg_return']:.2%}")
print(f"累计收益率:{report['total_return']:.2%}")
print(f"最大回撤:{report['max_drawdown']:.2%}")
print(f"盈亏比:{report['profit_loss_ratio']:.2f}")
print(f"每日平均收益率:{report['avg_daily_return']:.2%}")
print("=" * 80)
def test_backtest(self):
"""测试回测功能"""
logger.info("开始测试回测功能...")
# 测试回测初始化
logger.info("回测对象初始化完成")
# 测试获取交易日
test_dates = self._get_trading_dates('20241201', '20241231')
logger.info(f"测试日期范围获取完成,共 {len(test_dates)} 个交易日")
# 测试数据获取
test_data = self.data_manager.get_all_stock_data_for_backtest('20241201', '20241231')
logger.info(f"测试数据获取完成,共 {len(test_data)} 只股票有数据")
# 测试执行回测(使用简化参数)
test_params = {
"learning_window": "60日",
"profit_threshold": "3%",
"stock_count": "5只",
"volume_ratio_threshold": "1.2倍"
}
test_report = self.run(test_params, '20241201', '20241231')
logger.info("测试回测执行完成")
logger.info("回测功能测试完成")
if __name__ == "__main__":
# 配置日志
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
)
# 测试回测模块
backtester = Backtester()
backtester.test_backtest()
# report = backtester.run()
# backtester.save_report(report)
# backtester.print_report(report)