Save train stats
(train_losses, train_accs, val_losses, val_accs, args)
| 157 | |
| 158 | |
| 159 | def save_train_stats(train_losses, train_accs, val_losses, val_accs, args): |
| 160 | """Save train stats""" |
| 161 | |
| 162 | model_save_dir = constants.get_model_dir(args.dataset, args.model_type, args.model_name) |
| 163 | train_stats = { |
| 164 | "train_losses": train_losses, |
| 165 | "train_accs": train_accs, |
| 166 | "val_losses": val_losses, |
| 167 | "val_accs": val_accs, |
| 168 | "model_type": args.model_type, |
| 169 | "learning_rate": args.lr, |
| 170 | "optimizer": args.optimizer, |
| 171 | } |
| 172 | train_stats_save_path = os.path.join(model_save_dir, "train_stats.json") |
| 173 | print(f"Saving current train stats to {train_stats_save_path}...") |
| 174 | with open(train_stats_save_path, "w") as f: |
| 175 | json.dump(train_stats, f) |
| 176 | |
| 177 | |
| 178 | def main(): |