""" 中枢模块:识别价格在某个区间内的震荡模式 中枢定义:至少由三个连续同级别重叠的线段组成 """ import pandas as pd import numpy as np from typing import List, Tuple, Optional, Dict from dataclasses import dataclass import logging from .segment import SegmentElement logger = logging.getLogger(__name__) @dataclass class CentralBankElement: """中枢元素数据类""" segments: List[SegmentElement] # 组成中枢的线段 high_price: float # 中枢上边界 low_price: float # 中枢下边界 center_price: float # 中枢中心价格 start_time: pd.Timestamp # 中枢开始时间 end_time: pd.Timestamp # 中枢结束时间 duration: int # 中枢持续时间 strength: float # 中枢强度 level: str # 中枢级别 confirmed: bool = False # 是否已确认 class CentralBank: """中枢识别和处理类""" def __init__(self, segments: List[SegmentElement]): """ 初始化中枢处理器 Args: segments: 线段列表 """ self.segments = sorted(segments, key=lambda x: x.start_time) self.central_banks = [] self.min_segments = 3 # 形成中枢的最少线段数 def find_overlapping_segments(self, segments: List[SegmentElement]) -> List[SegmentElement]: """ 寻找重叠的线段组 Args: segments: 线段列表 Returns: 重叠线段组 """ if len(segments) < self.min_segments: return [] # 找到所有线段的价格区间 overlapping = [] for i, seg1 in enumerate(segments[:-2]): overlapping_group = [seg1] for j in range(i + 1, len(segments)): seg2 = segments[j] # 检查与group中任意线段是否重叠 has_overlap = False for existing_seg in overlapping_group: if self._segments_overlap(existing_seg, seg2): has_overlap = True break if has_overlap: overlapping_group.append(seg2) else: break # 一旦不重叠就停止扩展 # 如果找到足够的重叠线段,返回这组 if len(overlapping_group) >= self.min_segments: overlapping.extend(overlapping_group) return overlapping def _segments_overlap(self, seg1: SegmentElement, seg2: SegmentElement) -> bool: """ 判断两个线段是否重叠 Args: seg1: 线段1 seg2: 线段2 Returns: 是否重叠 """ # 获取每个线段的价格区间 seg1_high = max(seg1.start_price, seg1.end_price) seg1_low = min(seg1.start_price, seg1.end_price) seg2_high = max(seg2.start_price, seg2.end_price) seg2_low = min(seg2.start_price, seg2.end_price) # 检查区间是否重叠 return not (seg1_high < seg2_low or seg2_high < seg1_low) def calculate_overlap_zone(self, segments: List[SegmentElement]) -> Tuple[float, float]: """ 计算多个线段的重叠区域 Args: segments: 线段列表 Returns: (重叠区域下边界, 重叠区域上边界) """ if not segments: return 0, 0 # 计算所有线段的价格区间 all_highs = [] all_lows = [] for seg in segments: all_highs.append(max(seg.start_price, seg.end_price)) all_lows.append(min(seg.start_price, seg.end_price)) # 重叠区域是所有高点的最小值和所有低点的最大值 overlap_high = min(all_highs) overlap_low = max(all_lows) # 确保重叠区域有效 if overlap_high > overlap_low: return overlap_low, overlap_high else: return 0, 0 def create_central_bank(self, segments: List[SegmentElement]) -> Optional[CentralBankElement]: """ 创建中枢元素 Args: segments: 组成中枢的线段 Returns: 中枢元素,如果无效则返回None """ if len(segments) < self.min_segments: return None # 计算重叠区域 low_price, high_price = self.calculate_overlap_zone(segments) if low_price >= high_price: return None # 无有效重叠区域 center_price = (high_price + low_price) / 2 start_time = min(seg.start_time for seg in segments) end_time = max(seg.end_time for seg in segments) # 计算持续时间(简化为小时数) duration = int((end_time - start_time).total_seconds() / 3600) # 计算中枢强度 strength = self._calculate_central_bank_strength(segments, high_price - low_price, duration) # 确定中枢级别 level = self._determine_central_bank_level(segments) return CentralBankElement( segments=segments, high_price=high_price, low_price=low_price, center_price=center_price, start_time=start_time, end_time=end_time, duration=duration, strength=strength, level=level ) def _calculate_central_bank_strength(self, segments: List[SegmentElement], height: float, duration: int) -> float: """ 计算中枢强度 Args: segments: 组成中枢的线段 height: 中枢高度 duration: 持续时间 Returns: 中枢强度 """ # 基础强度:线段数量和平均强度 base_strength = len(segments) * np.mean([seg.strength for seg in segments]) # 高度因子:适中的高度得分更高 height_factor = 1 / (1 + height * 0.01) # 高度越大,因子越小 # 时间因子:持续时间适中得分更高 time_factor = min(duration / 24, 2.0) # 以24小时为基准,最多2倍 return base_strength * height_factor * (1 + time_factor * 0.1) def _determine_central_bank_level(self, segments: List[SegmentElement]) -> str: """ 确定中枢级别 Args: segments: 组成中枢的线段 Returns: 中枢级别 """ # 简化的级别判断:根据线段数量和强度 avg_strength = np.mean([seg.strength for seg in segments]) segment_count = len(segments) if segment_count >= 5 and avg_strength > 100: return "1日" elif segment_count >= 4 and avg_strength > 50: return "4小时" elif segment_count >= 3 and avg_strength > 20: return "1小时" else: return "30分钟" def detect_central_banks(self) -> List[CentralBankElement]: """ 检测所有中枢 Returns: 中枢列表 """ if len(self.segments) < self.min_segments: logger.warning("线段数量不足,无法形成中枢") return [] central_banks = [] # 使用滑动窗口寻找中枢 for i in range(len(self.segments) - self.min_segments + 1): # 尝试不同长度的窗口 for window_size in range(self.min_segments, min(8, len(self.segments) - i + 1)): window_segments = self.segments[i:i + window_size] # 检查这些线段是否能形成中枢 if self._can_form_central_bank(window_segments): central_bank = self.create_central_bank(window_segments) if central_bank: # 检查是否与已有中枢重复 if not self._is_duplicate_central_bank(central_bank, central_banks): central_banks.append(central_bank) self.central_banks = central_banks logger.info(f"检测到 {len(central_banks)} 个中枢") return central_banks def _can_form_central_bank(self, segments: List[SegmentElement]) -> bool: """ 判断线段组是否能形成中枢 Args: segments: 线段组 Returns: 是否能形成中枢 """ if len(segments) < self.min_segments: return False # 检查是否有足够的重叠 overlap_count = 0 for i in range(len(segments) - 1): for j in range(i + 1, len(segments)): if self._segments_overlap(segments[i], segments[j]): overlap_count += 1 # 至少需要一半的线段对重叠 required_overlaps = len(segments) // 2 return overlap_count >= required_overlaps def _is_duplicate_central_bank(self, new_cb: CentralBankElement, existing_cbs: List[CentralBankElement]) -> bool: """ 检查是否为重复的中枢 Args: new_cb: 新中枢 existing_cbs: 已有中枢列表 Returns: 是否重复 """ for existing_cb in existing_cbs: # 检查时间和价格区间是否大量重叠 time_overlap = (min(new_cb.end_time, existing_cb.end_time) - max(new_cb.start_time, existing_cb.start_time)).total_seconds() price_overlap = (min(new_cb.high_price, existing_cb.high_price) - max(new_cb.low_price, existing_cb.low_price)) if time_overlap > 0 and price_overlap > 0: # 计算重叠比例 new_duration = (new_cb.end_time - new_cb.start_time).total_seconds() new_height = new_cb.high_price - new_cb.low_price time_overlap_ratio = time_overlap / new_duration if new_duration > 0 else 0 price_overlap_ratio = price_overlap / new_height if new_height > 0 else 0 # 如果时间和价格重叠都超过70%,认为是重复 if time_overlap_ratio > 0.7 and price_overlap_ratio > 0.7: return True return False def analyze_central_bank_patterns(self) -> Dict: """ 分析中枢模式 Returns: 模式分析结果 """ if not self.central_banks: return {} # 统计不同级别的中枢 level_counts = {} for cb in self.central_banks: level_counts[cb.level] = level_counts.get(cb.level, 0) + 1 # 计算平均指标 avg_strength = np.mean([cb.strength for cb in self.central_banks]) avg_duration = np.mean([cb.duration for cb in self.central_banks]) avg_height = np.mean([cb.high_price - cb.low_price for cb in self.central_banks]) # 寻找最强中枢 strongest_cb = max(self.central_banks, key=lambda x: x.strength) if self.central_banks else None return { 'total_central_banks': len(self.central_banks), 'level_distribution': level_counts, 'avg_strength': avg_strength, 'avg_duration': avg_duration, 'avg_height': avg_height, 'strongest_central_bank': { 'strength': strongest_cb.strength, 'level': strongest_cb.level, 'duration': strongest_cb.duration } if strongest_cb else None, 'confirmed_central_banks': sum(1 for cb in self.central_banks if cb.confirmed) } def find_central_bank_breaks(self) -> List[Dict]: """ 寻找中枢突破 Returns: 突破信息列表 """ breaks = [] for cb in self.central_banks: # 检查中枢后续价格是否突破 post_segments = [seg for seg in self.segments if seg.start_time > cb.end_time] for seg in post_segments[:3]: # 只看后续3个线段 if seg.direction == 1 and seg.end_price > cb.high_price: # 向上突破 breaks.append({ 'central_bank': cb, 'break_type': 'upward', 'break_segment': seg, 'break_strength': seg.end_price - cb.high_price }) break elif seg.direction == -1 and seg.end_price < cb.low_price: # 向下突破 breaks.append({ 'central_bank': cb, 'break_type': 'downward', 'break_segment': seg, 'break_strength': cb.low_price - seg.end_price }) break return breaks def to_dataframe(self) -> pd.DataFrame: """ 将中枢转换为DataFrame Returns: 包含中枢信息的DataFrame """ if not self.central_banks: return pd.DataFrame() data = [] for i, cb in enumerate(self.central_banks): data.append({ 'central_bank_id': i, 'start_time': cb.start_time, 'end_time': cb.end_time, 'high_price': cb.high_price, 'low_price': cb.low_price, 'center_price': cb.center_price, 'height': cb.high_price - cb.low_price, 'duration': cb.duration, 'strength': cb.strength, 'level': cb.level, 'segment_count': len(cb.segments), 'confirmed': cb.confirmed }) return pd.DataFrame(data)