Files
2026-06-17 17:12:19 -04:00

1009 lines
36 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Comprehensive signal / performance visualization."""
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import matplotlib.dates as mdates
import matplotlib.ticker as mtick
import seaborn as sns
from matplotlib.gridspec import GridSpec
def visualize_signals(data, results, returns, metrics, save_path=None):
"""
Create comprehensive visualizations of trading signals and performance
Parameters:
data: Original price data DataFrame
results: Signal results DataFrame from predictor.predict()
returns: Strategy returns Series
metrics: Performance metrics dictionary
save_path: Optional path to save the plots
"""
# Set plotting style
plt.style.use('seaborn-v0_8-darkgrid')
sns.set_palette('Set1')
# Prepare data by aligning timeframes
aligned_data = data.loc[results.index].copy()
# Create a Figure with multiple subplots
fig = plt.figure(figsize=(20, 16))
gs = GridSpec(4, 4, figure=fig)
# 1. Main price chart with signals
ax_price = fig.add_subplot(gs[0:2, 0:3])
_plot_price_with_signals(ax_price, aligned_data, results)
# 2. Equity curve
ax_equity = fig.add_subplot(gs[2:3, 0:3])
_plot_equity_curve(ax_equity, returns)
# 3. Signal distribution
ax_signal_dist = fig.add_subplot(gs[0, 3])
_plot_signal_distribution(ax_signal_dist, results)
# 4. Signal strength heatmap
ax_signal_heatmap = fig.add_subplot(gs[1, 3])
_plot_signal_strength_heatmap(ax_signal_heatmap, results)
# 5. Market regime analysis
ax_regime = fig.add_subplot(gs[2, 3])
_plot_market_regime(ax_regime, results)
# 6. Performance metrics
ax_metrics = fig.add_subplot(gs[3, 3])
_plot_performance_metrics(ax_metrics, metrics)
# 7. Signal frequency over time
ax_frequency = fig.add_subplot(gs[3, 0:3])
_plot_signal_frequency(ax_frequency, results)
# Set the layout tight
fig.tight_layout()
fig.suptitle('Gold Price Trading Signal Analysis', fontsize=16, y=1.02)
# Save if path is provided
if save_path:
plt.savefig(save_path, bbox_inches='tight', dpi=300)
plt.show()
# Create a second figure for detailed analysis
fig2 = plt.figure(figsize=(20, 12))
gs2 = GridSpec(2, 3, figure=fig2)
# 1. Win/Loss by hour
ax_hour = fig2.add_subplot(gs2[0, 0])
_plot_win_loss_by_hour(ax_hour, results, returns)
# 2. Win/Loss by day of week
ax_day = fig2.add_subplot(gs2[0, 1])
_plot_win_loss_by_day(ax_day, results, returns)
# 3. Win/Loss by regime
ax_regime_perf = fig2.add_subplot(gs2[0, 2])
_plot_win_loss_by_regime(ax_regime_perf, results, returns)
# 4. Signal duration histogram
ax_duration = fig2.add_subplot(gs2[1, 0])
_plot_signal_duration(ax_duration, results)
# 5. Return distribution
ax_return_dist = fig2.add_subplot(gs2[1, 1])
_plot_return_distribution(ax_return_dist, returns, results)
# 6. Signal consistency
ax_consistency = fig2.add_subplot(gs2[1, 2])
_plot_signal_consistency(ax_consistency, results)
fig2.tight_layout()
fig2.suptitle('Detailed Signal Analysis', fontsize=16, y=1.02)
# Save if path is provided
if save_path:
detail_path = save_path.replace('.png', '_detail.png')
plt.savefig(detail_path, bbox_inches='tight', dpi=300)
plt.show()
# Create a third figure for model attribution analysis
fig3 = plt.figure(figsize=(20, 10))
gs3 = GridSpec(2, 2, figure=fig3)
# 1. Model agreement analysis
ax_agreement = fig3.add_subplot(gs3[0, 0])
_plot_model_agreement(ax_agreement, results, returns)
# 2. Signal probability analysis
ax_proba = fig3.add_subplot(gs3[0, 1])
_plot_signal_probability(ax_proba, results, returns)
# 3. Signal direction by strength
ax_strength = fig3.add_subplot(gs3[1, 0])
_plot_signal_strength_performance(ax_strength, results, returns)
# 4. Drawdown analysis
ax_drawdown = fig3.add_subplot(gs3[1, 1])
_plot_drawdown_analysis(ax_drawdown, returns)
fig3.tight_layout()
fig3.suptitle('Model Behavior Analysis', fontsize=16, y=1.02)
# Save if path is provided
if save_path:
model_path = save_path.replace('.png', '_model.png')
plt.savefig(model_path, bbox_inches='tight', dpi=300)
plt.show()
def _plot_price_with_signals(ax, data, results):
"""Plot price chart with buy/sell signals overlay"""
# Plot price
ax.plot(data.index, data['close'], color='#333333', linewidth=1, alpha=0.7, label='Price')
# Highlight buy/sell signals
buy_signals = results[results['signal'] == 1].index
sell_signals = results[results['signal'] == -1].index
# Get price values for the signals
buy_prices = data.loc[buy_signals, 'close']
sell_prices = data.loc[sell_signals, 'close']
# Plot signals with varying sizes based on strength
buy_sizes = results.loc[buy_signals, 'strength'].clip(lower=20, upper=100) / 2
sell_sizes = results.loc[sell_signals, 'strength'].clip(lower=20, upper=100) / 2
ax.scatter(buy_signals, buy_prices, color='green', s=buy_sizes, alpha=0.7, marker='^', label='Buy Signal')
ax.scatter(sell_signals, sell_prices, color='red', s=sell_sizes, alpha=0.7, marker='v', label='Sell Signal')
# Format x-axis for dates
ax.xaxis.set_major_formatter(mdates.DateFormatter('%Y-%m-%d'))
ax.xaxis.set_major_locator(mdates.WeekdayLocator(interval=2))
plt.setp(ax.xaxis.get_majorticklabels(), rotation=45, ha='right')
# Add labels and legend
ax.set_title('Gold Price with Trading Signals', fontsize=14)
ax.set_ylabel('Price', fontsize=12)
ax.legend(loc='best')
# Draw grid
ax.grid(True, alpha=0.3)
# Annotate some significant signals
top_buy = results[results['signal'] == 1].nlargest(3, 'strength')
top_sell = results[results['signal'] == -1].nlargest(3, 'strength')
for idx, row in pd.concat([top_buy, top_sell]).iterrows():
price = data.loc[idx, 'close']
strength = row['strength']
if row['signal'] == 1:
ax.annotate(f"{strength:.0f}%", (idx, price),
xytext=(0, 15), textcoords='offset points',
ha='center', va='bottom', fontsize=9,
arrowprops=dict(arrowstyle='->', color='green', alpha=0.7))
else:
ax.annotate(f"{strength:.0f}%", (idx, price),
xytext=(0, -15), textcoords='offset points',
ha='center', va='top', fontsize=9,
arrowprops=dict(arrowstyle='->', color='red', alpha=0.7))
def _plot_equity_curve(ax, returns):
"""Plot equity curve from strategy returns"""
# Calculate cumulative returns
cumulative_returns = (1 + returns).cumprod() - 1
# Plot the equity curve
ax.plot(cumulative_returns.index, cumulative_returns * 100, linewidth=2, color='#1f77b4')
# Draw the zero line
ax.axhline(y=0, color='black', linestyle='-', alpha=0.3)
# Format y-axis as percentage
ax.yaxis.set_major_formatter(mtick.PercentFormatter())
# Highlight drawdowns
underwater = cumulative_returns - cumulative_returns.cummax()
ax.fill_between(underwater.index, 0, underwater * 100, color='red', alpha=0.3)
# Add labels
ax.set_title('Strategy Equity Curve', fontsize=14)
ax.set_ylabel('Cumulative Return (%)', fontsize=12)
# Calculate and annotate key metrics directly on the chart
final_return = cumulative_returns.iloc[-1] * 100
max_drawdown = underwater.min() * 100
# Annotate final return
ax.annotate(f'Final Return: {final_return:.2f}%',
xy=(0.02, 0.85), xycoords='axes fraction',
bbox=dict(boxstyle="round,pad=0.3", fc="white", ec="gray", alpha=0.8))
# Annotate max drawdown
ax.annotate(f'Max Drawdown: {max_drawdown:.2f}%',
xy=(0.02, 0.7), xycoords='axes fraction',
bbox=dict(boxstyle="round,pad=0.3", fc="white", ec="gray", alpha=0.8))
def _plot_signal_distribution(ax, results):
"""Plot distribution of signal types"""
# Count signal types
signal_counts = results['signal'].value_counts()
# Create labels
labels = ['Buy (Long)', 'Neutral', 'Sell (Short)']
# Ensure we have all three categories (even if count is zero)
values = [signal_counts.get(1, 0), signal_counts.get(0, 0), signal_counts.get(-1, 0)]
# Calculate percentages
total = sum(values)
percentages = [v/total*100 for v in values]
# Custom color map
colors = ['green', 'gray', 'red']
# Create bar plot
bars = ax.bar(labels, values, color=colors, alpha=0.7)
# Add percentage labels on top of each bar
for bar, percentage in zip(bars, percentages):
height = bar.get_height()
ax.text(bar.get_x() + bar.get_width()/2., height + 0.1,
f'{percentage:.1f}%', ha='center', va='bottom', fontsize=9)
# Add title and labels
ax.set_title('Signal Distribution', fontsize=14)
ax.set_ylabel('Count', fontsize=12)
# Rotate x-labels for better readability
plt.setp(ax.get_xticklabels(), rotation=30, ha='right')
def _plot_signal_strength_heatmap(ax, results):
"""Plot heatmap of signal strength by direction"""
# Filter to get only actual signals
signals = results[results['signal'] != 0].copy()
# Create strength bins
signals['strength_bin'] = pd.cut(signals['strength'],
bins=[0, 20, 40, 60, 80, 100],
labels=['0-20', '20-40', '40-60', '60-80', '80-100'])
# Create direction labels
signals['direction'] = signals['signal'].map({1: 'Buy', -1: 'Sell'})
# Create count matrix
heatmap_data = pd.crosstab(signals['direction'], signals['strength_bin'])
# Plot heatmap
sns.heatmap(heatmap_data, annot=True, fmt='d', cmap='YlGnBu', ax=ax)
# Add title
ax.set_title('Signal Strength Distribution', fontsize=14)
ax.set_xlabel('Strength (%)', fontsize=12)
ax.set_ylabel('Signal Direction', fontsize=12)
def _plot_market_regime(ax, results):
"""Plot market regime distribution and signals per regime"""
# Map regime numbers to descriptive names
regime_map = {0: 'Low Vol', 1: 'Normal', 2: 'High Vol'}
# Create a copy with regime names
regime_data = results.copy()
regime_data['regime_name'] = regime_data['market_regime'].map(regime_map)
# Group by regime and count signals
regime_signals = pd.crosstab(regime_data['regime_name'], regime_data['signal'])
# Rename columns
regime_signals.columns = ['Neutral', 'Buy', 'Sell']
# Reorder columns
regime_signals = regime_signals[['Buy', 'Neutral', 'Sell']]
# Plot stacked bar chart
regime_signals.plot(kind='bar', stacked=True, color=['green', 'gray', 'red'],
alpha=0.7, ax=ax)
# Add title and labels
ax.set_title('Signals by Market Regime', fontsize=14)
ax.set_xlabel('Market Regime', fontsize=12)
ax.set_ylabel('Count', fontsize=12)
# Add total percentage annotation
for i, regime in enumerate(regime_signals.index):
total = regime_signals.iloc[i].sum()
percentage = total / len(results) * 100
ax.text(i, total + 5, f'{percentage:.1f}%', ha='center')
# Adjust legend
ax.legend(title='Signal Type')
def _plot_performance_metrics(ax, metrics):
"""Plot key performance metrics"""
# Remove axes
ax.axis('off')
# Create text content
metrics_text = (
f"Performance Metrics\n"
f"-------------------\n"
f"Total Return: {metrics['total_return']:.2%}\n"
f"Sharpe Ratio: {metrics['sharpe_ratio']:.2f}\n"
f"Win Rate: {metrics['win_rate']:.2%}\n"
f"Max Drawdown: {metrics['max_drawdown']:.2%}\n"
f"Signal Count: {metrics['signal_count']}\n"
f"Avg Signals/Day: {metrics['avg_signals_per_day']:.1f}"
)
# Add text box
ax.text(0.5, 0.5, metrics_text,
ha='center', va='center',
bbox=dict(boxstyle='round', facecolor='white', alpha=0.8),
fontsize=12, family='monospace')
ax.set_title('Performance Summary', fontsize=14)
def _plot_signal_frequency(ax, results):
"""Plot signal frequency over time"""
# Create a resampled view of signals per day
daily_signals = results['signal'].resample('D').apply(lambda x: (x != 0).sum())
# Plot as bar chart
ax.bar(daily_signals.index, daily_signals, alpha=0.7, color='#1f77b4')
# Add a trend line
z = np.polyfit(range(len(daily_signals)), daily_signals, 1)
p = np.poly1d(z)
ax.plot(daily_signals.index, p(range(len(daily_signals))),
linestyle='--', color='red', linewidth=2,
label=f'Trend: {"+" if z[0]>0 else ""}{z[0]:.4f}x + {z[1]:.1f}')
# Format x-axis for dates
ax.xaxis.set_major_formatter(mdates.DateFormatter('%Y-%m-%d'))
ax.xaxis.set_major_locator(mdates.WeekdayLocator(interval=7))
plt.setp(ax.xaxis.get_majorticklabels(), rotation=45, ha='right')
# Add labels and legend
ax.set_title('Signal Frequency Over Time', fontsize=14)
ax.set_ylabel('Number of Signals per Day', fontsize=12)
ax.legend()
# Calculate and display average signals per day
avg_signals = daily_signals.mean()
ax.axhline(y=avg_signals, color='gray', linestyle='--', alpha=0.7)
ax.text(daily_signals.index[10], avg_signals + 0.3,
f'Avg: {avg_signals:.2f} signals/day', fontsize=10)
def _plot_win_loss_by_hour(ax, results, returns):
"""Plot win/loss ratio by hour of day"""
# Combine signals and returns
performance = results.copy()
performance['return'] = returns
# Group by hour
hourly_perf = performance[performance['signal'] != 0].groupby(performance.index.hour)
# Calculate win rate and average return per hour
win_rates = hourly_perf['return'].apply(lambda x: (x > 0).mean())
avg_returns = hourly_perf['return'].mean()
# Create DataFrame for plotting
hourly_data = pd.DataFrame({
'Win Rate': win_rates,
'Avg Return': avg_returns
})
# Set up primary axis for win rate
hourly_data['Win Rate'].plot(kind='bar', color='skyblue', ax=ax, alpha=0.7)
ax.set_xlabel('Hour of Day', fontsize=12)
ax.set_ylabel('Win Rate', fontsize=12)
ax.set_ylim(0, 1)
# Set up secondary axis for average return
ax2 = ax.twinx()
hourly_data['Avg Return'].plot(kind='line', color='red', marker='o', ax=ax2)
ax2.set_ylabel('Average Return', fontsize=12, color='red')
ax2.tick_params(axis='y', colors='red')
# Add horizontal line at 0.5 for win rate
ax.axhline(y=0.5, color='gray', linestyle='--', alpha=0.5)
# Add horizontal line at 0 for average return
ax2.axhline(y=0, color='red', linestyle='--', alpha=0.5)
# Add title
ax.set_title('Performance by Hour of Day', fontsize=14)
# Add custom legend
from matplotlib.lines import Line2D
legend_elements = [
Line2D([0], [0], color='skyblue', lw=0, marker='s', markersize=10, label='Win Rate'),
Line2D([0], [0], color='red', marker='o', markersize=6, label='Avg Return')
]
ax.legend(handles=legend_elements, loc='upper right')
def _plot_win_loss_by_day(ax, results, returns):
"""Plot win/loss ratio by day of week"""
# Combine signals and returns
performance = results.copy()
performance['return'] = returns
# Convert day numbers to names
day_names = ['Monday', 'Tuesday', 'Wednesday', 'Thursday', 'Friday', 'Saturday', 'Sunday']
performance['day_name'] = [day_names[d] for d in performance.index.dayofweek]
# Group by day
daily_perf = performance[performance['signal'] != 0].groupby('day_name')
# Calculate win rate and average return per day
win_rates = daily_perf['return'].apply(lambda x: (x > 0).mean())
avg_returns = daily_perf['return'].mean()
counts = daily_perf.size()
# Reindex to ensure correct order
win_rates = win_rates.reindex(day_names)
avg_returns = avg_returns.reindex(day_names)
counts = counts.reindex(day_names)
# Create bar chart
bars = ax.bar(win_rates.index, win_rates, color='lightgreen', alpha=0.7)
# Add count annotations
for i, (bar, count) in enumerate(zip(bars, counts)):
height = bar.get_height()
ax.text(bar.get_x() + bar.get_width()/2., height + 0.02,
f'n={count}', ha='center', va='bottom', fontsize=9)
# Set up secondary axis for average return
ax2 = ax.twinx()
ax2.plot(avg_returns.index, avg_returns, color='purple', marker='d')
ax2.set_ylabel('Average Return', fontsize=12, color='purple')
# Add horizontal line at 0.5 for win rate
ax.axhline(y=0.5, color='gray', linestyle='--', alpha=0.5)
# Add horizontal line at 0 for average return
ax2.axhline(y=0, color='purple', linestyle='--', alpha=0.5)
# Add labels and title
ax.set_xlabel('Day of Week', fontsize=12)
ax.set_ylabel('Win Rate', fontsize=12)
ax.set_title('Performance by Day of Week', fontsize=14)
ax.set_ylim(0, 1)
# Rotate x-labels for better readability
plt.setp(ax.get_xticklabels(), rotation=30, ha='right')
# Add custom legend
from matplotlib.lines import Line2D
legend_elements = [
Line2D([0], [0], color='lightgreen', lw=0, marker='s', markersize=10, label='Win Rate'),
Line2D([0], [0], color='purple', marker='d', markersize=6, label='Avg Return')
]
ax.legend(handles=legend_elements, loc='upper right')
def _plot_win_loss_by_regime(ax, results, returns):
"""Plot win/loss by market regime"""
# Combine signals and returns
performance = results.copy()
performance['return'] = returns
# Map regime numbers to descriptive names
regime_map = {0: 'Low Vol', 1: 'Normal', 2: 'High Vol'}
performance['regime_name'] = performance['market_regime'].map(regime_map)
# Group by regime
regime_perf = performance[performance['signal'] != 0].groupby('regime_name')
# Calculate metrics
win_rates = regime_perf['return'].apply(lambda x: (x > 0).mean())
avg_returns = regime_perf['return'].mean()
sharpe_ratios = regime_perf['return'].apply(lambda x: x.mean() / x.std() if x.std() > 0 else 0)
counts = regime_perf.size()
# Create index for the bars
x = np.arange(len(win_rates))
width = 0.25
# Create grouped bar chart
ax.bar(x - width, win_rates, width, label='Win Rate', color='green', alpha=0.7)
ax.bar(x, avg_returns * 10, width, label='Avg Ret (×10)', color='blue', alpha=0.7)
ax.bar(x + width, sharpe_ratios, width, label='Sharpe', color='orange', alpha=0.7)
# Add count annotations
for i, count in enumerate(counts):
ax.text(i, 0.05, f'n={count}', ha='center', va='bottom', fontsize=9)
# Set x-tick labels
ax.set_xticks(x)
ax.set_xticklabels(win_rates.index)
# Add labels and title
ax.set_xlabel('Market Regime', fontsize=12)
ax.set_ylabel('Metric Value', fontsize=12)
ax.set_title('Performance by Market Regime', fontsize=14)
# Add horizontal line at 0.5 for reference
ax.axhline(y=0.5, color='gray', linestyle='--', alpha=0.5)
# Add legend
ax.legend()
def _plot_signal_duration(ax, results):
"""Plot histogram of signal duration"""
# Calculate signal duration
signal_changes = results['signal'].diff().abs()
signal_changes = signal_changes[signal_changes > 0]
# Create intervals between signal changes
durations = []
current_duration = 0
current_signal = 0
for idx, row in results.iterrows():
if row['signal'] != current_signal:
if current_signal != 0: # Only count actual signal durations
durations.append(current_duration)
current_duration = 1
current_signal = row['signal']
else:
current_duration += 1
# Add the last duration if it's a signal
if current_signal != 0:
durations.append(current_duration)
# Convert to 5-minute intervals
durations_minutes = [d * 5 for d in durations]
# Plot histogram
bins = [0, 15, 30, 60, 120, 240, 480, 720, 1440]
labels = ['0-15m', '15-30m', '30-60m', '1-2h', '2-4h', '4-8h', '8-12h', '12-24h']
ax.hist(durations_minutes, bins=bins, alpha=0.7, color='teal',
edgecolor='black', linewidth=1)
# Add labels and title
ax.set_xlabel('Signal Duration (minutes)', fontsize=12)
ax.set_ylabel('Frequency', fontsize=12)
ax.set_title('Signal Duration Distribution', fontsize=14)
# Set custom x-ticks
ax.set_xticks([b + (bins[i+1] - b)/2 for i, b in enumerate(bins[:-1])])
ax.set_xticklabels(labels)
plt.setp(ax.get_xticklabels(), rotation=30, ha='right')
# Add summary statistics
mean_duration = np.mean(durations_minutes)
median_duration = np.median(durations_minutes)
stats_text = (
f"Mean: {mean_duration:.1f} min\n"
f"Median: {median_duration:.1f} min"
)
ax.text(0.7, 0.8, stats_text,
transform=ax.transAxes,
bbox=dict(boxstyle='round', facecolor='white', alpha=0.8),
fontsize=10)
def _plot_return_distribution(ax, returns, results):
"""Plot distribution of strategy returns"""
# Separate returns by signal type
long_returns = returns[results['signal'] == 1]
short_returns = returns[results['signal'] == -1]
# Create histogram
n_bins = 30
ax.hist(long_returns, bins=n_bins, alpha=0.5, color='green', label='Long')
ax.hist(short_returns, bins=n_bins, alpha=0.5, color='red', label='Short')
# Add normal distribution for reference
from scipy import stats
x = np.linspace(min(returns), max(returns), 100)
all_returns = returns[results['signal'] != 0]
mu, std = all_returns.mean(), all_returns.std()
pdf = stats.norm.pdf(x, mu, std)
scaled_pdf = pdf * (len(all_returns) * (max(returns) - min(returns)) / n_bins)
ax.plot(x, scaled_pdf, 'k--', linewidth=1, label='Normal Dist.')
# Add vertical line at 0
ax.axvline(x=0, color='black', linestyle='-', alpha=0.3)
# Add labels and title
ax.set_xlabel('Return', fontsize=12)
ax.set_ylabel('Frequency', fontsize=12)
ax.set_title('Return Distribution by Signal Type', fontsize=14)
# Add summary statistics
long_stats = (
f"Long Signals:\n"
f"Mean: {long_returns.mean():.2%}\n"
f"Std: {long_returns.std():.2%}\n"
f"Win: {(long_returns > 0).mean():.1%}"
)
short_stats = (
f"Short Signals:\n"
f"Mean: {short_returns.mean():.2%}\n"
f"Std: {short_returns.std():.2%}\n"
f"Win: {(short_returns > 0).mean():.1%}"
)
ax.text(0.05, 0.95, long_stats,
transform=ax.transAxes, va='top',
bbox=dict(boxstyle='round', facecolor='white', alpha=0.8),
fontsize=9)
ax.text(0.95, 0.95, short_stats,
transform=ax.transAxes, va='top', ha='right',
bbox=dict(boxstyle='round', facecolor='white', alpha=0.8),
fontsize=9)
# Add legend
ax.legend()
def _plot_signal_consistency(ax, results):
"""Plot signal consistency over time"""
# Calculate rolling signal consistency
window = 10 # Number of signals to check
# Get signal direction changes
signal_data = results[results['signal'] != 0].copy()
signal_data['prev_signal'] = signal_data['signal'].shift(1)
signal_data['direction_change'] = (signal_data['signal'] != signal_data['prev_signal']) & (signal_data['prev_signal'] != 0)
# Calculate rolling consistency
signal_data['consistency'] = 1 - signal_data['direction_change'].rolling(window).mean()
# Plot
ax.plot(signal_data.index, signal_data['consistency'] * 100, color='purple')
# Add horizontal line at 50%
ax.axhline(y=50, color='gray', linestyle='--', alpha=0.5)
# Format y-axis as percentage
ax.yaxis.set_major_formatter(mtick.PercentFormatter())
# Format x-axis for dates
ax.xaxis.set_major_formatter(mdates.DateFormatter('%Y-%m-%d'))
ax.xaxis.set_major_locator(mdates.WeekdayLocator(interval=7))
plt.setp(ax.xaxis.get_majorticklabels(), rotation=45, ha='right')
# Add labels and title
ax.set_xlabel('Date', fontsize=12)
ax.set_ylabel('Signal Consistency (%)', fontsize=12)
ax.set_title(f'Signal Consistency (Window={window})', fontsize=14)
# Set y-limit
ax.set_ylim(0, 100)
# Add average line
avg_consistency = signal_data['consistency'].mean() * 100
ax.axhline(y=avg_consistency, color='red', linestyle='-', alpha=0.5)
ax.text(signal_data.index[10], avg_consistency + 5,
f'Avg: {avg_consistency:.1f}%', color='red')
def _plot_model_agreement(ax, results, returns):
"""Plot performance by model agreement level"""
# Create bins for model agreement
results['agreement_bin'] = pd.cut(results['model_agreement'],
bins=[0, 0.6, 0.7, 0.8, 0.9, 1.0],
labels=['0-60%', '60-70%', '70-80%', '80-90%', '90-100%'])
# Combine with returns
performance = results.copy()
performance['return'] = returns
# Keep only actual signals
performance = performance[performance['signal'] != 0]
# Group by agreement bin
agreement_perf = performance.groupby('agreement_bin')
# Calculate metrics
win_rates = agreement_perf['return'].apply(lambda x: (x > 0).mean() if len(x) > 0 else 0)
avg_returns = agreement_perf['return'].apply(lambda x: x.mean() if len(x) > 0 else 0)
counts = agreement_perf.size()
# Create bar plot
bars = ax.bar(win_rates.index, win_rates, color='skyblue', alpha=0.7)
# Add count annotations
for i, (bar, count) in enumerate(zip(bars, counts)):
if count > 0:
height = bar.get_height()
ax.text(bar.get_x() + bar.get_width()/2., height + 0.02,
f'n={count}', ha='center', va='bottom', fontsize=9)
# Set up secondary axis for average return
ax2 = ax.twinx()
ax2.plot(avg_returns.index, avg_returns, color='darkblue', marker='o')
ax2.set_ylabel('Average Return', fontsize=12, color='darkblue')
# Add horizontal reference lines
ax.axhline(y=0.5, color='gray', linestyle='--', alpha=0.5)
ax2.axhline(y=0, color='darkblue', linestyle='--', alpha=0.5)
# Add labels and title
ax.set_xlabel('Model Agreement Level', fontsize=12)
ax.set_ylabel('Win Rate', fontsize=12)
ax.set_title('Performance by Model Agreement', fontsize=14)
ax.set_ylim(0, 1)
# Rotate x-labels for better readability
plt.setp(ax.get_xticklabels(), rotation=30, ha='right')
# Add custom legend
from matplotlib.lines import Line2D
legend_elements = [
Line2D([0], [0], color='skyblue', lw=0, marker='s', markersize=10, label='Win Rate'),
Line2D([0], [0], color='darkblue', marker='o', markersize=6, label='Avg Return')
]
ax.legend(handles=legend_elements, loc='upper left')
def _plot_signal_probability(ax, results, returns):
"""Plot performance by signal probability"""
# Create bins for signal probability
results['proba_bin'] = pd.cut(
np.where(results['signal'] == 1, results['proba_up'],
np.where(results['signal'] == -1, results['proba_down'], 0)),
bins=[0, 0.65, 0.7, 0.75, 0.8, 0.85, 0.9, 0.95, 1.0],
labels=['0-65%', '65-70%', '70-75%', '75-80%', '80-85%', '85-90%', '90-95%', '95-100%']
)
# Combine with returns
performance = results.copy()
performance['return'] = returns
# Keep only actual signals
performance = performance[performance['signal'] != 0]
# Group by probability bin
proba_perf = performance.groupby('proba_bin')
# Calculate metrics
win_rates = proba_perf['return'].apply(lambda x: (x > 0).mean() if len(x) > 0 else 0)
avg_returns = proba_perf['return'].apply(lambda x: x.mean() if len(x) > 0 else 0)
counts = proba_perf.size()
# Create bar plot
bars = ax.bar(win_rates.index, win_rates, color='lightcoral', alpha=0.7)
# Add count annotations
for i, (bar, count) in enumerate(zip(bars, counts)):
if count > 0:
height = bar.get_height()
ax.text(bar.get_x() + bar.get_width()/2., height + 0.02,
f'n={count}', ha='center', va='bottom', fontsize=9)
# Set up secondary axis for average return
ax2 = ax.twinx()
ax2.plot(avg_returns.index, avg_returns, color='darkred', marker='o')
ax2.set_ylabel('Average Return', fontsize=12, color='darkred')
# Add horizontal reference lines
ax.axhline(y=0.5, color='gray', linestyle='--', alpha=0.5)
ax2.axhline(y=0, color='darkred', linestyle='--', alpha=0.5)
# Add labels and title
ax.set_xlabel('Signal Probability', fontsize=12)
ax.set_ylabel('Win Rate', fontsize=12)
ax.set_title('Performance by Signal Probability', fontsize=14)
ax.set_ylim(0, 1)
# Rotate x-labels for better readability
plt.setp(ax.get_xticklabels(), rotation=45, ha='right')
# Add custom legend
from matplotlib.lines import Line2D
legend_elements = [
Line2D([0], [0], color='lightcoral', lw=0, marker='s', markersize=10, label='Win Rate'),
Line2D([0], [0], color='darkred', marker='o', markersize=6, label='Avg Return')
]
ax.legend(handles=legend_elements, loc='upper left')
def _plot_signal_strength_performance(ax, results, returns):
"""Plot performance by signal strength"""
# Create bins for signal strength
results['strength_bin'] = pd.cut(results['strength'],
bins=[0, 20, 40, 60, 80, 100],
labels=['0-20', '20-40', '40-60', '60-80', '80-100'])
# Combine with returns
performance = results.copy()
performance['return'] = returns
# Separate long and short signals
long_data = performance[performance['signal'] == 1]
short_data = performance[performance['signal'] == -1]
# Group by strength bin
long_perf = long_data.groupby('strength_bin')
short_perf = short_data.groupby('strength_bin')
# Calculate win rates
long_wins = long_perf['return'].apply(lambda x: (x > 0).mean() if len(x) > 0 else 0)
short_wins = short_perf['return'].apply(lambda x: (x > 0).mean() if len(x) > 0 else 0)
long_counts = long_perf.size()
short_counts = short_perf.size()
# Set width for bars
width = 0.35
x = np.arange(len(long_wins))
# Create grouped bar chart
long_bars = ax.bar(x - width/2, long_wins, width, label='Long Signals', color='green', alpha=0.7)
short_bars = ax.bar(x + width/2, short_wins, width, label='Short Signals', color='red', alpha=0.7)
# Add count annotations
for i, (bar, count) in enumerate(zip(long_bars, long_counts)):
if count > 0:
height = bar.get_height()
ax.text(bar.get_x() + bar.get_width()/2., height + 0.02,
f'{count}', ha='center', va='bottom', fontsize=8, color='green')
for i, (bar, count) in enumerate(zip(short_bars, short_counts)):
if count > 0:
height = bar.get_height()
ax.text(bar.get_x() + bar.get_width()/2., height + 0.02,
f'{count}', ha='center', va='bottom', fontsize=8, color='red')
# Add horizontal reference line
ax.axhline(y=0.5, color='gray', linestyle='--', alpha=0.5)
# Set up x-ticks
ax.set_xticks(x)
ax.set_xticklabels(long_wins.index)
# Add labels and title
ax.set_xlabel('Signal Strength (%)', fontsize=12)
ax.set_ylabel('Win Rate', fontsize=12)
ax.set_title('Win Rate by Signal Strength', fontsize=14)
ax.set_ylim(0, 1)
# Add legend
ax.legend()
def _plot_drawdown_analysis(ax, returns):
"""Plot drawdown analysis"""
# Calculate cumulative returns and drawdowns
cumulative_returns = (1 + returns).cumprod() - 1
drawdown = cumulative_returns - cumulative_returns.cummax()
# Plot drawdown
ax.fill_between(drawdown.index, 0, drawdown * 100, color='red', alpha=0.3)
ax.plot(drawdown.index, drawdown * 100, color='red', linewidth=1)
# Format y-axis as percentage
ax.yaxis.set_major_formatter(mtick.PercentFormatter())
# Format x-axis for dates
ax.xaxis.set_major_formatter(mdates.DateFormatter('%Y-%m-%d'))
ax.xaxis.set_major_locator(mdates.WeekdayLocator(interval=7))
plt.setp(ax.xaxis.get_majorticklabels(), rotation=45, ha='right')
# Find worst drawdowns
def find_drawdown_periods(drawdown_series, top_n=5):
periods = []
current_dd = 0
start_date = None
end_date = None
for date, value in drawdown_series.items():
if value < current_dd:
current_dd = value
end_date = date
elif value == 0 and current_dd < 0:
# Drawdown ended
periods.append((start_date, end_date, current_dd))
current_dd = 0
start_date = None
end_date = None
elif current_dd == 0 and value < 0:
# New drawdown started
start_date = date
current_dd = value
end_date = date
# Add any ongoing drawdown
if current_dd < 0:
periods.append((start_date, end_date, current_dd))
# Sort by largest drawdown and return top_n
return sorted(periods, key=lambda x: x[2])[:top_n]
# Get top drawdowns
top_drawdowns = find_drawdown_periods(drawdown, top_n=3)
# Highlight top drawdowns
colors = ['darkred', 'firebrick', 'indianred']
for i, (start, end, magnitude) in enumerate(top_drawdowns):
if start and end:
ax.axvspan(start, end, color=colors[i], alpha=0.2)
# Add annotation
mid_point = start + (end - start) / 2
ax.annotate(f"{magnitude*100:.1f}%", (mid_point, magnitude*100 - 0.5),
ha='center', fontsize=10, color=colors[i])
# Add labels and title
ax.set_ylabel('Drawdown (%)', fontsize=12)
ax.set_title('Drawdown Analysis', fontsize=14)
# Calculate and display statistics
max_dd = drawdown.min() * 100
avg_dd = drawdown[drawdown < 0].mean() * 100
stats_text = (
f"Max Drawdown: {max_dd:.2f}%\n"
f"Avg Drawdown: {avg_dd:.2f}%\n"
f"# of DDs >1%: {(drawdown < -0.01).sum()}"
)
ax.text(0.02, 0.05, stats_text,
transform=ax.transAxes,
bbox=dict(boxstyle='round', facecolor='white', alpha=0.8),
fontsize=10)
def analyze_trading_model(data_path, forecast_bars=24, confidence_threshold=0.65, save_path=None):
"""
Complete function to run model and create visualizations
Parameters:
data_path: Path to CSV file with OHLCV data
forecast_bars: Number of 5-min bars to forecast
confidence_threshold: Threshold for signal generation
save_path: Optional path to save visualization images
Returns:
predictor: Trained model
metrics: Performance metrics
results: Signal results
returns: Strategy returns
"""
# Import run_model function (assuming it's in your environment)
from tradingbot.models.tree_ensemble import run_model
# Run the model
predictor, metrics, results, returns = run_model(
data_path, forecast_bars, confidence_threshold
)
# Load original data
data = pd.read_csv(data_path)
data['timestamp'] = pd.to_datetime(data['timestamp'])
data = data.drop_duplicates(subset=['timestamp'])
data.set_index('timestamp', inplace=True)
# Create visualizations
print("\nGenerating visualizations...")
visualize_signals(data, results, returns, metrics, save_path)
return predictor, metrics, results, returns
def plot_price_signals(data, signals, price_col="close", title="Price with Trading Signals"):
"""Quick scatter of buy (1) / sell (-1) signals over the price series.
``signals`` is a Series (or array) aligned with ``data`` holding -1/0/1.
"""
df = data.copy()
df["signal"] = signals
plt.figure(figsize=(12, 6))
plt.plot(df.index, df[price_col], label="Price", color="blue", alpha=0.7)
buy = df[df["signal"] == 1][price_col]
plt.scatter(buy.index, buy, label="Buy Signal", color="green", marker="^", s=100)
sell = df[df["signal"] == -1][price_col]
plt.scatter(sell.index, sell, label="Sell Signal", color="red", marker="v", s=100)
plt.title(title)
plt.xlabel("Date")
plt.ylabel("Price")
plt.legend()
plt.grid(True)
plt.show()