Belle II Software light-2609-luna
plot_training Namespace Reference

Functions

 load_checkpoint (path)
 
 plot_loss_curves (ckpt, name, ax_loss, ax_lr)
 
 plot_weight_distributions (ckpt, name, fig)
 
 plot_input_importance (ckpt, name, ax)
 
 summarise (ckpt, name)
 
 main ()
 

Detailed Description

Plot training diagnostics from a ModeSelector checkpoint.

Usage:
    python3 plot_training.py networks/net_category.pt
    python3 plot_training.py networks/net_category.pt networks/net_main.pt
    python3 plot_training.py networks/net_category.pt --output plots/

Function Documentation

◆ load_checkpoint()

load_checkpoint ( path)
Load a PyTorch checkpoint and derive a display name from the filename.

Definition at line 19 of file plot_training.py.

19def load_checkpoint(path):
20 """Load a PyTorch checkpoint and derive a display name from the filename."""
21 ckpt = torch.load(path, map_location='cpu')
22 name = os.path.splitext(os.path.basename(path))[0]
23 return ckpt, name
24
25

◆ main()

main ( )
Parse command-line arguments and produce diagnostic plots for each checkpoint.

Definition at line 136 of file plot_training.py.

136def main():
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()
143
144 if args.output:
145 os.makedirs(args.output, exist_ok=True)
146
147 for path in args.checkpoints:
148 ckpt, name = load_checkpoint(path)
149 summarise(ckpt, name)
150
151 # --- Figure 1: loss curves + LR ---
152 fig1, (ax_loss, ax_lr) = plt.subplots(1, 2, figsize=(12, 4))
153 fig1.suptitle(name)
154 plot_loss_curves(ckpt, name, ax_loss, ax_lr)
155 fig1.tight_layout()
156 if args.output:
157 out = os.path.join(args.output, f'{name}_loss.pdf')
158 fig1.savefig(out, bbox_inches='tight')
159 print(f"Saved {out}")
160
161 # --- Figure 2: weight distributions ---
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))
165 plot_weight_distributions(ckpt, name, fig2)
166 fig2.tight_layout()
167 if args.output:
168 out = os.path.join(args.output, f'{name}_weights.pdf')
169 fig2.savefig(out, bbox_inches='tight')
170 print(f"Saved {out}")
171
172 # --- Figure 3: input importance ---
173 fig3, ax3 = plt.subplots(figsize=(14, 4))
174 plot_input_importance(ckpt, name, ax3)
175 fig3.tight_layout()
176 if args.output:
177 out = os.path.join(args.output, f'{name}_importance.pdf')
178 fig3.savefig(out, bbox_inches='tight')
179 print(f"Saved {out}")
180
181 if not args.output:
182 plt.show()
183 else:
184 plt.close('all')
185
186
Definition main.py:1

◆ plot_input_importance()

plot_input_importance ( ckpt,
name,
ax )
L2 norm of first-layer weights per input feature (top features highlighted).

Definition at line 86 of file plot_training.py.

86def plot_input_importance(ckpt, name, ax):
87 """L2 norm of first-layer weights per input feature (top features highlighted)."""
88 sd = ckpt['model_state_dict']
89 # First layer weight: shape (256, n_inputs)
90 first_weight = None
91 for k, v in sd.items():
92 if 'weight' in k:
93 first_weight = v.numpy()
94 break
95
96 if first_weight is None:
97 ax.set_visible(False)
98 return
99
100 importance = np.linalg.norm(first_weight, axis=0) # L2 norm over output units
101 has_inputs = ckpt.get('has_inputs', list(range(len(importance))))
102
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)
108
109 # Annotate top 10
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)
116
117

◆ plot_loss_curves()

plot_loss_curves ( ckpt,
name,
ax_loss,
ax_lr )
Plot train/val loss and learning rate per epoch.

Definition at line 26 of file plot_training.py.

26def plot_loss_curves(ckpt, name, ax_loss, ax_lr):
27 """Plot train/val loss and learning rate per epoch."""
28 history = ckpt.get('history')
29
30 if history is None:
31 # Show only the best-epoch point when per-epoch history is unavailable.
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)
39 return
40
41 epochs = range(1, len(history['train_loss']) + 1)
42 best_epoch = ckpt['epoch'] + 1 # 0-indexed in checkpoint
43
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')
52 ax_loss.legend()
53 ax_loss.grid(True, alpha=0.3)
54
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)
61
62

◆ plot_weight_distributions()

plot_weight_distributions ( ckpt,
name,
fig )
Histogram of weights per layer.

Definition at line 63 of file plot_training.py.

63def plot_weight_distributions(ckpt, name, fig):
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]
67
68 n = len(weight_layers)
69 axes = fig.subplots(1, n)
70 if n == 1:
71 axes = [axes]
72
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)
78 std = weights.std()
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)
82
83 fig.suptitle(f'{name} -- weight distributions')
84
85

◆ summarise()

summarise ( ckpt,
name )
Print a summary of checkpoint metadata to stdout.

Definition at line 118 of file plot_training.py.

118def summarise(ckpt, name):
119 """Print a summary of checkpoint metadata to stdout."""
120 sd = ckpt['model_state_dict']
121 n_params = sum(v.numel() for v in sd.values())
122 cfg = ckpt.get('config', {})
123 print(f"\n{'='*50}")
124 print(f"Checkpoint: {name}")
125 print(f" Best epoch: {ckpt['epoch'] + 1}")
126 print(f" Best val loss: {ckpt['val_loss']:.6f}")
127 print(f" Train loss: {ckpt['train_loss']:.6f}")
128 print(f" Parameters: {n_params:,}")
129 print(f" Input size: {cfg.get('input_size', '?')}")
130 print(f" Num labels: {cfg.get('num_labels', '?')}")
131 print(f" Network type: {cfg.get('network_type', '?')}")
132 print(f" has_inputs: {len(ckpt.get('has_inputs', []))} features")
133 print(f" History: {'yes' if 'history' in ckpt else 'no'}")
134
135