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
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
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
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