vcp_strategy.py

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
VCP策略实现类 - Volatility Contraction Pattern
"""

import logging
import pandas as pd
import numpy as np
from typing import Dict, Any

from src.strategies.base_strategy import BaseStrategy

# 配置日志
logger = logging.getLogger(__name__)

class VCPStrategy(BaseStrategy):
    """VCP(波动率收缩模式)策略实现"""
    
    def _get_description(self) -> str:
        """获取策略描述"""
        return "VCP(波动率收缩模式)策略:通过寻找股价波动率持续收缩的股票"
    
    def _get_default_params(self) -> Dict[str, Any]:
        """获取策略默认参数"""
        return {
            "min_contraction_period": 20,
            "max_contraction_period": 60,
            "volatility_threshold": 0.15,
            "breakout_threshold": 0.05,
            "min_price": 5,
            "max_price": 200,
            "min_volume_ratio": 1.5,
            "min_macd_signal": 0.01,
            "rsi_upper_bound": 70,
            "rsi_lower_bound": 30
        }
    
    def check_stock(self, ts_code: str, df_stock: pd.DataFrame) -> bool:
        """
        检查单只股票是否符合VCP策略条件
        
        参数:
            ts_code: 股票代码
            df_stock: 股票历史数据
            
        返回:
            bool: 是否符合条件
        """
        try:
            # 检查价格范围
            latest_close = df_stock['close'].iloc[-1]
            if latest_close < self.params['min_price'] or latest_close > self.params['max_price']:
                return False
            
            # 计算价格波动率
            df_stock['volatility'] = df_stock['close'].pct_change().rolling(20).std()
            
            # 寻找波动率收缩的时期
            contraction_periods = self._find_volatility_contraction(df_stock['volatility'])
            
            if not contraction_periods:
                return False
            
            latest_contraction = contraction_periods[-1]
            
            # 检查收缩期长度
            if (latest_contraction['length'] < self.params['min_contraction_period'] or 
                latest_contraction['length'] > self.params['max_contraction_period']):
                return False
            
            # 检查波动率收缩幅度
            if latest_contraction['contraction_ratio'] < self.params['volatility_threshold']:
                return False
            
            # 检查突破幅度
            breakout_range = self._calculate_breakout_range(df_stock, latest_contraction['end_index'])
            if breakout_range < self.params['breakout_threshold']:
                return False
            
            # 计算成交量比
            avg_volume = df_stock['vol'].rolling(window=20).mean().iloc[-1]
            latest_volume = df_stock['vol'].iloc[-1]
            volume_ratio = latest_volume / avg_volume if avg_volume != 0 else 0
            
            if volume_ratio < self.params['min_volume_ratio']:
                return False
            
            # 计算RSI指标
            rsi = self._calculate_rsi(df_stock['close'], 14).iloc[-1]
            if rsi < self.params['rsi_lower_bound'] or rsi > self.params['rsi_upper_bound']:
                return False
            
            # 计算MACD指标
            macd, _, _ = self._calculate_macd(df_stock['close'])
            if macd[-1] < self.params['min_macd_signal']:
                return False
            
            logger.debug(f"股票 {ts_code} 符合VCP策略条件")
            return True
            
        except Exception as e:
            logger.error(f"检查股票 {ts_code} 时出错: {e}")
            return False
    
    def _find_volatility_contraction(self, volatility_series):
        """寻找波动率收缩的时期"""
        contraction_periods = []
        start_idx = None
        
        for i in range(1, len(volatility_series)):
            if start_idx is None and volatility_series[i] < volatility_series[i-1]:
                start_idx = i-1
            
            if start_idx is not None and volatility_series[i] > volatility_series[i-1]:
                contraction_length = i - start_idx
                contraction_ratio = (volatility_series[start_idx] - volatility_series[i-1]) / volatility_series[start_idx]
                
                contraction_periods.append({
                    "start_index": start_idx,
                    "end_index": i-1,
                    "length": contraction_length,
                    "contraction_ratio": contraction_ratio,
                    "start_vol": volatility_series[start_idx],
                    "end_vol": volatility_series[i-1]
                })
                
                start_idx = None
        
        return contraction_periods
    
    def _calculate_breakout_range(self, df_stock, end_index):
        """计算突破幅度"""
        try:
            breakout_data = df_stock.iloc[end_index:end_index+5]
            max_price = breakout_data['high'].max()
            min_price = breakout_data['low'].min()
            breakout_range = (max_price - min_price) / min_price
            
            return breakout_range
        except Exception as e:
            logger.error(f"计算突破幅度失败: {e}")
            return 0
    
    def _calculate_rsi(self, prices, period=14):
        """计算RSI指标"""
        delta = prices.diff()
        gain = (delta.where(delta > 0, 0)).rolling(window=period).mean()
        loss = (-delta.where(delta < 0, 0)).rolling(window=period).mean()
        
        rs = gain / loss
        rsi = 100 - (100 / (1 + rs))
        
        return rsi
    
    def _calculate_macd(self, prices, fast=12, slow=26, signal_period=9):
        """计算MACD指标"""
        ema_fast = prices.ewm(span=fast, adjust=False).mean()
        ema_slow = prices.ewm(span=slow, adjust=False).mean()
        
        macd = ema_fast - ema_slow
        signal = macd.ewm(span=signal_period, adjust=False).mean()
        hist = macd - signal
        
        return macd.values, signal.values, hist.values