From 09e17986975a4264f9f9b719beb1a9304530fd31 Mon Sep 17 00:00:00 2001 From: shiyi20060618-cmd <202430841089@mail.scut.edu.cn> Date: Wed, 1 Jul 2026 17:54:58 +0800 Subject: [PATCH] Add configurable checkpoint outputs --- main.py | 16 +++++++++++----- 1 file changed, 11 insertions(+), 5 deletions(-) diff --git a/main.py b/main.py index 129edaeba..20343141d 100644 --- a/main.py +++ b/main.py @@ -63,6 +63,10 @@ def get_argparser(): parser.add_argument("--ckpt", default=None, type=str, help="restore from checkpoint") + parser.add_argument("--ckpt_dir", default="checkpoints", type=str, + help="checkpoint output directory") + parser.add_argument("--step_ckpt_interval", type=int, default=0, + help="save step_XXXXXX.pth every N iterations; 0 disables it") parser.add_argument("--continue_training", action='store_true', default=False) parser.add_argument("--loss_type", type=str, default='cross_entropy', @@ -283,7 +287,7 @@ def save_ckpt(path): }, path) print("Model saved as %s" % path) - utils.mkdir('checkpoints') + utils.mkdir(opts.ckpt_dir) # Restore best_score = 0.0 cur_itrs = 0 @@ -348,8 +352,10 @@ def save_ckpt(path): interval_loss = 0.0 if (cur_itrs) % opts.val_interval == 0: - save_ckpt('checkpoints/latest_%s_%s_os%d.pth' % - (opts.model, opts.dataset, opts.output_stride)) + save_ckpt(os.path.join(opts.ckpt_dir, 'latest_%s_%s_os%d.pth' % + (opts.model, opts.dataset, opts.output_stride))) + if opts.step_ckpt_interval > 0 and cur_itrs % opts.step_ckpt_interval == 0: + save_ckpt(os.path.join(opts.ckpt_dir, 'step_%06d.pth' % cur_itrs)) print("validation...") model.eval() val_score, ret_samples = validate( @@ -358,8 +364,8 @@ def save_ckpt(path): print(metrics.to_str(val_score)) if val_score['Mean IoU'] > best_score: # save best model best_score = val_score['Mean IoU'] - save_ckpt('checkpoints/best_%s_%s_os%d.pth' % - (opts.model, opts.dataset, opts.output_stride)) + save_ckpt(os.path.join(opts.ckpt_dir, 'best_%s_%s_os%d.pth' % + (opts.model, opts.dataset, opts.output_stride))) if vis is not None: # visualize validation score and samples vis.vis_scalar("[Val] Overall Acc", cur_itrs, val_score['Overall Acc'])