""" 笔模块:连接相邻的顶分型和底分型形成笔 笔的定义:由一个顶分型和一个底分型连接而成,且中间不能有其他分型 """ import pandas as pd import numpy as np from typing import List, Tuple, Optional, Dict from dataclasses import dataclass import logging from .fractal import FractalPoint logger = logging.getLogger(__name__) @dataclass class StrokeElement: """笔元素数据类""" start_fractal: FractalPoint # 起始分型 end_fractal: FractalPoint # 结束分型 direction: int # 方向:1为上升笔,-1为下降笔 length: float # 笔的长度(价格差) duration: int # 持续时间(K线数量) strength: float # 笔的强度 class Stroke: """笔识别和处理类""" def __init__(self, fractals: List[FractalPoint], kline_data: pd.DataFrame): """ 初始化笔处理器 Args: fractals: 分型点列表 kline_data: K线数据 """ self.fractals = fractals self.kline_data = kline_data self.strokes = [] self.up_strokes = [] # 上升笔 self.down_strokes = [] # 下降笔 def create_stroke(self, start_fractal: FractalPoint, end_fractal: FractalPoint) -> StrokeElement: """ 创建笔元素 Args: start_fractal: 起始分型 end_fractal: 结束分型 Returns: 笔元素 """ # 计算方向 if start_fractal.fractal_type == 'bottom' and end_fractal.fractal_type == 'top': direction = 1 # 上升笔 elif start_fractal.fractal_type == 'top' and end_fractal.fractal_type == 'bottom': direction = -1 # 下降笔 else: raise ValueError("无效的分型组合") # 计算长度 length = abs(end_fractal.price - start_fractal.price) # 计算持续时间 duration = end_fractal.index - start_fractal.index # 计算强度(可以基于多个因素) strength = self._calculate_stroke_strength(start_fractal, end_fractal, length, duration) return StrokeElement( start_fractal=start_fractal, end_fractal=end_fractal, direction=direction, length=length, duration=duration, strength=strength ) def _calculate_stroke_strength(self, start_fractal: FractalPoint, end_fractal: FractalPoint, length: float, duration: int) -> float: """ 计算笔的强度 Args: start_fractal: 起始分型 end_fractal: 结束分型 length: 价格长度 duration: 时间长度 Returns: 笔的强度值 """ # 基础强度:价格变化幅度 price_strength = length # 分型强度加权 fractal_strength = (start_fractal.strength + end_fractal.strength) / 2 # 时间因子:适中的时间长度得分更高 time_factor = min(duration / 10, 1.0) if duration > 0 else 0 # 综合强度 strength = price_strength * (1 + fractal_strength * 0.1) * (1 + time_factor * 0.1) return strength def validate_stroke(self, start_fractal: FractalPoint, end_fractal: FractalPoint) -> bool: """ 验证笔的有效性 Args: start_fractal: 起始分型 end_fractal: 结束分型 Returns: 是否为有效笔 """ # 检查分型类型是否正确 valid_combinations = [ ('bottom', 'top'), # 上升笔 ('top', 'bottom') # 下降笔 ] combination = (start_fractal.fractal_type, end_fractal.fractal_type) if combination not in valid_combinations: return False # 检查时间顺序 if start_fractal.index >= end_fractal.index: return False # 检查价格关系 if start_fractal.fractal_type == 'bottom': # 上升笔:结束价格应该高于起始价格 if end_fractal.price <= start_fractal.price: return False else: # 下降笔:结束价格应该低于起始价格 if end_fractal.price >= start_fractal.price: return False return True def detect_strokes(self) -> List[StrokeElement]: """ 检测所有笔 Returns: 笔列表 """ if len(self.fractals) < 2: logger.warning("分型数量不足,无法生成笔") return [] strokes = [] # 按时间顺序排序分型 sorted_fractals = sorted(self.fractals, key=lambda x: x.index) i = 0 while i < len(sorted_fractals) - 1: start_fractal = sorted_fractals[i] # 寻找下一个有效的分型来形成笔 j = i + 1 while j < len(sorted_fractals): end_fractal = sorted_fractals[j] # 检查是否可以形成有效笔 if self.validate_stroke(start_fractal, end_fractal): stroke = self.create_stroke(start_fractal, end_fractal) strokes.append(stroke) # 分类存储 if stroke.direction == 1: self.up_strokes.append(stroke) else: self.down_strokes.append(stroke) # 从结束分型继续寻找下一笔 i = j break j += 1 else: # 没有找到有效的结束分型,跳到下一个分型 i += 1 self.strokes = strokes logger.info(f"检测到 {len(strokes)} 笔:上升笔 {len(self.up_strokes)} 个,下降笔 {len(self.down_strokes)} 个") return strokes def get_stroke_sequence(self) -> List[StrokeElement]: """ 获取连续的笔序列 Returns: 连续笔序列 """ return sorted(self.strokes, key=lambda x: x.start_fractal.index) def find_stroke_extremes(self) -> Dict[str, List[StrokeElement]]: """ 寻找笔的极值点 Returns: 包含最长、最短、最强笔的字典 """ if not self.strokes: return {} # 按长度排序 by_length = sorted(self.strokes, key=lambda x: x.length, reverse=True) # 按强度排序 by_strength = sorted(self.strokes, key=lambda x: x.strength, reverse=True) # 按持续时间排序 by_duration = sorted(self.strokes, key=lambda x: x.duration, reverse=True) return { 'longest': by_length[:5], # 最长的5笔 'strongest': by_strength[:5], # 最强的5笔 'longest_duration': by_duration[:5] # 持续时间最长的5笔 } def analyze_stroke_patterns(self) -> Dict: """ 分析笔的模式 Returns: 模式分析结果 """ if len(self.strokes) < 3: return {} sequence = self.get_stroke_sequence() # 分析连续同向笔(可能的延伸) extensions = [] i = 0 while i < len(sequence) - 1: current = sequence[i] next_stroke = sequence[i + 1] # 检查是否为同向延伸 if current.direction == next_stroke.direction: extensions.append((current, next_stroke)) i += 1 # 分析笔的趋势强度 trend_strength = self._calculate_trend_strength(sequence) # 统计平均笔长度 avg_length = np.mean([s.length for s in self.strokes]) if self.strokes else 0 avg_duration = np.mean([s.duration for s in self.strokes]) if self.strokes else 0 return { 'total_strokes': len(self.strokes), 'up_strokes': len(self.up_strokes), 'down_strokes': len(self.down_strokes), 'extensions': len(extensions), 'avg_length': avg_length, 'avg_duration': avg_duration, 'trend_strength': trend_strength } def _calculate_trend_strength(self, sequence: List[StrokeElement]) -> float: """ 计算趋势强度 Args: sequence: 笔序列 Returns: 趋势强度值 """ if len(sequence) < 2: return 0 # 计算方向变化的频率 direction_changes = 0 for i in range(1, len(sequence)): if sequence[i].direction != sequence[i-1].direction: direction_changes += 1 # 趋势强度与方向变化成反比 change_ratio = direction_changes / (len(sequence) - 1) if len(sequence) > 1 else 1 trend_strength = 1 - change_ratio return trend_strength def filter_strokes_by_strength(self, min_strength: float) -> List[StrokeElement]: """ 按强度过滤笔 Args: min_strength: 最小强度要求 Returns: 过滤后的笔列表 """ return [s for s in self.strokes if s.strength >= min_strength] def filter_strokes_by_length(self, min_length: float) -> List[StrokeElement]: """ 按长度过滤笔 Args: min_length: 最小长度要求 Returns: 过滤后的笔列表 """ return [s for s in self.strokes if s.length >= min_length] def to_dataframe(self) -> pd.DataFrame: """ 将笔转换为DataFrame Returns: 包含笔信息的DataFrame """ if not self.strokes: return pd.DataFrame() data = [] for i, stroke in enumerate(self.strokes): data.append({ 'stroke_id': i, 'start_time': stroke.start_fractal.timestamp, 'end_time': stroke.end_fractal.timestamp, 'start_price': stroke.start_fractal.price, 'end_price': stroke.end_fractal.price, 'direction': stroke.direction, 'length': stroke.length, 'duration': stroke.duration, 'strength': stroke.strength, 'start_fractal_type': stroke.start_fractal.fractal_type, 'end_fractal_type': stroke.end_fractal.fractal_type }) return pd.DataFrame(data)