27 """Plot train/val loss and learning rate per epoch."""
28 history = ckpt.get(
'history')
32 best_epoch = ckpt[
'epoch']
33 ax_loss.scatter([best_epoch + 1], [ckpt[
'train_loss']], marker=
'o',
34 label=
'train (best epoch only)')
35 ax_loss.scatter([best_epoch + 1], [ckpt[
'val_loss']], marker=
's',
36 label=
'val (best epoch only)')
37 ax_loss.set_title(f
'{name} -- no loss history in checkpoint')
38 ax_lr.set_visible(
False)
41 epochs = range(1, len(history[
'train_loss']) + 1)
42 best_epoch = ckpt[
'epoch'] + 1
44 ax_loss.plot(epochs, history[
'train_loss'], label=
'train')
45 ax_loss.plot(epochs, history[
'val_loss'], label=
'val')
46 if any(x > 0
for x
in history[
'disco_loss']):
47 ax_loss.plot(epochs, history[
'disco_loss'], label=
'disco', linestyle=
'--', alpha=0.7)
48 ax_loss.axvline(best_epoch, color=
'gray', linestyle=
':', alpha=0.8, label=f
'best (ep {best_epoch})')
49 ax_loss.set_xlabel(
'epoch')
50 ax_loss.set_ylabel(
'loss')
51 ax_loss.set_title(f
'{name} -- loss curves')
53 ax_loss.grid(
True, alpha=0.3)
55 ax_lr.plot(epochs, history[
'lr'])
56 ax_lr.set_xlabel(
'epoch')
57 ax_lr.set_ylabel(
'learning rate')
58 ax_lr.set_title(
'learning rate schedule')
59 ax_lr.set_yscale(
'log')
60 ax_lr.grid(
True, alpha=0.3)
64 """Histogram of weights per layer."""
65 sd = ckpt[
'model_state_dict']
66 weight_layers = [(k, v.numpy().ravel())
for k, v
in sd.items()
if 'weight' in k]
68 n = len(weight_layers)
69 axes = fig.subplots(1, n)
73 for ax, (layer_name, weights)
in zip(axes, weight_layers):
74 ax.hist(weights, bins=60, density=
True)
75 ax.set_title(layer_name.replace(
'network.',
'layer '), fontsize=8)
76 ax.set_xlabel(
'weight value', fontsize=7)
77 ax.tick_params(labelsize=7)
79 ax.axvline(0, color=
'k', linewidth=0.5)
80 ax.text(0.97, 0.97, f
'std={std:.3f}', transform=ax.transAxes,
81 ha=
'right', va=
'top', fontsize=7)
83 fig.suptitle(f
'{name} -- weight distributions')
87 """L2 norm of first-layer weights per input feature (top features highlighted)."""
88 sd = ckpt[
'model_state_dict']
91 for k, v
in sd.items():
93 first_weight = v.numpy()
96 if first_weight
is None:
100 importance = np.linalg.norm(first_weight, axis=0)
101 has_inputs = ckpt.get(
'has_inputs', list(range(len(importance))))
103 ax.bar(range(len(importance)), importance, width=1.0, linewidth=0)
104 ax.set_xlabel(
'feature index (within selected features)')
105 ax.set_ylabel(
'L2 norm of input weights')
106 ax.set_title(f
'{name} -- input feature importance (first layer)')
107 ax.grid(
True, axis=
'y', alpha=0.3)
110 top10 = np.argsort(importance)[-10:][::-1]
111 for rank, idx
in enumerate(top10):
112 global_idx = has_inputs[idx]
if idx < len(has_inputs)
else idx
113 ax.annotate(f
'{global_idx}', xy=(idx, importance[idx]),
114 xytext=(0, 3), textcoords=
'offset points',
115 ha=
'center', fontsize=6, rotation=90)
137 """Parse command-line arguments and produce diagnostic plots for each checkpoint."""
138 parser = argparse.ArgumentParser(description=
'Plot ModeSelector training diagnostics')
139 parser.add_argument(
'checkpoints', nargs=
'+', help=
'Path(s) to .pt checkpoint file(s)')
140 parser.add_argument(
'--output', default=
None,
141 help=
'Output directory for plots (default: show interactively)')
142 args = parser.parse_args()
145 os.makedirs(args.output, exist_ok=
True)
147 for path
in args.checkpoints:
152 fig1, (ax_loss, ax_lr) = plt.subplots(1, 2, figsize=(12, 4))
157 out = os.path.join(args.output, f
'{name}_loss.pdf')
158 fig1.savefig(out, bbox_inches=
'tight')
159 print(f
"Saved {out}")
162 sd = ckpt[
'model_state_dict']
163 n_layers = sum(1
for k
in sd
if 'weight' in k)
164 fig2 = plt.figure(figsize=(3 * n_layers, 3))
168 out = os.path.join(args.output, f
'{name}_weights.pdf')
169 fig2.savefig(out, bbox_inches=
'tight')
170 print(f
"Saved {out}")
173 fig3, ax3 = plt.subplots(figsize=(14, 4))
177 out = os.path.join(args.output, f
'{name}_importance.pdf')
178 fig3.savefig(out, bbox_inches=
'tight')
179 print(f
"Saved {out}")