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)