| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625 |
- """
- Training visualization module for creating loss/accuracy plots and training reports.
- This module provides comprehensive visualization capabilities for training metrics,
- including loss curves, accuracy plots, learning rate schedules, and performance analysis.
- """
- import os
- import time
- from pathlib import Path
- from typing import Dict, List, Any, Optional, Tuple
- import numpy as np
- try:
- import matplotlib
- matplotlib.use('Agg') # Use non-interactive backend
- import matplotlib.pyplot as plt
- import matplotlib.dates as mdates
- from matplotlib.gridspec import GridSpec
- MATPLOTLIB_AVAILABLE = True
- except ImportError:
- MATPLOTLIB_AVAILABLE = False
- plt = None
- from .base import TrainingMetrics
- class TrainingVisualizer:
- """
- Creates comprehensive visualizations for training progress and results.
-
- Features:
- - Loss/accuracy curves
- - Learning rate schedules
- - Training time analysis
- - Performance metrics
- - Model comparison plots
- """
-
- def __init__(self, output_dir: str, model_name: str):
- """
- Initialize the training visualizer.
-
- Args:
- output_dir: Directory to save plots
- model_name: Name of the model for file naming
- """
- self.output_dir = Path(output_dir)
- self.model_name = model_name
- self.plots_dir = self.output_dir / "plots"
- self.plots_dir.mkdir(parents=True, exist_ok=True)
-
- if not MATPLOTLIB_AVAILABLE:
- print("Warning: matplotlib not available. Plots will not be generated.")
-
- def create_training_plots(self, metrics: TrainingMetrics,
- additional_metrics: Dict[str, Any] = None) -> List[str]:
- """
- Create comprehensive training visualization plots.
-
- Args:
- metrics: TrainingMetrics object with training history
- additional_metrics: Additional metrics to visualize
-
- Returns:
- List of generated plot file paths
- """
- if not MATPLOTLIB_AVAILABLE:
- return []
-
- generated_plots = []
-
- # Main training curves plot
- main_plot = self._create_main_training_plot(metrics)
- if main_plot:
- generated_plots.append(main_plot)
-
- # Learning rate schedule plot
- lr_plot = self._create_learning_rate_plot(metrics)
- if lr_plot:
- generated_plots.append(lr_plot)
-
- # Training time analysis
- time_plot = self._create_time_analysis_plot(metrics)
- if time_plot:
- generated_plots.append(time_plot)
-
- # Loss distribution plot
- loss_dist_plot = self._create_loss_distribution_plot(metrics)
- if loss_dist_plot:
- generated_plots.append(loss_dist_plot)
-
- # Additional metrics plots
- if additional_metrics:
- additional_plots = self._create_additional_metrics_plots(additional_metrics)
- generated_plots.extend(additional_plots)
-
- return generated_plots
-
- def _create_main_training_plot(self, metrics: TrainingMetrics) -> Optional[str]:
- """Create the main training curves plot (loss + accuracy)."""
- if not metrics.metrics['train_loss']:
- return None
-
- fig = plt.figure(figsize=(14, 8))
- gs = GridSpec(2, 2, hspace=0.3, wspace=0.3)
-
- epochs = list(range(len(metrics.metrics['train_loss'])))
-
- # Training and validation loss
- ax1 = fig.add_subplot(gs[0, :])
- ax1.plot(epochs, metrics.metrics['train_loss'], 'b-', label='Train Loss', linewidth=2)
- if metrics.metrics['val_loss']:
- # Filter out None values for validation loss
- val_epochs = [i for i, loss in enumerate(metrics.metrics['val_loss']) if loss is not None]
- val_losses = [loss for loss in metrics.metrics['val_loss'] if loss is not None]
- ax1.plot(val_epochs, val_losses, 'r-', label='Validation Loss', linewidth=2)
-
- ax1.set_xlabel('Epoch')
- ax1.set_ylabel('Loss')
- ax1.set_title(f'{self.model_name} - Training Progress', fontsize=14, fontweight='bold')
- ax1.legend()
- ax1.grid(True, alpha=0.3)
- ax1.set_yscale('log') # Log scale for better loss visualization
-
- # Training and validation accuracy
- ax2 = fig.add_subplot(gs[1, 0])
- if metrics.metrics['train_accuracy']:
- ax2.plot(epochs, metrics.metrics['train_accuracy'], 'b-', label='Train Accuracy', linewidth=2)
- if metrics.metrics['val_accuracy']:
- val_epochs = [i for i, acc in enumerate(metrics.metrics['val_accuracy']) if acc is not None]
- val_accs = [acc for acc in metrics.metrics['val_accuracy'] if acc is not None]
- ax2.plot(val_epochs, val_accs, 'r-', label='Validation Accuracy', linewidth=2)
-
- ax2.set_xlabel('Epoch')
- ax2.set_ylabel('Accuracy')
- ax2.set_title('Accuracy Progress')
- ax2.legend()
- ax2.grid(True, alpha=0.3)
- ax2.set_ylim(0, 1)
-
- # Training summary box
- ax3 = fig.add_subplot(gs[1, 1])
- ax3.axis('off')
-
- # Create summary text
- summary_text = self._create_training_summary(metrics)
- ax3.text(0.05, 0.95, summary_text, transform=ax3.transAxes, fontsize=10,
- verticalalignment='top', bbox=dict(boxstyle="round,pad=0.3", facecolor="lightblue", alpha=0.5))
-
- plt.suptitle(f'{self.model_name} Training Overview', fontsize=16, fontweight='bold')
-
- # Save plot
- plot_path = self.plots_dir / f"{self.model_name}_training_curves.png"
- plt.savefig(plot_path, dpi=300, bbox_inches='tight')
- plt.close()
-
- return str(plot_path)
-
- def _create_learning_rate_plot(self, metrics: TrainingMetrics) -> Optional[str]:
- """Create learning rate schedule visualization."""
- if not metrics.metrics['learning_rates']:
- return None
-
- fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(12, 8))
-
- epochs = list(range(len(metrics.metrics['learning_rates'])))
-
- # Learning rate over time
- ax1.plot(epochs, metrics.metrics['learning_rates'], 'g-', linewidth=2)
- ax1.set_xlabel('Epoch')
- ax1.set_ylabel('Learning Rate')
- ax1.set_title(f'{self.model_name} - Learning Rate Schedule')
- ax1.grid(True, alpha=0.3)
- ax1.set_yscale('log')
-
- # Learning rate vs loss correlation
- if metrics.metrics['train_loss']:
- ax2.scatter(metrics.metrics['learning_rates'], metrics.metrics['train_loss'],
- alpha=0.6, c=epochs, cmap='viridis')
- ax2.set_xlabel('Learning Rate')
- ax2.set_ylabel('Training Loss')
- ax2.set_title('Learning Rate vs Training Loss')
- ax2.set_xscale('log')
- ax2.set_yscale('log')
- ax2.grid(True, alpha=0.3)
-
- # Add colorbar
- cbar = plt.colorbar(ax2.collections[0], ax=ax2)
- cbar.set_label('Epoch')
-
- plt.tight_layout()
-
- # Save plot
- plot_path = self.plots_dir / f"{self.model_name}_learning_rate.png"
- plt.savefig(plot_path, dpi=300, bbox_inches='tight')
- plt.close()
-
- return str(plot_path)
-
- def _create_time_analysis_plot(self, metrics: TrainingMetrics) -> Optional[str]:
- """Create training time analysis visualization."""
- if not metrics.metrics['epoch_times']:
- return None
-
- fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(14, 10))
-
- epochs = list(range(len(metrics.metrics['epoch_times'])))
- epoch_times = metrics.metrics['epoch_times']
-
- # Epoch time progression
- ax1.plot(epochs, epoch_times, 'purple', linewidth=2)
- ax1.set_xlabel('Epoch')
- ax1.set_ylabel('Time (seconds)')
- ax1.set_title('Training Time per Epoch')
- ax1.grid(True, alpha=0.3)
-
- # Cumulative time
- cumulative_time = np.cumsum(epoch_times)
- ax2.plot(epochs, cumulative_time / 3600, 'orange', linewidth=2) # Convert to hours
- ax2.set_xlabel('Epoch')
- ax2.set_ylabel('Cumulative Time (hours)')
- ax2.set_title('Cumulative Training Time')
- ax2.grid(True, alpha=0.3)
-
- # Time distribution histogram
- ax3.hist(epoch_times, bins=20, alpha=0.7, color='skyblue', edgecolor='black')
- ax3.set_xlabel('Epoch Time (seconds)')
- ax3.set_ylabel('Frequency')
- ax3.set_title('Epoch Time Distribution')
- ax3.grid(True, alpha=0.3)
-
- # Training efficiency (loss reduction per time)
- if metrics.metrics['train_loss'] and len(metrics.metrics['train_loss']) > 1:
- loss_improvements = []
- time_ratios = []
-
- for i in range(1, len(metrics.metrics['train_loss'])):
- if metrics.metrics['train_loss'][i-1] > 0:
- loss_improvement = (metrics.metrics['train_loss'][i-1] - metrics.metrics['train_loss'][i]) / metrics.metrics['train_loss'][i-1]
- time_ratio = epoch_times[i] / np.mean(epoch_times[:i+1])
-
- loss_improvements.append(loss_improvement)
- time_ratios.append(time_ratio)
-
- if loss_improvements:
- ax4.scatter(time_ratios, loss_improvements, alpha=0.6, c=range(len(loss_improvements)), cmap='viridis')
- ax4.set_xlabel('Relative Epoch Time')
- ax4.set_ylabel('Loss Improvement Ratio')
- ax4.set_title('Training Efficiency')
- ax4.grid(True, alpha=0.3)
-
- plt.suptitle(f'{self.model_name} Training Time Analysis', fontsize=14, fontweight='bold')
- plt.tight_layout()
-
- # Save plot
- plot_path = self.plots_dir / f"{self.model_name}_time_analysis.png"
- plt.savefig(plot_path, dpi=300, bbox_inches='tight')
- plt.close()
-
- return str(plot_path)
-
- def _create_loss_distribution_plot(self, metrics: TrainingMetrics) -> Optional[str]:
- """Create loss distribution and convergence analysis."""
- if not metrics.metrics['train_loss']:
- return None
-
- fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(14, 10))
-
- train_losses = metrics.metrics['train_loss']
- epochs = list(range(len(train_losses)))
-
- # Loss convergence (moving average)
- window_size = max(5, len(train_losses) // 20)
- if len(train_losses) >= window_size:
- moving_avg = np.convolve(train_losses, np.ones(window_size)/window_size, mode='valid')
- moving_epochs = epochs[window_size-1:]
-
- ax1.plot(epochs, train_losses, alpha=0.3, color='blue', label='Raw Loss')
- ax1.plot(moving_epochs, moving_avg, color='red', linewidth=2, label=f'Moving Average ({window_size})')
- ax1.set_xlabel('Epoch')
- ax1.set_ylabel('Loss')
- ax1.set_title('Loss Convergence Analysis')
- ax1.legend()
- ax1.grid(True, alpha=0.3)
- ax1.set_yscale('log')
-
- # Loss distribution histogram
- ax2.hist(train_losses, bins=30, alpha=0.7, color='green', edgecolor='black')
- ax2.set_xlabel('Loss Value')
- ax2.set_ylabel('Frequency')
- ax2.set_title('Training Loss Distribution')
- ax2.grid(True, alpha=0.3)
-
- # Loss gradient (rate of change)
- if len(train_losses) > 1:
- loss_gradients = np.diff(train_losses)
- ax3.plot(epochs[1:], loss_gradients, color='purple', linewidth=1)
- ax3.axhline(y=0, color='red', linestyle='--', alpha=0.7)
- ax3.set_xlabel('Epoch')
- ax3.set_ylabel('Loss Change')
- ax3.set_title('Loss Gradient (Rate of Change)')
- ax3.grid(True, alpha=0.3)
-
- # Loss stability (rolling standard deviation)
- if len(train_losses) >= 10:
- window = 10
- rolling_std = []
- for i in range(window, len(train_losses)):
- std = np.std(train_losses[i-window:i])
- rolling_std.append(std)
-
- ax4.plot(epochs[window:], rolling_std, color='orange', linewidth=2)
- ax4.set_xlabel('Epoch')
- ax4.set_ylabel('Rolling Std Dev')
- ax4.set_title(f'Loss Stability (Rolling {window}-Epoch Window)')
- ax4.grid(True, alpha=0.3)
-
- plt.suptitle(f'{self.model_name} Loss Analysis', fontsize=14, fontweight='bold')
- plt.tight_layout()
-
- # Save plot
- plot_path = self.plots_dir / f"{self.model_name}_loss_analysis.png"
- plt.savefig(plot_path, dpi=300, bbox_inches='tight')
- plt.close()
-
- return str(plot_path)
-
- def _create_additional_metrics_plots(self, additional_metrics: Dict[str, Any]) -> List[str]:
- """Create plots for additional metrics like speaker verification, etc."""
- generated_plots = []
-
- # Speaker verification metrics (if available)
- if 'speaker_verification' in additional_metrics:
- plot_path = self._create_speaker_verification_plot(additional_metrics['speaker_verification'])
- if plot_path:
- generated_plots.append(plot_path)
-
- # Embedding analysis (if available)
- if 'embedding_analysis' in additional_metrics:
- plot_path = self._create_embedding_analysis_plot(additional_metrics['embedding_analysis'])
- if plot_path:
- generated_plots.append(plot_path)
-
- return generated_plots
-
- def _create_speaker_verification_plot(self, verification_metrics: Dict[str, Any]) -> Optional[str]:
- """Create speaker verification performance visualization."""
- fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(14, 10))
-
- # EER visualization
- eer = verification_metrics.get('equal_error_rate', 0)
- eer_threshold = verification_metrics.get('eer_threshold', 0)
-
- ax1.bar(['EER'], [eer], color='red', alpha=0.7)
- ax1.set_ylabel('Equal Error Rate')
- ax1.set_title(f'Speaker Verification EER: {eer:.4f}')
- ax1.grid(True, alpha=0.3)
-
- # Similarity distributions
- pos_sim = verification_metrics.get('mean_positive_similarity', 0)
- neg_sim = verification_metrics.get('mean_negative_similarity', 0)
-
- ax2.bar(['Positive Pairs', 'Negative Pairs'], [pos_sim, neg_sim],
- color=['green', 'red'], alpha=0.7)
- ax2.set_ylabel('Mean Similarity')
- ax2.set_title('Speaker Similarity Distribution')
- ax2.grid(True, alpha=0.3)
-
- # Threshold visualization
- thresholds = np.linspace(neg_sim - 0.2, pos_sim + 0.2, 100)
- ax3.axvline(x=eer_threshold, color='red', linestyle='--', linewidth=2, label=f'EER Threshold: {eer_threshold:.3f}')
- ax3.axvline(x=eer_threshold + 0.1, color='orange', linestyle='--', alpha=0.7, label='Conservative: +0.1')
- ax3.axvline(x=eer_threshold - 0.1, color='blue', linestyle='--', alpha=0.7, label='Liberal: -0.1')
- ax3.set_xlabel('Similarity Threshold')
- ax3.set_title('Recommended Thresholds')
- ax3.legend()
- ax3.grid(True, alpha=0.3)
-
- # Performance summary
- ax4.axis('off')
- summary_text = f"""Speaker Verification Summary:
-
- EER: {eer:.4f}
- EER Threshold: {eer_threshold:.4f}
- Positive Pairs: {verification_metrics.get('num_positive_pairs', 0):,}
- Negative Pairs: {verification_metrics.get('num_negative_pairs', 0):,}
- Mean Positive Similarity: {pos_sim:.4f}
- Mean Negative Similarity: {neg_sim:.4f}
- Separability: {pos_sim - neg_sim:.4f}
- """
-
- ax4.text(0.05, 0.95, summary_text, transform=ax4.transAxes, fontsize=12,
- verticalalignment='top', bbox=dict(boxstyle="round,pad=0.3", facecolor="lightgreen", alpha=0.5))
-
- plt.suptitle(f'{self.model_name} Speaker Verification Analysis', fontsize=14, fontweight='bold')
- plt.tight_layout()
-
- # Save plot
- plot_path = self.plots_dir / f"{self.model_name}_speaker_verification.png"
- plt.savefig(plot_path, dpi=300, bbox_inches='tight')
- plt.close()
-
- return str(plot_path)
-
- def _create_embedding_analysis_plot(self, embedding_metrics: Dict[str, Any]) -> Optional[str]:
- """Create embedding quality analysis visualization."""
- fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(14, 10))
-
- # Distance comparison
- intra_dist = embedding_metrics.get('mean_intra_distance', 0)
- inter_dist = embedding_metrics.get('mean_inter_distance', 0)
- separability = embedding_metrics.get('separability_ratio', 0)
-
- distances = ['Intra-Speaker', 'Inter-Speaker']
- values = [intra_dist, inter_dist]
- colors = ['blue', 'red']
-
- bars = ax1.bar(distances, values, color=colors, alpha=0.7)
- ax1.set_ylabel('Distance')
- ax1.set_title('Speaker Embedding Distances')
- ax1.grid(True, alpha=0.3)
-
- # Add value labels on bars
- for bar, value in zip(bars, values):
- height = bar.get_height()
- ax1.text(bar.get_x() + bar.get_width()/2., height + height*0.01,
- f'{value:.4f}', ha='center', va='bottom')
-
- # Separability ratio
- ax2.bar(['Separability Ratio'], [separability], color='green', alpha=0.7)
- ax2.set_ylabel('Ratio (Inter/Intra)')
- ax2.set_title(f'Embedding Separability: {separability:.2f}')
- ax2.grid(True, alpha=0.3)
-
- # Add separability quality indicator
- if separability > 2.0:
- quality = "Excellent"
- color = 'green'
- elif separability > 1.5:
- quality = "Good"
- color = 'orange'
- else:
- quality = "Poor"
- color = 'red'
-
- ax2.text(0, separability + separability*0.05, quality, ha='center',
- fontweight='bold', color=color, fontsize=14)
-
- # Speaker count information
- num_speakers = embedding_metrics.get('num_speakers', 0)
- ax3.pie([num_speakers, max(1, 10 - num_speakers)], labels=['Trained Speakers', 'Potential'],
- autopct='%1.0f', startangle=90, colors=['lightblue', 'lightgray'])
- ax3.set_title(f'Speaker Coverage ({num_speakers} speakers)')
-
- # Quality assessment
- ax4.axis('off')
-
- # Determine overall quality
- if separability > 2.0 and intra_dist < inter_dist:
- overall_quality = "EXCELLENT"
- quality_color = "lightgreen"
- elif separability > 1.5:
- overall_quality = "GOOD"
- quality_color = "lightyellow"
- else:
- overall_quality = "NEEDS IMPROVEMENT"
- quality_color = "lightcoral"
-
- summary_text = f"""Embedding Quality Assessment:
-
- Overall Quality: {overall_quality}
- Metrics:
- • Number of Speakers: {num_speakers}
- • Intra-Speaker Distance: {intra_dist:.4f}
- • Inter-Speaker Distance: {inter_dist:.4f}
- • Separability Ratio: {separability:.2f}
- Quality Indicators:
- • Distance Separation: {'✓' if inter_dist > intra_dist else '✗'}
- • Good Separability: {'✓' if separability > 1.5 else '✗'}
- • Excellent Separability: {'✓' if separability > 2.0 else '✗'}
- """
-
- ax4.text(0.05, 0.95, summary_text, transform=ax4.transAxes, fontsize=11,
- verticalalignment='top', bbox=dict(boxstyle="round,pad=0.3", facecolor=quality_color, alpha=0.7))
-
- plt.suptitle(f'{self.model_name} Embedding Quality Analysis', fontsize=14, fontweight='bold')
- plt.tight_layout()
-
- # Save plot
- plot_path = self.plots_dir / f"{self.model_name}_embedding_analysis.png"
- plt.savefig(plot_path, dpi=300, bbox_inches='tight')
- plt.close()
-
- return str(plot_path)
-
- def _create_training_summary(self, metrics: TrainingMetrics) -> str:
- """Create a text summary of training results."""
- total_epochs = len(metrics.metrics['train_loss'])
- final_train_loss = metrics.metrics['train_loss'][-1] if metrics.metrics['train_loss'] else 0
- final_train_acc = metrics.metrics['train_accuracy'][-1] if metrics.metrics['train_accuracy'] else 0
-
- val_losses = [loss for loss in metrics.metrics['val_loss'] if loss is not None]
- final_val_loss = val_losses[-1] if val_losses else 0
-
- val_accs = [acc for acc in metrics.metrics['val_accuracy'] if acc is not None]
- final_val_acc = val_accs[-1] if val_accs else 0
-
- total_time = sum(metrics.metrics['epoch_times']) if metrics.metrics['epoch_times'] else 0
- avg_epoch_time = total_time / max(1, total_epochs)
-
- summary = f"""Training Summary:
-
- Total Epochs: {total_epochs}
- Total Time: {total_time/3600:.2f} hours
- Avg Time/Epoch: {avg_epoch_time:.2f} sec
- Final Metrics:
- Train Loss: {final_train_loss:.6f}
- Train Accuracy: {final_train_acc:.4f}
- Val Loss: {final_val_loss:.6f}
- Val Accuracy: {final_val_acc:.4f}
- Best Validation:
- Loss: {metrics.best_val_loss:.6f}
- Epoch: {metrics.best_epoch}
- """
-
- return summary
-
- def create_comparison_plot(self, multiple_metrics: Dict[str, TrainingMetrics],
- title: str = "Model Comparison") -> Optional[str]:
- """
- Create comparison plots for multiple training runs.
-
- Args:
- multiple_metrics: Dictionary mapping model names to their metrics
- title: Title for the comparison plot
-
- Returns:
- Path to generated comparison plot
- """
- if not MATPLOTLIB_AVAILABLE or not multiple_metrics:
- return None
-
- fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(16, 12))
-
- colors = plt.cm.Set1(np.linspace(0, 1, len(multiple_metrics)))
-
- for (model_name, metrics), color in zip(multiple_metrics.items(), colors):
- epochs = list(range(len(metrics.metrics['train_loss'])))
-
- # Training loss comparison
- ax1.plot(epochs, metrics.metrics['train_loss'],
- label=model_name, color=color, linewidth=2)
-
- # Validation loss comparison
- if metrics.metrics['val_loss']:
- val_epochs = [i for i, loss in enumerate(metrics.metrics['val_loss']) if loss is not None]
- val_losses = [loss for loss in metrics.metrics['val_loss'] if loss is not None]
- ax2.plot(val_epochs, val_losses,
- label=model_name, color=color, linewidth=2, linestyle='--')
-
- # Training accuracy comparison
- if metrics.metrics['train_accuracy']:
- ax3.plot(epochs, metrics.metrics['train_accuracy'],
- label=model_name, color=color, linewidth=2)
-
- # Training time comparison
- if metrics.metrics['epoch_times']:
- ax4.plot(epochs, np.cumsum(metrics.metrics['epoch_times']) / 3600,
- label=model_name, color=color, linewidth=2)
-
- # Configure subplots
- ax1.set_xlabel('Epoch')
- ax1.set_ylabel('Training Loss')
- ax1.set_title('Training Loss Comparison')
- ax1.legend()
- ax1.grid(True, alpha=0.3)
- ax1.set_yscale('log')
-
- ax2.set_xlabel('Epoch')
- ax2.set_ylabel('Validation Loss')
- ax2.set_title('Validation Loss Comparison')
- ax2.legend()
- ax2.grid(True, alpha=0.3)
- ax2.set_yscale('log')
-
- ax3.set_xlabel('Epoch')
- ax3.set_ylabel('Accuracy')
- ax3.set_title('Training Accuracy Comparison')
- ax3.legend()
- ax3.grid(True, alpha=0.3)
-
- ax4.set_xlabel('Epoch')
- ax4.set_ylabel('Cumulative Time (hours)')
- ax4.set_title('Training Time Comparison')
- ax4.legend()
- ax4.grid(True, alpha=0.3)
-
- plt.suptitle(title, fontsize=16, fontweight='bold')
- plt.tight_layout()
-
- # Save plot
- plot_path = self.plots_dir / f"model_comparison_{int(time.time())}.png"
- plt.savefig(plot_path, dpi=300, bbox_inches='tight')
- plt.close()
-
- return str(plot_path)
- def create_training_visualizer(output_dir: str, model_name: str) -> TrainingVisualizer:
- """
- Factory function to create a training visualizer.
-
- Args:
- output_dir: Directory to save plots
- model_name: Name of the model
-
- Returns:
- TrainingVisualizer instance
- """
- return TrainingVisualizer(output_dir=output_dir, model_name=model_name)
|