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