Add files to chanlun_1
This commit is contained in:
@@ -0,0 +1,416 @@
|
||||
"""
|
||||
中枢模块:识别价格在某个区间内的震荡模式
|
||||
中枢定义:至少由三个连续同级别重叠的线段组成
|
||||
"""
|
||||
|
||||
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)
|
||||
Reference in New Issue
Block a user