visualization.py 25 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625
  1. """
  2. Training visualization module for creating loss/accuracy plots and training reports.
  3. This module provides comprehensive visualization capabilities for training metrics,
  4. including loss curves, accuracy plots, learning rate schedules, and performance analysis.
  5. """
  6. import os
  7. import time
  8. from pathlib import Path
  9. from typing import Dict, List, Any, Optional, Tuple
  10. import numpy as np
  11. try:
  12. import matplotlib
  13. matplotlib.use('Agg') # Use non-interactive backend
  14. import matplotlib.pyplot as plt
  15. import matplotlib.dates as mdates
  16. from matplotlib.gridspec import GridSpec
  17. MATPLOTLIB_AVAILABLE = True
  18. except ImportError:
  19. MATPLOTLIB_AVAILABLE = False
  20. plt = None
  21. from .base import TrainingMetrics
  22. class TrainingVisualizer:
  23. """
  24. Creates comprehensive visualizations for training progress and results.
  25. Features:
  26. - Loss/accuracy curves
  27. - Learning rate schedules
  28. - Training time analysis
  29. - Performance metrics
  30. - Model comparison plots
  31. """
  32. def __init__(self, output_dir: str, model_name: str):
  33. """
  34. Initialize the training visualizer.
  35. Args:
  36. output_dir: Directory to save plots
  37. model_name: Name of the model for file naming
  38. """
  39. self.output_dir = Path(output_dir)
  40. self.model_name = model_name
  41. self.plots_dir = self.output_dir / "plots"
  42. self.plots_dir.mkdir(parents=True, exist_ok=True)
  43. if not MATPLOTLIB_AVAILABLE:
  44. print("Warning: matplotlib not available. Plots will not be generated.")
  45. def create_training_plots(self, metrics: TrainingMetrics,
  46. additional_metrics: Dict[str, Any] = None) -> List[str]:
  47. """
  48. Create comprehensive training visualization plots.
  49. Args:
  50. metrics: TrainingMetrics object with training history
  51. additional_metrics: Additional metrics to visualize
  52. Returns:
  53. List of generated plot file paths
  54. """
  55. if not MATPLOTLIB_AVAILABLE:
  56. return []
  57. generated_plots = []
  58. # Main training curves plot
  59. main_plot = self._create_main_training_plot(metrics)
  60. if main_plot:
  61. generated_plots.append(main_plot)
  62. # Learning rate schedule plot
  63. lr_plot = self._create_learning_rate_plot(metrics)
  64. if lr_plot:
  65. generated_plots.append(lr_plot)
  66. # Training time analysis
  67. time_plot = self._create_time_analysis_plot(metrics)
  68. if time_plot:
  69. generated_plots.append(time_plot)
  70. # Loss distribution plot
  71. loss_dist_plot = self._create_loss_distribution_plot(metrics)
  72. if loss_dist_plot:
  73. generated_plots.append(loss_dist_plot)
  74. # Additional metrics plots
  75. if additional_metrics:
  76. additional_plots = self._create_additional_metrics_plots(additional_metrics)
  77. generated_plots.extend(additional_plots)
  78. return generated_plots
  79. def _create_main_training_plot(self, metrics: TrainingMetrics) -> Optional[str]:
  80. """Create the main training curves plot (loss + accuracy)."""
  81. if not metrics.metrics['train_loss']:
  82. return None
  83. fig = plt.figure(figsize=(14, 8))
  84. gs = GridSpec(2, 2, hspace=0.3, wspace=0.3)
  85. epochs = list(range(len(metrics.metrics['train_loss'])))
  86. # Training and validation loss
  87. ax1 = fig.add_subplot(gs[0, :])
  88. ax1.plot(epochs, metrics.metrics['train_loss'], 'b-', label='Train Loss', linewidth=2)
  89. if metrics.metrics['val_loss']:
  90. # Filter out None values for validation loss
  91. val_epochs = [i for i, loss in enumerate(metrics.metrics['val_loss']) if loss is not None]
  92. val_losses = [loss for loss in metrics.metrics['val_loss'] if loss is not None]
  93. ax1.plot(val_epochs, val_losses, 'r-', label='Validation Loss', linewidth=2)
  94. ax1.set_xlabel('Epoch')
  95. ax1.set_ylabel('Loss')
  96. ax1.set_title(f'{self.model_name} - Training Progress', fontsize=14, fontweight='bold')
  97. ax1.legend()
  98. ax1.grid(True, alpha=0.3)
  99. ax1.set_yscale('log') # Log scale for better loss visualization
  100. # Training and validation accuracy
  101. ax2 = fig.add_subplot(gs[1, 0])
  102. if metrics.metrics['train_accuracy']:
  103. ax2.plot(epochs, metrics.metrics['train_accuracy'], 'b-', label='Train Accuracy', linewidth=2)
  104. if metrics.metrics['val_accuracy']:
  105. val_epochs = [i for i, acc in enumerate(metrics.metrics['val_accuracy']) if acc is not None]
  106. val_accs = [acc for acc in metrics.metrics['val_accuracy'] if acc is not None]
  107. ax2.plot(val_epochs, val_accs, 'r-', label='Validation Accuracy', linewidth=2)
  108. ax2.set_xlabel('Epoch')
  109. ax2.set_ylabel('Accuracy')
  110. ax2.set_title('Accuracy Progress')
  111. ax2.legend()
  112. ax2.grid(True, alpha=0.3)
  113. ax2.set_ylim(0, 1)
  114. # Training summary box
  115. ax3 = fig.add_subplot(gs[1, 1])
  116. ax3.axis('off')
  117. # Create summary text
  118. summary_text = self._create_training_summary(metrics)
  119. ax3.text(0.05, 0.95, summary_text, transform=ax3.transAxes, fontsize=10,
  120. verticalalignment='top', bbox=dict(boxstyle="round,pad=0.3", facecolor="lightblue", alpha=0.5))
  121. plt.suptitle(f'{self.model_name} Training Overview', fontsize=16, fontweight='bold')
  122. # Save plot
  123. plot_path = self.plots_dir / f"{self.model_name}_training_curves.png"
  124. plt.savefig(plot_path, dpi=300, bbox_inches='tight')
  125. plt.close()
  126. return str(plot_path)
  127. def _create_learning_rate_plot(self, metrics: TrainingMetrics) -> Optional[str]:
  128. """Create learning rate schedule visualization."""
  129. if not metrics.metrics['learning_rates']:
  130. return None
  131. fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(12, 8))
  132. epochs = list(range(len(metrics.metrics['learning_rates'])))
  133. # Learning rate over time
  134. ax1.plot(epochs, metrics.metrics['learning_rates'], 'g-', linewidth=2)
  135. ax1.set_xlabel('Epoch')
  136. ax1.set_ylabel('Learning Rate')
  137. ax1.set_title(f'{self.model_name} - Learning Rate Schedule')
  138. ax1.grid(True, alpha=0.3)
  139. ax1.set_yscale('log')
  140. # Learning rate vs loss correlation
  141. if metrics.metrics['train_loss']:
  142. ax2.scatter(metrics.metrics['learning_rates'], metrics.metrics['train_loss'],
  143. alpha=0.6, c=epochs, cmap='viridis')
  144. ax2.set_xlabel('Learning Rate')
  145. ax2.set_ylabel('Training Loss')
  146. ax2.set_title('Learning Rate vs Training Loss')
  147. ax2.set_xscale('log')
  148. ax2.set_yscale('log')
  149. ax2.grid(True, alpha=0.3)
  150. # Add colorbar
  151. cbar = plt.colorbar(ax2.collections[0], ax=ax2)
  152. cbar.set_label('Epoch')
  153. plt.tight_layout()
  154. # Save plot
  155. plot_path = self.plots_dir / f"{self.model_name}_learning_rate.png"
  156. plt.savefig(plot_path, dpi=300, bbox_inches='tight')
  157. plt.close()
  158. return str(plot_path)
  159. def _create_time_analysis_plot(self, metrics: TrainingMetrics) -> Optional[str]:
  160. """Create training time analysis visualization."""
  161. if not metrics.metrics['epoch_times']:
  162. return None
  163. fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(14, 10))
  164. epochs = list(range(len(metrics.metrics['epoch_times'])))
  165. epoch_times = metrics.metrics['epoch_times']
  166. # Epoch time progression
  167. ax1.plot(epochs, epoch_times, 'purple', linewidth=2)
  168. ax1.set_xlabel('Epoch')
  169. ax1.set_ylabel('Time (seconds)')
  170. ax1.set_title('Training Time per Epoch')
  171. ax1.grid(True, alpha=0.3)
  172. # Cumulative time
  173. cumulative_time = np.cumsum(epoch_times)
  174. ax2.plot(epochs, cumulative_time / 3600, 'orange', linewidth=2) # Convert to hours
  175. ax2.set_xlabel('Epoch')
  176. ax2.set_ylabel('Cumulative Time (hours)')
  177. ax2.set_title('Cumulative Training Time')
  178. ax2.grid(True, alpha=0.3)
  179. # Time distribution histogram
  180. ax3.hist(epoch_times, bins=20, alpha=0.7, color='skyblue', edgecolor='black')
  181. ax3.set_xlabel('Epoch Time (seconds)')
  182. ax3.set_ylabel('Frequency')
  183. ax3.set_title('Epoch Time Distribution')
  184. ax3.grid(True, alpha=0.3)
  185. # Training efficiency (loss reduction per time)
  186. if metrics.metrics['train_loss'] and len(metrics.metrics['train_loss']) > 1:
  187. loss_improvements = []
  188. time_ratios = []
  189. for i in range(1, len(metrics.metrics['train_loss'])):
  190. if metrics.metrics['train_loss'][i-1] > 0:
  191. loss_improvement = (metrics.metrics['train_loss'][i-1] - metrics.metrics['train_loss'][i]) / metrics.metrics['train_loss'][i-1]
  192. time_ratio = epoch_times[i] / np.mean(epoch_times[:i+1])
  193. loss_improvements.append(loss_improvement)
  194. time_ratios.append(time_ratio)
  195. if loss_improvements:
  196. ax4.scatter(time_ratios, loss_improvements, alpha=0.6, c=range(len(loss_improvements)), cmap='viridis')
  197. ax4.set_xlabel('Relative Epoch Time')
  198. ax4.set_ylabel('Loss Improvement Ratio')
  199. ax4.set_title('Training Efficiency')
  200. ax4.grid(True, alpha=0.3)
  201. plt.suptitle(f'{self.model_name} Training Time Analysis', fontsize=14, fontweight='bold')
  202. plt.tight_layout()
  203. # Save plot
  204. plot_path = self.plots_dir / f"{self.model_name}_time_analysis.png"
  205. plt.savefig(plot_path, dpi=300, bbox_inches='tight')
  206. plt.close()
  207. return str(plot_path)
  208. def _create_loss_distribution_plot(self, metrics: TrainingMetrics) -> Optional[str]:
  209. """Create loss distribution and convergence analysis."""
  210. if not metrics.metrics['train_loss']:
  211. return None
  212. fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(14, 10))
  213. train_losses = metrics.metrics['train_loss']
  214. epochs = list(range(len(train_losses)))
  215. # Loss convergence (moving average)
  216. window_size = max(5, len(train_losses) // 20)
  217. if len(train_losses) >= window_size:
  218. moving_avg = np.convolve(train_losses, np.ones(window_size)/window_size, mode='valid')
  219. moving_epochs = epochs[window_size-1:]
  220. ax1.plot(epochs, train_losses, alpha=0.3, color='blue', label='Raw Loss')
  221. ax1.plot(moving_epochs, moving_avg, color='red', linewidth=2, label=f'Moving Average ({window_size})')
  222. ax1.set_xlabel('Epoch')
  223. ax1.set_ylabel('Loss')
  224. ax1.set_title('Loss Convergence Analysis')
  225. ax1.legend()
  226. ax1.grid(True, alpha=0.3)
  227. ax1.set_yscale('log')
  228. # Loss distribution histogram
  229. ax2.hist(train_losses, bins=30, alpha=0.7, color='green', edgecolor='black')
  230. ax2.set_xlabel('Loss Value')
  231. ax2.set_ylabel('Frequency')
  232. ax2.set_title('Training Loss Distribution')
  233. ax2.grid(True, alpha=0.3)
  234. # Loss gradient (rate of change)
  235. if len(train_losses) > 1:
  236. loss_gradients = np.diff(train_losses)
  237. ax3.plot(epochs[1:], loss_gradients, color='purple', linewidth=1)
  238. ax3.axhline(y=0, color='red', linestyle='--', alpha=0.7)
  239. ax3.set_xlabel('Epoch')
  240. ax3.set_ylabel('Loss Change')
  241. ax3.set_title('Loss Gradient (Rate of Change)')
  242. ax3.grid(True, alpha=0.3)
  243. # Loss stability (rolling standard deviation)
  244. if len(train_losses) >= 10:
  245. window = 10
  246. rolling_std = []
  247. for i in range(window, len(train_losses)):
  248. std = np.std(train_losses[i-window:i])
  249. rolling_std.append(std)
  250. ax4.plot(epochs[window:], rolling_std, color='orange', linewidth=2)
  251. ax4.set_xlabel('Epoch')
  252. ax4.set_ylabel('Rolling Std Dev')
  253. ax4.set_title(f'Loss Stability (Rolling {window}-Epoch Window)')
  254. ax4.grid(True, alpha=0.3)
  255. plt.suptitle(f'{self.model_name} Loss Analysis', fontsize=14, fontweight='bold')
  256. plt.tight_layout()
  257. # Save plot
  258. plot_path = self.plots_dir / f"{self.model_name}_loss_analysis.png"
  259. plt.savefig(plot_path, dpi=300, bbox_inches='tight')
  260. plt.close()
  261. return str(plot_path)
  262. def _create_additional_metrics_plots(self, additional_metrics: Dict[str, Any]) -> List[str]:
  263. """Create plots for additional metrics like speaker verification, etc."""
  264. generated_plots = []
  265. # Speaker verification metrics (if available)
  266. if 'speaker_verification' in additional_metrics:
  267. plot_path = self._create_speaker_verification_plot(additional_metrics['speaker_verification'])
  268. if plot_path:
  269. generated_plots.append(plot_path)
  270. # Embedding analysis (if available)
  271. if 'embedding_analysis' in additional_metrics:
  272. plot_path = self._create_embedding_analysis_plot(additional_metrics['embedding_analysis'])
  273. if plot_path:
  274. generated_plots.append(plot_path)
  275. return generated_plots
  276. def _create_speaker_verification_plot(self, verification_metrics: Dict[str, Any]) -> Optional[str]:
  277. """Create speaker verification performance visualization."""
  278. fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(14, 10))
  279. # EER visualization
  280. eer = verification_metrics.get('equal_error_rate', 0)
  281. eer_threshold = verification_metrics.get('eer_threshold', 0)
  282. ax1.bar(['EER'], [eer], color='red', alpha=0.7)
  283. ax1.set_ylabel('Equal Error Rate')
  284. ax1.set_title(f'Speaker Verification EER: {eer:.4f}')
  285. ax1.grid(True, alpha=0.3)
  286. # Similarity distributions
  287. pos_sim = verification_metrics.get('mean_positive_similarity', 0)
  288. neg_sim = verification_metrics.get('mean_negative_similarity', 0)
  289. ax2.bar(['Positive Pairs', 'Negative Pairs'], [pos_sim, neg_sim],
  290. color=['green', 'red'], alpha=0.7)
  291. ax2.set_ylabel('Mean Similarity')
  292. ax2.set_title('Speaker Similarity Distribution')
  293. ax2.grid(True, alpha=0.3)
  294. # Threshold visualization
  295. thresholds = np.linspace(neg_sim - 0.2, pos_sim + 0.2, 100)
  296. ax3.axvline(x=eer_threshold, color='red', linestyle='--', linewidth=2, label=f'EER Threshold: {eer_threshold:.3f}')
  297. ax3.axvline(x=eer_threshold + 0.1, color='orange', linestyle='--', alpha=0.7, label='Conservative: +0.1')
  298. ax3.axvline(x=eer_threshold - 0.1, color='blue', linestyle='--', alpha=0.7, label='Liberal: -0.1')
  299. ax3.set_xlabel('Similarity Threshold')
  300. ax3.set_title('Recommended Thresholds')
  301. ax3.legend()
  302. ax3.grid(True, alpha=0.3)
  303. # Performance summary
  304. ax4.axis('off')
  305. summary_text = f"""Speaker Verification Summary:
  306. EER: {eer:.4f}
  307. EER Threshold: {eer_threshold:.4f}
  308. Positive Pairs: {verification_metrics.get('num_positive_pairs', 0):,}
  309. Negative Pairs: {verification_metrics.get('num_negative_pairs', 0):,}
  310. Mean Positive Similarity: {pos_sim:.4f}
  311. Mean Negative Similarity: {neg_sim:.4f}
  312. Separability: {pos_sim - neg_sim:.4f}
  313. """
  314. ax4.text(0.05, 0.95, summary_text, transform=ax4.transAxes, fontsize=12,
  315. verticalalignment='top', bbox=dict(boxstyle="round,pad=0.3", facecolor="lightgreen", alpha=0.5))
  316. plt.suptitle(f'{self.model_name} Speaker Verification Analysis', fontsize=14, fontweight='bold')
  317. plt.tight_layout()
  318. # Save plot
  319. plot_path = self.plots_dir / f"{self.model_name}_speaker_verification.png"
  320. plt.savefig(plot_path, dpi=300, bbox_inches='tight')
  321. plt.close()
  322. return str(plot_path)
  323. def _create_embedding_analysis_plot(self, embedding_metrics: Dict[str, Any]) -> Optional[str]:
  324. """Create embedding quality analysis visualization."""
  325. fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(14, 10))
  326. # Distance comparison
  327. intra_dist = embedding_metrics.get('mean_intra_distance', 0)
  328. inter_dist = embedding_metrics.get('mean_inter_distance', 0)
  329. separability = embedding_metrics.get('separability_ratio', 0)
  330. distances = ['Intra-Speaker', 'Inter-Speaker']
  331. values = [intra_dist, inter_dist]
  332. colors = ['blue', 'red']
  333. bars = ax1.bar(distances, values, color=colors, alpha=0.7)
  334. ax1.set_ylabel('Distance')
  335. ax1.set_title('Speaker Embedding Distances')
  336. ax1.grid(True, alpha=0.3)
  337. # Add value labels on bars
  338. for bar, value in zip(bars, values):
  339. height = bar.get_height()
  340. ax1.text(bar.get_x() + bar.get_width()/2., height + height*0.01,
  341. f'{value:.4f}', ha='center', va='bottom')
  342. # Separability ratio
  343. ax2.bar(['Separability Ratio'], [separability], color='green', alpha=0.7)
  344. ax2.set_ylabel('Ratio (Inter/Intra)')
  345. ax2.set_title(f'Embedding Separability: {separability:.2f}')
  346. ax2.grid(True, alpha=0.3)
  347. # Add separability quality indicator
  348. if separability > 2.0:
  349. quality = "Excellent"
  350. color = 'green'
  351. elif separability > 1.5:
  352. quality = "Good"
  353. color = 'orange'
  354. else:
  355. quality = "Poor"
  356. color = 'red'
  357. ax2.text(0, separability + separability*0.05, quality, ha='center',
  358. fontweight='bold', color=color, fontsize=14)
  359. # Speaker count information
  360. num_speakers = embedding_metrics.get('num_speakers', 0)
  361. ax3.pie([num_speakers, max(1, 10 - num_speakers)], labels=['Trained Speakers', 'Potential'],
  362. autopct='%1.0f', startangle=90, colors=['lightblue', 'lightgray'])
  363. ax3.set_title(f'Speaker Coverage ({num_speakers} speakers)')
  364. # Quality assessment
  365. ax4.axis('off')
  366. # Determine overall quality
  367. if separability > 2.0 and intra_dist < inter_dist:
  368. overall_quality = "EXCELLENT"
  369. quality_color = "lightgreen"
  370. elif separability > 1.5:
  371. overall_quality = "GOOD"
  372. quality_color = "lightyellow"
  373. else:
  374. overall_quality = "NEEDS IMPROVEMENT"
  375. quality_color = "lightcoral"
  376. summary_text = f"""Embedding Quality Assessment:
  377. Overall Quality: {overall_quality}
  378. Metrics:
  379. • Number of Speakers: {num_speakers}
  380. • Intra-Speaker Distance: {intra_dist:.4f}
  381. • Inter-Speaker Distance: {inter_dist:.4f}
  382. • Separability Ratio: {separability:.2f}
  383. Quality Indicators:
  384. • Distance Separation: {'✓' if inter_dist > intra_dist else '✗'}
  385. • Good Separability: {'✓' if separability > 1.5 else '✗'}
  386. • Excellent Separability: {'✓' if separability > 2.0 else '✗'}
  387. """
  388. ax4.text(0.05, 0.95, summary_text, transform=ax4.transAxes, fontsize=11,
  389. verticalalignment='top', bbox=dict(boxstyle="round,pad=0.3", facecolor=quality_color, alpha=0.7))
  390. plt.suptitle(f'{self.model_name} Embedding Quality Analysis', fontsize=14, fontweight='bold')
  391. plt.tight_layout()
  392. # Save plot
  393. plot_path = self.plots_dir / f"{self.model_name}_embedding_analysis.png"
  394. plt.savefig(plot_path, dpi=300, bbox_inches='tight')
  395. plt.close()
  396. return str(plot_path)
  397. def _create_training_summary(self, metrics: TrainingMetrics) -> str:
  398. """Create a text summary of training results."""
  399. total_epochs = len(metrics.metrics['train_loss'])
  400. final_train_loss = metrics.metrics['train_loss'][-1] if metrics.metrics['train_loss'] else 0
  401. final_train_acc = metrics.metrics['train_accuracy'][-1] if metrics.metrics['train_accuracy'] else 0
  402. val_losses = [loss for loss in metrics.metrics['val_loss'] if loss is not None]
  403. final_val_loss = val_losses[-1] if val_losses else 0
  404. val_accs = [acc for acc in metrics.metrics['val_accuracy'] if acc is not None]
  405. final_val_acc = val_accs[-1] if val_accs else 0
  406. total_time = sum(metrics.metrics['epoch_times']) if metrics.metrics['epoch_times'] else 0
  407. avg_epoch_time = total_time / max(1, total_epochs)
  408. summary = f"""Training Summary:
  409. Total Epochs: {total_epochs}
  410. Total Time: {total_time/3600:.2f} hours
  411. Avg Time/Epoch: {avg_epoch_time:.2f} sec
  412. Final Metrics:
  413. Train Loss: {final_train_loss:.6f}
  414. Train Accuracy: {final_train_acc:.4f}
  415. Val Loss: {final_val_loss:.6f}
  416. Val Accuracy: {final_val_acc:.4f}
  417. Best Validation:
  418. Loss: {metrics.best_val_loss:.6f}
  419. Epoch: {metrics.best_epoch}
  420. """
  421. return summary
  422. def create_comparison_plot(self, multiple_metrics: Dict[str, TrainingMetrics],
  423. title: str = "Model Comparison") -> Optional[str]:
  424. """
  425. Create comparison plots for multiple training runs.
  426. Args:
  427. multiple_metrics: Dictionary mapping model names to their metrics
  428. title: Title for the comparison plot
  429. Returns:
  430. Path to generated comparison plot
  431. """
  432. if not MATPLOTLIB_AVAILABLE or not multiple_metrics:
  433. return None
  434. fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(16, 12))
  435. colors = plt.cm.Set1(np.linspace(0, 1, len(multiple_metrics)))
  436. for (model_name, metrics), color in zip(multiple_metrics.items(), colors):
  437. epochs = list(range(len(metrics.metrics['train_loss'])))
  438. # Training loss comparison
  439. ax1.plot(epochs, metrics.metrics['train_loss'],
  440. label=model_name, color=color, linewidth=2)
  441. # Validation loss comparison
  442. if metrics.metrics['val_loss']:
  443. val_epochs = [i for i, loss in enumerate(metrics.metrics['val_loss']) if loss is not None]
  444. val_losses = [loss for loss in metrics.metrics['val_loss'] if loss is not None]
  445. ax2.plot(val_epochs, val_losses,
  446. label=model_name, color=color, linewidth=2, linestyle='--')
  447. # Training accuracy comparison
  448. if metrics.metrics['train_accuracy']:
  449. ax3.plot(epochs, metrics.metrics['train_accuracy'],
  450. label=model_name, color=color, linewidth=2)
  451. # Training time comparison
  452. if metrics.metrics['epoch_times']:
  453. ax4.plot(epochs, np.cumsum(metrics.metrics['epoch_times']) / 3600,
  454. label=model_name, color=color, linewidth=2)
  455. # Configure subplots
  456. ax1.set_xlabel('Epoch')
  457. ax1.set_ylabel('Training Loss')
  458. ax1.set_title('Training Loss Comparison')
  459. ax1.legend()
  460. ax1.grid(True, alpha=0.3)
  461. ax1.set_yscale('log')
  462. ax2.set_xlabel('Epoch')
  463. ax2.set_ylabel('Validation Loss')
  464. ax2.set_title('Validation Loss Comparison')
  465. ax2.legend()
  466. ax2.grid(True, alpha=0.3)
  467. ax2.set_yscale('log')
  468. ax3.set_xlabel('Epoch')
  469. ax3.set_ylabel('Accuracy')
  470. ax3.set_title('Training Accuracy Comparison')
  471. ax3.legend()
  472. ax3.grid(True, alpha=0.3)
  473. ax4.set_xlabel('Epoch')
  474. ax4.set_ylabel('Cumulative Time (hours)')
  475. ax4.set_title('Training Time Comparison')
  476. ax4.legend()
  477. ax4.grid(True, alpha=0.3)
  478. plt.suptitle(title, fontsize=16, fontweight='bold')
  479. plt.tight_layout()
  480. # Save plot
  481. plot_path = self.plots_dir / f"model_comparison_{int(time.time())}.png"
  482. plt.savefig(plot_path, dpi=300, bbox_inches='tight')
  483. plt.close()
  484. return str(plot_path)
  485. def create_training_visualizer(output_dir: str, model_name: str) -> TrainingVisualizer:
  486. """
  487. Factory function to create a training visualizer.
  488. Args:
  489. output_dir: Directory to save plots
  490. model_name: Name of the model
  491. Returns:
  492. TrainingVisualizer instance
  493. """
  494. return TrainingVisualizer(output_dir=output_dir, model_name=model_name)