mirror of
https://github.com/Arianhgh/fx-quant-research.git
synced 2026-08-16 12:18:04 +00:00
384 lines
16 KiB
Python
384 lines
16 KiB
Python
"""Convert model predictions into trading signals with risk management."""
|
|
import numpy as np
|
|
import pandas as pd
|
|
|
|
from tradingbot.models.model_manager import GoldModelManager
|
|
|
|
|
|
class GoldSignalGenerator:
|
|
"""
|
|
Generates trading signals from model predictions with risk management
|
|
"""
|
|
def __init__(
|
|
self,
|
|
confidence_threshold=0.7,
|
|
risk_reward_min=1.5,
|
|
stop_atr_factor=1.5,
|
|
target_atr_factor=2.25,
|
|
model_manager=None,
|
|
model_path='models'
|
|
):
|
|
"""
|
|
Initialize the signal generator
|
|
|
|
Parameters:
|
|
-----------
|
|
confidence_threshold : float
|
|
Minimum probability threshold for generating signals
|
|
risk_reward_min : float
|
|
Minimum risk/reward ratio for valid trades
|
|
stop_atr_factor : float
|
|
Factor to multiply ATR for stop loss calculation
|
|
target_atr_factor : float
|
|
Factor to multiply ATR for take profit calculation
|
|
model_manager : GoldModelManager, optional
|
|
Model manager instance (if None, will load from model_path)
|
|
model_path : str
|
|
Directory to load models from (if model_manager is None)
|
|
"""
|
|
self.confidence_threshold = confidence_threshold
|
|
self.risk_reward_min = risk_reward_min
|
|
self.stop_atr_factor = stop_atr_factor
|
|
self.target_atr_factor = target_atr_factor
|
|
|
|
# Use provided model manager or create a new one
|
|
if model_manager is not None:
|
|
self.model_manager = model_manager
|
|
else:
|
|
self.model_manager = GoldModelManager(model_path=model_path)
|
|
self.model_manager.load_models()
|
|
|
|
def generate_signals(self, data):
|
|
"""
|
|
Generate trading signals from data with improved error handling
|
|
|
|
Parameters:
|
|
-----------
|
|
data : pd.DataFrame
|
|
Data with features
|
|
|
|
Returns:
|
|
--------
|
|
signals : pd.DataFrame
|
|
DataFrame with trading signals and risk management
|
|
"""
|
|
try:
|
|
# Prepare features
|
|
X, _ = self.model_manager.prepare_data(data, remove_cols=None)
|
|
|
|
# Check if model manager has a trained meta model
|
|
if self.model_manager.meta_model is None:
|
|
print("Warning: No trained meta-model available for prediction")
|
|
return pd.DataFrame(index=data.index)
|
|
|
|
# Get model predictions
|
|
predictions, probabilities = self.model_manager.predict(X)
|
|
|
|
# Check if predictions or probabilities are empty
|
|
if predictions.empty or probabilities.empty:
|
|
print("Warning: Empty predictions or probabilities")
|
|
return pd.DataFrame(index=data.index)
|
|
|
|
# Initialize signals DataFrame
|
|
signals = pd.DataFrame(index=data.index)
|
|
signals['prediction'] = predictions
|
|
|
|
# Add class probabilities
|
|
for col in probabilities.columns:
|
|
signals[col] = probabilities[col]
|
|
|
|
# Calculate signal confidence with NaN handling
|
|
if probabilities.values.size > 0:
|
|
confidence_values = np.nanmax(probabilities.values, axis=1)
|
|
# Replace any NaN confidence values with 0
|
|
confidence_values = np.nan_to_num(confidence_values, nan=0)
|
|
signals['confidence'] = confidence_values
|
|
else:
|
|
signals['confidence'] = 0
|
|
|
|
# Generate directional signals
|
|
signals['signal'] = 0 # Default: no signal
|
|
|
|
# Long signals (Strong Up or Weak Up with high confidence)
|
|
long_mask = (
|
|
((signals['prediction'] == 2) & (signals['confidence'] > self.confidence_threshold * 1.1)) | # Higher threshold for strong up
|
|
((signals['prediction'] == 1) & (signals['confidence'] > self.confidence_threshold))
|
|
)
|
|
if not long_mask.empty:
|
|
signals.loc[long_mask, 'signal'] = 1
|
|
|
|
# Short signals (Strong Down or Weak Down with high confidence)
|
|
short_mask = (
|
|
((signals['prediction'] == -2) & (signals['confidence'] > self.confidence_threshold * 1.1)) | # Higher threshold for strong down
|
|
((signals['prediction'] == -1) & (signals['confidence'] > self.confidence_threshold))
|
|
)
|
|
if not short_mask.empty:
|
|
signals.loc[short_mask, 'signal'] = -1
|
|
|
|
# Add risk management
|
|
if 'atr_10' in data.columns:
|
|
# Use ATR for stop loss and take profit calculation
|
|
signals['atr'] = data['atr_10']
|
|
|
|
# Calculate stops and targets
|
|
signals['stop_distance'] = signals['atr'] * self.stop_atr_factor
|
|
signals['target_distance'] = signals['atr'] * self.target_atr_factor
|
|
|
|
# Set specific stop and target levels
|
|
signals['stop_price'] = np.where(
|
|
signals['signal'] == 1,
|
|
data['close'] - signals['stop_distance'], # Long stop
|
|
np.where(
|
|
signals['signal'] == -1,
|
|
data['close'] + signals['stop_distance'], # Short stop
|
|
np.nan
|
|
)
|
|
)
|
|
|
|
signals['target_price'] = np.where(
|
|
signals['signal'] == 1,
|
|
data['close'] + signals['target_distance'], # Long target
|
|
np.where(
|
|
signals['signal'] == -1,
|
|
data['close'] - signals['target_distance'], # Short target
|
|
np.nan
|
|
)
|
|
)
|
|
|
|
# Calculate risk-reward ratio
|
|
signals['risk_reward'] = np.where(
|
|
signals['signal'] == 1,
|
|
signals['target_distance'] / signals['stop_distance'], # Long R:R
|
|
np.where(
|
|
signals['signal'] == -1,
|
|
signals['target_distance'] / signals['stop_distance'], # Short R:R
|
|
np.nan
|
|
)
|
|
)
|
|
|
|
# Filter signals by risk-reward ratio
|
|
poor_rr_mask = (signals['signal'] != 0) & (signals['risk_reward'] < self.risk_reward_min)
|
|
if not poor_rr_mask.empty:
|
|
signals.loc[poor_rr_mask, 'signal'] = 0
|
|
|
|
# Add signal strength (1-3)
|
|
signals['signal_strength'] = 0
|
|
|
|
# Strength 3: Very high confidence predictions
|
|
strong_mask = (signals['signal'] != 0) & (signals['confidence'] > 0.85)
|
|
if not strong_mask.empty:
|
|
signals.loc[strong_mask, 'signal_strength'] = 3
|
|
|
|
# Strength 2: High confidence predictions
|
|
medium_mask = (signals['signal'] != 0) & (signals['confidence'] > 0.75) & (signals['confidence'] <= 0.85)
|
|
if not medium_mask.empty:
|
|
signals.loc[medium_mask, 'signal_strength'] = 2
|
|
|
|
# Strength 1: Moderate confidence predictions
|
|
weak_mask = (signals['signal'] != 0) & (signals['confidence'] <= 0.75)
|
|
if not weak_mask.empty:
|
|
signals.loc[weak_mask, 'signal_strength'] = 1
|
|
|
|
# Add market context
|
|
if 'volatility_regime' in data.columns:
|
|
signals['volatility_regime'] = data['volatility_regime']
|
|
|
|
# Add key price levels
|
|
signals['close'] = data['close']
|
|
|
|
# Add signal label for easier interpretation
|
|
signals['signal_label'] = 'NO_SIGNAL'
|
|
long_label_mask = signals['signal'] == 1
|
|
short_label_mask = signals['signal'] == -1
|
|
|
|
if not long_label_mask.empty:
|
|
signals.loc[long_label_mask, 'signal_label'] = 'LONG'
|
|
|
|
if not short_label_mask.empty:
|
|
signals.loc[short_label_mask, 'signal_label'] = 'SHORT'
|
|
|
|
# Count active signals
|
|
signal_count = (signals['signal'] != 0).sum()
|
|
print(f"Generated {signal_count} active signals out of {len(signals)} bars")
|
|
|
|
return signals
|
|
|
|
except Exception as e:
|
|
import traceback
|
|
print(f"Signal generation error: {str(e)}")
|
|
print(f"Traceback: {traceback.format_exc()}")
|
|
|
|
# Return empty DataFrame with same index as data
|
|
return pd.DataFrame(index=data.index)
|
|
|
|
def analyze_signals(self, signals, data):
|
|
"""
|
|
Analyze generated signals performance with improved error handling
|
|
|
|
Parameters:
|
|
-----------
|
|
signals : pd.DataFrame
|
|
DataFrame with trading signals
|
|
data : pd.DataFrame
|
|
Original data with price information
|
|
|
|
Returns:
|
|
--------
|
|
analysis : dict
|
|
Dictionary with signal statistics
|
|
"""
|
|
try:
|
|
# Ensure we have price data
|
|
if 'close' not in data.columns:
|
|
raise ValueError("Price data required for signal analysis")
|
|
|
|
# Check if signals DataFrame is empty or has no signal column
|
|
if signals.empty or 'signal' not in signals.columns:
|
|
print("Warning: Empty signals DataFrame or missing 'signal' column")
|
|
return {
|
|
'total_signals': 0,
|
|
'signal_frequency': 0,
|
|
'long_count': 0,
|
|
'short_count': 0,
|
|
'overall_win_rate': np.nan,
|
|
'overall_avg_return': np.nan
|
|
}
|
|
|
|
# Copy signals to avoid modifying the original
|
|
signals_copy = signals.copy()
|
|
|
|
# Calculate forward returns for performance assessment
|
|
for period in [1, 3, 6, 12]: # Multiple forward periods
|
|
signals_copy[f'fwd_return_{period}'] = data['close'].pct_change(period).shift(-period)
|
|
|
|
# Count signals
|
|
total_signals = (signals_copy['signal'] != 0).sum()
|
|
|
|
# If no signals were generated, return empty stats
|
|
if total_signals == 0:
|
|
print("No active signals found for analysis")
|
|
return {
|
|
'total_signals': 0,
|
|
'signal_frequency': 0,
|
|
'long_count': 0,
|
|
'short_count': 0,
|
|
'overall_win_rate': np.nan,
|
|
'overall_avg_return': np.nan
|
|
}
|
|
|
|
# Separate long and short signals
|
|
long_signals = signals_copy[signals_copy['signal'] == 1]
|
|
short_signals = signals_copy[signals_copy['signal'] == -1]
|
|
|
|
# Calculate win rates with error handling
|
|
if len(long_signals) > 0 and 'fwd_return_6' in long_signals.columns:
|
|
long_win_rate = (long_signals['fwd_return_6'] > 0).mean()
|
|
long_avg_return = long_signals['fwd_return_6'].mean()
|
|
else:
|
|
long_win_rate = np.nan
|
|
long_avg_return = np.nan
|
|
|
|
if len(short_signals) > 0 and 'fwd_return_6' in short_signals.columns:
|
|
short_win_rate = (short_signals['fwd_return_6'] < 0).mean()
|
|
short_avg_return = -short_signals['fwd_return_6'].mean()
|
|
else:
|
|
short_win_rate = np.nan
|
|
short_avg_return = np.nan
|
|
|
|
# Calculate overall metrics
|
|
win_rates = [r for r in [long_win_rate, short_win_rate] if not np.isnan(r)]
|
|
returns = [r for r in [long_avg_return, short_avg_return] if not np.isnan(r)]
|
|
|
|
overall_win_rate = np.mean(win_rates) if win_rates else np.nan
|
|
overall_avg_return = np.mean(returns) if returns else np.nan
|
|
|
|
# Signal frequency
|
|
signal_frequency = total_signals / len(signals_copy)
|
|
|
|
# Analyze by volatility regime if available
|
|
regime_stats = None
|
|
if 'volatility_regime' in signals_copy.columns:
|
|
regime_stats = {}
|
|
for regime in signals_copy['volatility_regime'].unique():
|
|
regime_signals = signals_copy[signals_copy['volatility_regime'] == regime]
|
|
|
|
# Skip if too few signals
|
|
if (regime_signals['signal'] != 0).sum() < 5:
|
|
continue
|
|
|
|
regime_long = regime_signals[regime_signals['signal'] == 1]
|
|
regime_short = regime_signals[regime_signals['signal'] == -1]
|
|
|
|
# Calculate regime metrics with error handling
|
|
if len(regime_long) > 0 and 'fwd_return_6' in regime_long.columns:
|
|
regime_long_win_rate = (regime_long['fwd_return_6'] > 0).mean()
|
|
regime_long_avg_return = regime_long['fwd_return_6'].mean()
|
|
else:
|
|
regime_long_win_rate = np.nan
|
|
regime_long_avg_return = np.nan
|
|
|
|
if len(regime_short) > 0 and 'fwd_return_6' in regime_short.columns:
|
|
regime_short_win_rate = (regime_short['fwd_return_6'] < 0).mean()
|
|
regime_short_avg_return = -regime_short['fwd_return_6'].mean()
|
|
else:
|
|
regime_short_win_rate = np.nan
|
|
regime_short_avg_return = np.nan
|
|
|
|
regime_stats[int(regime)] = {
|
|
'count': (regime_signals['signal'] != 0).sum(),
|
|
'frequency': (regime_signals['signal'] != 0).sum() / len(regime_signals),
|
|
'long_win_rate': regime_long_win_rate,
|
|
'short_win_rate': regime_short_win_rate,
|
|
'long_avg_return': regime_long_avg_return,
|
|
'short_avg_return': regime_short_avg_return
|
|
}
|
|
|
|
# Compile analysis results
|
|
analysis = {
|
|
'total_signals': total_signals,
|
|
'signal_frequency': signal_frequency,
|
|
'long_count': len(long_signals),
|
|
'short_count': len(short_signals),
|
|
'long_win_rate': long_win_rate,
|
|
'short_win_rate': short_win_rate,
|
|
'overall_win_rate': overall_win_rate,
|
|
'long_avg_return': long_avg_return,
|
|
'short_avg_return': short_avg_return,
|
|
'overall_avg_return': overall_avg_return,
|
|
'regime_stats': regime_stats
|
|
}
|
|
|
|
# Print summary
|
|
print("\nSignal Analysis:")
|
|
print(f"Total Signals: {total_signals} ({signal_frequency:.2%} of bars)")
|
|
print(f"Long Signals: {len(long_signals)}, Short Signals: {len(short_signals)}")
|
|
|
|
win_rate_str = f"{overall_win_rate:.2%}" if not np.isnan(overall_win_rate) else "N/A"
|
|
long_win_rate_str = f"{long_win_rate:.2%}" if not np.isnan(long_win_rate) else "N/A"
|
|
short_win_rate_str = f"{short_win_rate:.2%}" if not np.isnan(short_win_rate) else "N/A"
|
|
|
|
long_return_str = f"{long_avg_return:.2%}" if not np.isnan(long_avg_return) else "N/A"
|
|
short_return_str = f"{short_avg_return:.2%}" if not np.isnan(short_avg_return) else "N/A"
|
|
overall_return_str = f"{overall_avg_return:.2%}" if not np.isnan(overall_avg_return) else "N/A"
|
|
|
|
print(f"Win Rates - Long: {long_win_rate_str}, Short: {short_win_rate_str}, Overall: {win_rate_str}")
|
|
print(f"Avg Returns - Long: {long_return_str}, Short: {short_return_str}, Overall: {overall_return_str}")
|
|
|
|
return analysis
|
|
|
|
except Exception as e:
|
|
import traceback
|
|
print(f"Signal analysis error: {str(e)}")
|
|
print(f"Traceback: {traceback.format_exc()}")
|
|
|
|
# Return basic metrics
|
|
return {
|
|
'total_signals': 0,
|
|
'signal_frequency': 0,
|
|
'long_count': 0,
|
|
'short_count': 0,
|
|
'overall_win_rate': np.nan,
|
|
'overall_avg_return': np.nan,
|
|
'error': str(e)
|
|
}
|