| 226 | ) |
| 227 | |
| 228 | def write_log_stats(epoch, train_stats, test_stats): |
| 229 | if accelerator.is_main_process: |
| 230 | if log_writer is not None: |
| 231 | log_writer.flush() |
| 232 | |
| 233 | log_stats = dict( |
| 234 | epoch=epoch, **{f"train_{k}": v for k, v in train_stats.items()} |
| 235 | ) |
| 236 | for test_name in data_loader_test: |
| 237 | if test_name not in test_stats: |
| 238 | continue |
| 239 | log_stats.update( |
| 240 | {test_name + "_" + k: v for k, v in test_stats[test_name].items()} |
| 241 | ) |
| 242 | |
| 243 | with open( |
| 244 | os.path.join(args.output_dir, "log.txt"), mode="a", encoding="utf-8" |
| 245 | ) as f: |
| 246 | f.write(json.dumps(log_stats) + "\n") |
| 247 | |
| 248 | def save_model(epoch, fname, best_so_far): |
| 249 | misc.save_model( |