You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Detectron2每迭代保存训练进度遇断言错误求解决方案

问题:Detectron2每轮迭代保存训练进度并断点恢复的实现问题

请问是否可实现Detectron2模型每轮迭代时保存训练进度,以便中断后从断点恢复训练?我已尝试修改基于Layout Parser的train_net.py中的Trainer类,但保存时均触发断言错误。

修改后的代码如下:

class Trainer(DefaultTrainer):
    """
    We use the "DefaultTrainer" which contains pre-defined default logic for
    standard training workflow. They may not work for you, especially if you
    are working on a new research project. In that case you can use the cleaner
    "SimpleTrainer", or write your own training loop. You can use
    "tools/plain_train_net.py" as an example.

    Adapted from:
        https://github.com/facebookresearch/detectron2/blob/master/projects/DeepLab/train_net.py
    """

    @classmethod
    def build_train_loader(cls, cfg):
        mapper = DatasetMapper(cfg, is_train=True, augmentations=get_augs(cfg))
        return build_detection_train_loader(cfg, mapper=mapper)

    @classmethod
    def build_evaluator(cls, cfg, dataset_name, output_folder=None):
        """
        Returns:
            DatasetEvaluator or None

        It is not implemented by default.
        """
        return COCOEvaluator(dataset_name, cfg, True, output_folder)

    @classmethod
    def test_with_TTA(cls, cfg, model):
        logger = logging.getLogger("detectron2.trainer")
        # In the end of training, run an evaluation with TTA
        # Only support some R-CNN models.
        logger.info("Running inference with test-time augmentation ...")
        model = GeneralizedRCNNWithTTA(cfg, model)
        evaluators = [
            cls.build_evaluator(
                cfg, name, output_folder=os.path.join(cfg.OUTPUT_DIR, "inference_TTA")
            )
            for name in cfg.DATASETS.TEST
        ]
        res = cls.test(cfg, model, evaluators)
        res = OrderedDict({k + "_TTA": v for k, v in res.items()})
        return res

    @classmethod
    def eval_and_save(cls, cfg, model):
        evaluators = [
            cls.build_evaluator(
                cfg, name, output_folder=os.path.join(cfg.OUTPUT_DIR, "inference")
            )
            for name in cfg.DATASETS.TEST
        ]
        res = cls.test(cfg, model, evaluators)
        pd.DataFrame(res).to_csv(os.path.join(cfg.OUTPUT_DIR, "eval.csv"))
        return res
    
    @classmethod
    def save_model(cls, trainer):
        trainer.checkpointer.save(f"{trainer.cfg.OUTPUT_DIR}/model_{trainer.iter}")
        return None

    def run_step(self):
        loss_dict = super().run_step()
        Trainer.save_model(self)
        return loss_dict

报错信息如下:

Traceback (most recent call last):
  File "train_net.py", line 236, in <module>
    launch(
  File "/usr/local/lib/python3.8/dist-packages/detectron2/engine/launch.py", line 62, in launch
    main_func(*args)
  File "train_net.py", line 204, in main
    return trainer.train()
  File "/usr/local/lib/python3.8/dist-packages/detectron2/engine/defaults.py", line 431, in train
    super().train(self.start_iter, self.max_iter)
  File "/usr/local/lib/python3.8/dist-packages/detectron2/engine/train_loop.py", line 138, in train
    self.run_step()
  File "train_net.py", line 124, in run_step
    Trainer.save_model(self)
  File "train_net.py", line 119, in save_model
    trainer.checkpointer.save(f"{trainer.cfg.OUTPUT_DIR}/model_{trainer.iter}")
  File "/usr/local/lib/python3.8/dist-packages/fvcore/common/checkpoint.py", line 111, in save
    assert os.path.basename(save_file) == basename, basename
AssertionError: /content/layout-model-training/outputs/bib/fast_rcnn_R_50_FPN_3x/model_0.pth

问题原因与解决方案

原因分析

断言错误的根源是调用checkpointer.save()时传入了完整路径,但Detectron2的Checkpointer要求仅传入文件名(basename),它会自动将文件保存到cfg.OUTPUT_DIR指定的目录下,无需手动拼接完整路径。

修正后的代码

修改save_model方法,只传入文件名部分:

@classmethod
def save_model(cls, trainer):
    # 仅传入文件名,Checkpointer会自动保存到OUTPUT_DIR目录下
    trainer.checkpointer.save(f"model_{trainer.iter}")
    return None

额外优化建议

  1. 控制保存频率:每轮迭代都保存会生成大量模型文件,占用磁盘空间。建议设置间隔保存,比如每100轮保存一次:
def run_step(self):
    loss_dict = super().run_step()
    # 每100轮保存一次训练进度
    if self.iter % 100 == 0:
        Trainer.save_model(self)
    return loss_dict
  1. 断点恢复训练:Detectron2默认支持断点恢复,只需在重启训练时指定同一个OUTPUT_DIR,Trainer会自动加载目录下最新的checkpoint继续训练,无需额外修改代码。

内容的提问来源于stack exchange,提问作者Sewatech

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.31 21:15:35