Files
2025-05-23 19:09:55 +08:00

416 lines
14 KiB
Python

"""
中枢模块:识别价格在某个区间内的震荡模式
中枢定义:至少由三个连续同级别重叠的线段组成
"""
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)