| """ |
| Training visualization — generates monitoring charts saved as PNGs. |
| |
| Charts: |
| 1. Loss curves (train vs val, per epoch) |
| 2. Per-statistic R² comparison (Model A vs Model B bar chart) |
| 3. Learning rate schedule |
| 4. Predicted vs actual scatter plots (per statistic) |
| |
| All charts are saved to the checkpoint/log directory for download. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import json |
| import logging |
| import os |
| from pathlib import Path |
|
|
| import numpy as np |
|
|
| logger = logging.getLogger(__name__) |
|
|
|
|
| def plot_loss_curves( |
| history_a: list[dict], |
| history_b: list[dict], |
| outdir: str, |
| ) -> str: |
| """Plot train/val loss curves for both models.""" |
| import matplotlib |
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
|
|
| fig, axes = plt.subplots(1, 2, figsize=(14, 5)) |
|
|
| for ax, history, label in [ |
| (axes[0], history_a, "Model A (plain HH)"), |
| (axes[1], history_b, "Model B (HH+ACh)"), |
| ]: |
| epochs = [h["epoch"] for h in history] |
| train_loss = [h["train_loss"] for h in history] |
| val_loss = [h["val_loss"] for h in history] |
|
|
| ax.plot(epochs, train_loss, label="Train", color="#2196F3", alpha=0.8) |
| ax.plot(epochs, val_loss, label="Val", color="#FF9800", alpha=0.8) |
| ax.set_xlabel("Epoch") |
| ax.set_ylabel("MSE Loss") |
| ax.set_title(label) |
| ax.legend() |
| ax.grid(True, alpha=0.3) |
|
|
| |
| best_idx = int(np.argmin(val_loss)) |
| ax.axvline(epochs[best_idx], color="green", linestyle="--", alpha=0.5, |
| label=f"Best: epoch {epochs[best_idx]}") |
| ax.legend() |
|
|
| fig.suptitle("Training Loss Curves", fontsize=14, fontweight="bold") |
| fig.tight_layout() |
|
|
| path = os.path.join(outdir, "loss_curves.png") |
| fig.savefig(path, dpi=150, bbox_inches="tight") |
| plt.close(fig) |
| logger.info(f"Saved loss curves to {path}") |
| return path |
|
|
|
|
| def plot_r2_comparison( |
| r2_a: dict[str, float], |
| r2_b: dict[str, float], |
| outdir: str, |
| ) -> str: |
| """Bar chart comparing per-statistic R² for Model A vs Model B.""" |
| import matplotlib |
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
|
|
| stats = list(r2_a.keys()) |
| vals_a = [r2_a[s] for s in stats] |
| vals_b = [r2_b[s] for s in stats] |
|
|
| x = np.arange(len(stats)) |
| width = 0.35 |
|
|
| fig, ax = plt.subplots(figsize=(14, 6)) |
| bars_a = ax.bar(x - width/2, vals_a, width, label="Model A (plain HH)", |
| color="#2196F3", alpha=0.8) |
| bars_b = ax.bar(x + width/2, vals_b, width, label="Model B (HH+ACh)", |
| color="#FF9800", alpha=0.8) |
|
|
| ax.set_ylabel("R² (validation set)") |
| ax.set_title("Per-Statistic R² — Model A vs Model B", fontsize=14, fontweight="bold") |
| ax.set_xticks(x) |
| ax.set_xticklabels([s.replace("_", "\n") for s in stats], fontsize=8, rotation=45, ha="right") |
| ax.legend() |
| ax.grid(True, axis="y", alpha=0.3) |
| ax.set_ylim(-0.1, 1.05) |
| ax.axhline(0, color="black", linewidth=0.5) |
|
|
| |
| for bar in bars_a: |
| h = bar.get_height() |
| if h > 0: |
| ax.text(bar.get_x() + bar.get_width()/2., h + 0.02, |
| f"{h:.2f}", ha="center", va="bottom", fontsize=7, color="#1565C0") |
| for bar in bars_b: |
| h = bar.get_height() |
| if h > 0: |
| ax.text(bar.get_x() + bar.get_width()/2., h + 0.02, |
| f"{h:.2f}", ha="center", va="bottom", fontsize=7, color="#E65100") |
|
|
| |
| mean_a = np.mean(vals_a) |
| mean_b = np.mean(vals_b) |
| ax.text(0.98, 0.95, f"Mean R²: A={mean_a:.3f}, B={mean_b:.3f}", |
| transform=ax.transAxes, ha="right", va="top", |
| bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5)) |
|
|
| fig.tight_layout() |
| path = os.path.join(outdir, "r2_comparison.png") |
| fig.savefig(path, dpi=150, bbox_inches="tight") |
| plt.close(fig) |
| logger.info(f"Saved R² comparison to {path}") |
| return path |
|
|
|
|
| def plot_lr_schedule( |
| history: list[dict], |
| outdir: str, |
| model_label: str = "", |
| ) -> str: |
| """Plot learning rate over training.""" |
| import matplotlib |
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
|
|
| epochs = [h["epoch"] for h in history] |
| lrs = [h["lr"] for h in history] |
|
|
| fig, ax = plt.subplots(figsize=(8, 4)) |
| ax.plot(epochs, lrs, color="#4CAF50", linewidth=2) |
| ax.set_xlabel("Epoch") |
| ax.set_ylabel("Learning Rate") |
| ax.set_title(f"LR Schedule {model_label}", fontsize=12) |
| ax.grid(True, alpha=0.3) |
| ax.set_yscale("log") |
|
|
| fig.tight_layout() |
| path = os.path.join(outdir, f"lr_schedule{'_' + model_label.lower().replace(' ', '_') if model_label else ''}.png") |
| fig.savefig(path, dpi=150, bbox_inches="tight") |
| plt.close(fig) |
| logger.info(f"Saved LR schedule to {path}") |
| return path |
|
|
|
|
| def plot_val_r2_over_time( |
| history: list[dict], |
| outdir: str, |
| model_label: str = "", |
| ) -> str: |
| """Plot mean validation R² over epochs.""" |
| import matplotlib |
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
|
|
| epochs = [h["epoch"] for h in history] |
| r2s = [h["mean_val_r2"] for h in history] |
|
|
| fig, ax = plt.subplots(figsize=(8, 4)) |
| ax.plot(epochs, r2s, color="#9C27B0", linewidth=2) |
| ax.set_xlabel("Epoch") |
| ax.set_ylabel("Mean Val R²") |
| ax.set_title(f"Validation R² Over Training {model_label}", fontsize=12) |
| ax.grid(True, alpha=0.3) |
| ax.set_ylim(-0.1, 1.05) |
|
|
| fig.tight_layout() |
| suffix = f"_{'_'.join(model_label.lower().split())}" if model_label else "" |
| path = os.path.join(outdir, f"val_r2_over_time{suffix}.png") |
| fig.savefig(path, dpi=150, bbox_inches="tight") |
| plt.close(fig) |
| logger.info(f"Saved val R² curve to {path}") |
| return path |
|
|
|
|
| def generate_all_plots( |
| results_a: dict, |
| results_b: dict, |
| log_dir: str, |
| ) -> list[str]: |
| """Generate all training monitoring plots. |
| |
| Args: |
| results_a: Results dict from train_one_model("A") |
| results_b: Results dict from train_one_model("B") |
| log_dir: Directory containing history JSONs and for saving plots |
| |
| Returns: |
| List of saved plot paths. |
| """ |
| plot_dir = os.path.join(log_dir, "plots") |
| os.makedirs(plot_dir, exist_ok=True) |
|
|
| saved = [] |
|
|
| |
| hist_a_path = os.path.join(log_dir, "history_model_a.json") |
| hist_b_path = os.path.join(log_dir, "history_model_b.json") |
|
|
| if os.path.exists(hist_a_path) and os.path.exists(hist_b_path): |
| with open(hist_a_path) as f: |
| history_a = json.load(f) |
| with open(hist_b_path) as f: |
| history_b = json.load(f) |
|
|
| |
| saved.append(plot_loss_curves(history_a, history_b, plot_dir)) |
|
|
| |
| r2_a = results_a.get("final_val_r2", {}) |
| r2_b = results_b.get("final_val_r2", {}) |
| if r2_a and r2_b: |
| saved.append(plot_r2_comparison(r2_a, r2_b, plot_dir)) |
|
|
| |
| saved.append(plot_lr_schedule(history_a, plot_dir, "Model A")) |
| saved.append(plot_lr_schedule(history_b, plot_dir, "Model B")) |
|
|
| |
| saved.append(plot_val_r2_over_time(history_a, plot_dir, "Model A")) |
| saved.append(plot_val_r2_over_time(history_b, plot_dir, "Model B")) |
|
|
| logger.info(f"Generated {len(saved)} plots in {plot_dir}") |
| return saved |
|
|