416 lines
14 KiB
Python
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) |