576 lines
23 KiB
Python
576 lines
23 KiB
Python
"""
|
|
买卖点信号模块:识别缠论中的各类买卖点
|
|
包括一类买卖点、二类买卖点、三类买卖点的识别逻辑
|
|
"""
|
|
|
|
import pandas as pd
|
|
import numpy as np
|
|
from typing import List, Tuple, Optional, Dict
|
|
from dataclasses import dataclass
|
|
import logging
|
|
from .central_bank import CentralBankElement
|
|
from .segment import SegmentElement
|
|
from .stroke import StrokeElement
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
@dataclass
|
|
class TradingPoint:
|
|
"""买卖点数据类"""
|
|
timestamp: pd.Timestamp # 信号时间
|
|
price: float # 信号价格
|
|
signal_type: str # 'buy' 或 'sell'
|
|
point_class: str # 'first', 'second', 'third'
|
|
strength: float # 信号强度
|
|
description: str # 信号描述
|
|
related_central_bank: Optional[CentralBankElement] = None # 相关中枢
|
|
confirmed: bool = False # 是否已确认
|
|
|
|
|
|
class TradingSignal:
|
|
"""买卖点信号识别类"""
|
|
|
|
def __init__(self, central_banks: List[CentralBankElement],
|
|
segments: List[SegmentElement],
|
|
strokes: List[StrokeElement]):
|
|
"""
|
|
初始化买卖点识别器
|
|
|
|
Args:
|
|
central_banks: 中枢列表
|
|
segments: 线段列表
|
|
strokes: 笔列表
|
|
"""
|
|
self.central_banks = sorted(central_banks, key=lambda x: x.start_time)
|
|
self.segments = sorted(segments, key=lambda x: x.start_time)
|
|
self.strokes = sorted(strokes, key=lambda x: x.start_fractal.timestamp)
|
|
self.trading_points = []
|
|
|
|
def detect_first_class_points(self) -> List[TradingPoint]:
|
|
"""
|
|
识别一类买卖点:中枢突破点
|
|
|
|
Returns:
|
|
一类买卖点列表
|
|
"""
|
|
first_class_points = []
|
|
|
|
for cb in self.central_banks:
|
|
# 寻找中枢后的突破
|
|
post_segments = [seg for seg in self.segments
|
|
if seg.start_time > cb.end_time]
|
|
|
|
if not post_segments:
|
|
continue
|
|
|
|
# 检查前3个线段中的突破
|
|
for seg in post_segments[:3]:
|
|
if seg.direction == 1 and seg.end_price > cb.high_price:
|
|
# 向上突破 - 一类买点
|
|
strength = self._calculate_breakout_strength(cb, seg, 'up')
|
|
|
|
point = TradingPoint(
|
|
timestamp=seg.end_time,
|
|
price=seg.end_price,
|
|
signal_type='buy',
|
|
point_class='first',
|
|
strength=strength,
|
|
description=f"一类买点:向上突破{cb.level}中枢",
|
|
related_central_bank=cb
|
|
)
|
|
first_class_points.append(point)
|
|
break
|
|
|
|
elif seg.direction == -1 and seg.end_price < cb.low_price:
|
|
# 向下突破 - 一类卖点
|
|
strength = self._calculate_breakout_strength(cb, seg, 'down')
|
|
|
|
point = TradingPoint(
|
|
timestamp=seg.end_time,
|
|
price=seg.end_price,
|
|
signal_type='sell',
|
|
point_class='first',
|
|
strength=strength,
|
|
description=f"一类卖点:向下突破{cb.level}中枢",
|
|
related_central_bank=cb
|
|
)
|
|
first_class_points.append(point)
|
|
break
|
|
|
|
return first_class_points
|
|
|
|
def detect_second_class_points(self) -> List[TradingPoint]:
|
|
"""
|
|
识别二类买卖点:回拉不进入中枢的确认点
|
|
|
|
Returns:
|
|
二类买卖点列表
|
|
"""
|
|
second_class_points = []
|
|
|
|
# 先获取一类买卖点
|
|
first_class_points = self.detect_first_class_points()
|
|
|
|
for first_point in first_class_points:
|
|
cb = first_point.related_central_bank
|
|
if not cb:
|
|
continue
|
|
|
|
# 寻找一类点之后的回拉
|
|
post_segments = [seg for seg in self.segments
|
|
if seg.start_time > first_point.timestamp]
|
|
|
|
for i, seg in enumerate(post_segments[:5]): # 检查后续5个线段
|
|
if first_point.signal_type == 'buy':
|
|
# 一类买点后的回拉测试
|
|
if (seg.direction == -1 and
|
|
seg.end_price > cb.high_price and # 没有跌破中枢上边界
|
|
i < len(post_segments) - 1): # 确保有后续线段
|
|
|
|
next_seg = post_segments[i + 1]
|
|
if next_seg.direction == 1: # 回拉后再次上涨
|
|
strength = self._calculate_pullback_strength(cb, seg, next_seg)
|
|
|
|
point = TradingPoint(
|
|
timestamp=next_seg.start_time,
|
|
price=seg.end_price,
|
|
signal_type='buy',
|
|
point_class='second',
|
|
strength=strength,
|
|
description=f"二类买点:回拉确认{cb.level}中枢支撑",
|
|
related_central_bank=cb
|
|
)
|
|
second_class_points.append(point)
|
|
|
|
elif first_point.signal_type == 'sell':
|
|
# 一类卖点后的反弹测试
|
|
if (seg.direction == 1 and
|
|
seg.end_price < cb.low_price and # 没有涨破中枢下边界
|
|
i < len(post_segments) - 1): # 确保有后续线段
|
|
|
|
next_seg = post_segments[i + 1]
|
|
if next_seg.direction == -1: # 反弹后再次下跌
|
|
strength = self._calculate_pullback_strength(cb, seg, next_seg)
|
|
|
|
point = TradingPoint(
|
|
timestamp=next_seg.start_time,
|
|
price=seg.end_price,
|
|
signal_type='sell',
|
|
point_class='second',
|
|
strength=strength,
|
|
description=f"二类卖点:反弹确认{cb.level}中枢阻力",
|
|
related_central_bank=cb
|
|
)
|
|
second_class_points.append(point)
|
|
|
|
return second_class_points
|
|
|
|
def detect_third_class_points(self) -> List[TradingPoint]:
|
|
"""
|
|
识别三类买卖点:次级别背驰点
|
|
|
|
Returns:
|
|
三类买卖点列表
|
|
"""
|
|
third_class_points = []
|
|
|
|
# 在中枢内部寻找次级别的背驰机会
|
|
for cb in self.central_banks:
|
|
# 获取中枢时间范围内的笔
|
|
internal_strokes = [stroke for stroke in self.strokes
|
|
if (stroke.start_fractal.timestamp >= cb.start_time and
|
|
stroke.end_fractal.timestamp <= cb.end_time)]
|
|
|
|
if len(internal_strokes) < 5: # 需要足够的笔进行分析
|
|
continue
|
|
|
|
# 寻找背驰模式
|
|
divergence_points = self._find_divergence_in_strokes(internal_strokes, cb)
|
|
third_class_points.extend(divergence_points)
|
|
|
|
return third_class_points
|
|
|
|
def _find_divergence_in_strokes(self, strokes: List[StrokeElement],
|
|
cb: CentralBankElement) -> List[TradingPoint]:
|
|
"""
|
|
在笔中寻找背驰模式
|
|
|
|
Args:
|
|
strokes: 笔列表
|
|
cb: 相关中枢
|
|
|
|
Returns:
|
|
背驰点列表
|
|
"""
|
|
divergence_points = []
|
|
|
|
# 按方向分组
|
|
up_strokes = [s for s in strokes if s.direction == 1]
|
|
down_strokes = [s for s in strokes if s.direction == -1]
|
|
|
|
# 检查上升笔的顶背驰
|
|
if len(up_strokes) >= 3:
|
|
for i in range(2, len(up_strokes)):
|
|
current_stroke = up_strokes[i]
|
|
prev_stroke = up_strokes[i-1]
|
|
|
|
# 价格创新高但力度减弱
|
|
if (current_stroke.end_fractal.price > prev_stroke.end_fractal.price and
|
|
current_stroke.strength < prev_stroke.strength * 0.8):
|
|
|
|
strength = self._calculate_divergence_strength(current_stroke, prev_stroke)
|
|
|
|
point = TradingPoint(
|
|
timestamp=current_stroke.end_fractal.timestamp,
|
|
price=current_stroke.end_fractal.price,
|
|
signal_type='sell',
|
|
point_class='third',
|
|
strength=strength,
|
|
description=f"三类卖点:{cb.level}中枢内顶背驰",
|
|
related_central_bank=cb
|
|
)
|
|
divergence_points.append(point)
|
|
|
|
# 检查下降笔的底背驰
|
|
if len(down_strokes) >= 3:
|
|
for i in range(2, len(down_strokes)):
|
|
current_stroke = down_strokes[i]
|
|
prev_stroke = down_strokes[i-1]
|
|
|
|
# 价格创新低但力度减弱
|
|
if (current_stroke.end_fractal.price < prev_stroke.end_fractal.price and
|
|
current_stroke.strength < prev_stroke.strength * 0.8):
|
|
|
|
strength = self._calculate_divergence_strength(current_stroke, prev_stroke)
|
|
|
|
point = TradingPoint(
|
|
timestamp=current_stroke.end_fractal.timestamp,
|
|
price=current_stroke.end_fractal.price,
|
|
signal_type='buy',
|
|
point_class='third',
|
|
strength=strength,
|
|
description=f"三类买点:{cb.level}中枢内底背驰",
|
|
related_central_bank=cb
|
|
)
|
|
divergence_points.append(point)
|
|
|
|
return divergence_points
|
|
|
|
def _calculate_breakout_strength(self, cb: CentralBankElement,
|
|
segment: SegmentElement, direction: str) -> float:
|
|
"""
|
|
计算突破强度
|
|
|
|
Args:
|
|
cb: 中枢
|
|
segment: 突破线段
|
|
direction: 突破方向
|
|
|
|
Returns:
|
|
突破强度
|
|
"""
|
|
# 基础强度:突破幅度
|
|
if direction == 'up':
|
|
breakout_distance = segment.end_price - cb.high_price
|
|
base_strength = breakout_distance / cb.high_price
|
|
else:
|
|
breakout_distance = cb.low_price - segment.end_price
|
|
base_strength = breakout_distance / cb.low_price
|
|
|
|
# 中枢强度加权
|
|
cb_strength_factor = min(cb.strength / 100, 2.0)
|
|
|
|
# 线段强度加权
|
|
segment_strength_factor = min(segment.strength / 50, 2.0)
|
|
|
|
# 综合强度
|
|
total_strength = base_strength * (1 + cb_strength_factor * 0.2) * (1 + segment_strength_factor * 0.3)
|
|
|
|
return max(0, min(total_strength, 1.0)) # 限制在0-1之间
|
|
|
|
def _calculate_pullback_strength(self, cb: CentralBankElement,
|
|
pullback_seg: SegmentElement,
|
|
resume_seg: SegmentElement) -> float:
|
|
"""
|
|
计算回拉强度
|
|
|
|
Args:
|
|
cb: 中枢
|
|
pullback_seg: 回拉线段
|
|
resume_seg: 恢复线段
|
|
|
|
Returns:
|
|
回拉强度
|
|
"""
|
|
# 回拉深度(越浅越好)
|
|
cb_height = cb.high_price - cb.low_price
|
|
if pullback_seg.direction == -1: # 向下回拉
|
|
pullback_depth = (cb.high_price - pullback_seg.end_price) / cb_height
|
|
else: # 向上回拉
|
|
pullback_depth = (pullback_seg.end_price - cb.low_price) / cb_height
|
|
|
|
# 回拉强度:深度越小越好
|
|
pullback_strength = max(0, 1 - pullback_depth)
|
|
|
|
# 恢复强度
|
|
resume_strength = min(resume_seg.strength / 30, 1.0)
|
|
|
|
# 综合强度
|
|
total_strength = (pullback_strength + resume_strength) / 2
|
|
|
|
return max(0, min(total_strength, 1.0))
|
|
|
|
def _calculate_divergence_strength(self, current_stroke: StrokeElement,
|
|
prev_stroke: StrokeElement) -> float:
|
|
"""
|
|
计算背驰强度
|
|
|
|
Args:
|
|
current_stroke: 当前笔
|
|
prev_stroke: 前一笔
|
|
|
|
Returns:
|
|
背驰强度
|
|
"""
|
|
# 力度差异
|
|
strength_ratio = prev_stroke.strength / current_stroke.strength if current_stroke.strength > 0 else 1
|
|
|
|
# 价格差异
|
|
price_change = abs(current_stroke.end_fractal.price - prev_stroke.end_fractal.price)
|
|
price_ratio = price_change / prev_stroke.start_fractal.price if prev_stroke.start_fractal.price > 0 else 0
|
|
|
|
# 背驰强度
|
|
divergence_strength = min(strength_ratio * 0.3 + price_ratio * 10, 1.0)
|
|
|
|
return max(0, divergence_strength)
|
|
|
|
def detect_fractal_based_signals(self, fractals: List, processed_klines) -> List[TradingPoint]:
|
|
"""
|
|
基于分型识别买卖点(不依赖中枢)
|
|
|
|
Args:
|
|
fractals: 分型列表
|
|
processed_klines: 处理后的K线数据
|
|
|
|
Returns:
|
|
分型买卖点列表
|
|
"""
|
|
signals = []
|
|
|
|
if not fractals or len(fractals) < 3:
|
|
return signals
|
|
|
|
for i, fractal in enumerate(fractals):
|
|
# 分析分型强度和位置
|
|
strength_score = self._calculate_fractal_strength_score(fractal, fractals, i)
|
|
position_score = self._calculate_fractal_position_score(fractal, processed_klines)
|
|
|
|
# 综合评分
|
|
total_score = (strength_score + position_score) / 2
|
|
|
|
if total_score > 0.6: # 阈值可调整
|
|
signal_type = 'buy' if fractal.fractal_type == 'bottom' else 'sell'
|
|
point_class = self._determine_fractal_point_class(total_score)
|
|
|
|
description = f"分型{signal_type}点:强度{fractal.strength},评分{total_score:.2f}"
|
|
|
|
signal = TradingPoint(
|
|
timestamp=fractal.timestamp,
|
|
price=fractal.price,
|
|
signal_type=signal_type,
|
|
point_class=point_class,
|
|
strength=total_score,
|
|
description=description,
|
|
confirmed=fractal.confirmed
|
|
)
|
|
signals.append(signal)
|
|
|
|
return signals
|
|
|
|
def _calculate_fractal_strength_score(self, fractal, all_fractals: List, index: int) -> float:
|
|
"""计算分型强度评分(避免使用未来数据)"""
|
|
base_score = min(fractal.strength / 5.0, 1.0) # 基础强度评分
|
|
|
|
# 只使用过去的分型进行比较,避免使用未来数据
|
|
past_fractals = []
|
|
for i in range(max(0, index-5), index): # 只看过去5个分型
|
|
if i < len(all_fractals) and all_fractals[i].fractal_type == fractal.fractal_type:
|
|
past_fractals.append(all_fractals[i])
|
|
|
|
if past_fractals:
|
|
relative_strength = fractal.strength / max(f.strength for f in past_fractals)
|
|
relative_score = min(relative_strength, 1.0)
|
|
else:
|
|
relative_score = 1.0
|
|
|
|
return (base_score + relative_score) / 2
|
|
|
|
def _calculate_fractal_position_score(self, fractal, processed_klines) -> float:
|
|
"""计算分型位置评分(避免使用未来数据)"""
|
|
if processed_klines is None or len(processed_klines) < 20:
|
|
return 0.5
|
|
|
|
# 寻找分型在K线数据中的位置
|
|
fractal_index = None
|
|
for i, (timestamp, kline) in enumerate(processed_klines.iterrows()):
|
|
# 修复时间戳计算兼容性问题
|
|
time_diff = abs((pd.Timestamp(timestamp) - pd.Timestamp(fractal.timestamp)).total_seconds())
|
|
if time_diff < 3600: # 1小时内
|
|
fractal_index = i
|
|
break
|
|
|
|
if fractal_index is None:
|
|
return 0.5
|
|
|
|
# 计算相对位置(只使用过去的数据)
|
|
start_idx = max(0, fractal_index - 20) # 只看过去20个K线
|
|
end_idx = fractal_index + 1 # 包含当前K线,但不包含未来K线
|
|
recent_data = processed_klines.iloc[start_idx:end_idx]
|
|
|
|
if fractal.fractal_type == 'bottom':
|
|
# 底分型:看是否接近过去一段时间的低点
|
|
min_price = recent_data['low'].min()
|
|
score = 1.0 if abs(fractal.price - min_price) / min_price < 0.02 else 0.5
|
|
else:
|
|
# 顶分型:看是否接近过去一段时间的高点
|
|
max_price = recent_data['high'].max()
|
|
score = 1.0 if abs(fractal.price - max_price) / max_price < 0.02 else 0.5
|
|
|
|
return score
|
|
|
|
def _determine_fractal_point_class(self, score: float) -> str:
|
|
"""根据评分确定买卖点类别"""
|
|
if score >= 0.8:
|
|
return "second" # 高质量分型当作二类买卖点
|
|
else:
|
|
return "third" # 一般分型当作三类买卖点
|
|
|
|
def detect_all_trading_points(self) -> List[TradingPoint]:
|
|
"""
|
|
检测所有买卖点(包括中枢相关和分型相关)
|
|
"""
|
|
all_signals = []
|
|
|
|
# 原有的中枢相关买卖点
|
|
all_signals.extend(self.detect_first_class_points())
|
|
all_signals.extend(self.detect_second_class_points())
|
|
all_signals.extend(self.detect_third_class_points())
|
|
|
|
# 新增:基于分型的买卖点(当没有足够中枢时)
|
|
if len(self.central_banks) < 2:
|
|
logger.info("中枢数量不足,启用分型买卖点识别")
|
|
# 需要获取分型和K线数据,这里需要从外部传入
|
|
# 暂时先返回现有信号
|
|
pass
|
|
|
|
# 按时间排序
|
|
all_signals.sort(key=lambda x: x.timestamp)
|
|
|
|
# 去重(同一时间点可能有多个信号)
|
|
unique_signals = []
|
|
seen_times = set()
|
|
for signal in all_signals:
|
|
time_key = signal.timestamp.strftime('%Y%m%d%H%M')
|
|
if time_key not in seen_times:
|
|
unique_signals.append(signal)
|
|
seen_times.add(time_key)
|
|
|
|
class_counts = {'first': 0, 'second': 0, 'third': 0}
|
|
for signal in unique_signals:
|
|
class_counts[signal.point_class] += 1
|
|
|
|
total_signals = len(unique_signals)
|
|
logger.info(f"检测到买卖点:一类 {class_counts['first']} 个,二类 {class_counts['second']} 个,"
|
|
f"三类 {class_counts['third']} 个,总计 {total_signals} 个")
|
|
|
|
return unique_signals
|
|
|
|
def get_signal_statistics(self) -> Dict:
|
|
"""
|
|
获取信号统计信息
|
|
|
|
Returns:
|
|
统计信息字典
|
|
"""
|
|
if not self.trading_points:
|
|
return {}
|
|
|
|
stats = {
|
|
'total_signals': len(self.trading_points),
|
|
'buy_signals': len([p for p in self.trading_points if p.signal_type == 'buy']),
|
|
'sell_signals': len([p for p in self.trading_points if p.signal_type == 'sell']),
|
|
'by_class': {},
|
|
'avg_strength': np.mean([p.strength for p in self.trading_points]),
|
|
'max_strength': max([p.strength for p in self.trading_points]),
|
|
'confirmed_signals': len([p for p in self.trading_points if p.confirmed])
|
|
}
|
|
|
|
# 按类别统计
|
|
for point_class in ['first', 'second', 'third']:
|
|
class_points = [p for p in self.trading_points if p.point_class == point_class]
|
|
stats['by_class'][f'{point_class}_class'] = {
|
|
'total': len(class_points),
|
|
'buy': len([p for p in class_points if p.signal_type == 'buy']),
|
|
'sell': len([p for p in class_points if p.signal_type == 'sell']),
|
|
'avg_strength': np.mean([p.strength for p in class_points]) if class_points else 0
|
|
}
|
|
|
|
return stats
|
|
|
|
def filter_signals_by_strength(self, min_strength: float) -> List[TradingPoint]:
|
|
"""
|
|
按强度过滤信号
|
|
|
|
Args:
|
|
min_strength: 最小强度要求
|
|
|
|
Returns:
|
|
过滤后的信号列表
|
|
"""
|
|
return [p for p in self.trading_points if p.strength >= min_strength]
|
|
|
|
def get_latest_signals(self, hours: int = 24) -> List[TradingPoint]:
|
|
"""
|
|
获取最新的信号
|
|
|
|
Args:
|
|
hours: 最近多少小时
|
|
|
|
Returns:
|
|
最新信号列表
|
|
"""
|
|
if not self.trading_points:
|
|
return []
|
|
|
|
latest_time = max([p.timestamp for p in self.trading_points])
|
|
cutoff_time = latest_time - pd.Timedelta(hours=hours)
|
|
|
|
return [p for p in self.trading_points if p.timestamp >= cutoff_time]
|
|
|
|
def to_dataframe(self) -> pd.DataFrame:
|
|
"""
|
|
将买卖点转换为DataFrame
|
|
|
|
Returns:
|
|
包含买卖点信息的DataFrame
|
|
"""
|
|
if not self.trading_points:
|
|
return pd.DataFrame()
|
|
|
|
data = []
|
|
for point in self.trading_points:
|
|
data.append({
|
|
'timestamp': point.timestamp,
|
|
'price': point.price,
|
|
'signal_type': point.signal_type,
|
|
'point_class': point.point_class,
|
|
'strength': point.strength,
|
|
'description': point.description,
|
|
'confirmed': point.confirmed,
|
|
'related_central_bank_level': point.related_central_bank.level if point.related_central_bank else None
|
|
})
|
|
|
|
df = pd.DataFrame(data)
|
|
df.set_index('timestamp', inplace=True)
|
|
return df |