PyTorch model saving error: "Can't pickle local object"

Viewed 8692

When I try to save the PyTorch model with this piece of code:

checkpoint = {'model': Net(), 'state_dict': model.state_dict(),'optimizer' :optimizer.state_dict()}
torch.save(checkpoint, 'Checkpoint.pth')

I get the following error:

    E:\PROGRAM FILES\Anaconda\envs\staj_projesi\lib\site-packages\torch\serialization.py:251: UserWarning: Couldn't retrieve source code for container of type Net. It won't be checked for correctness upon loading.
...

      "type " + obj.__name__ + ". It won't be checked "
    Can't pickle local object 'trainModel.<locals>.Net'

When I try to save the PyTorch model with this piece of code:

checkpoint = {'state_dict': model.state_dict(),'optimizer' :optimizer.state_dict()}
torch.save(checkpoint, 'Checkpoint.pth')

I don't don't get any errors, but I want to save the ANN class. How can I solve this problem? Also, I could save the model with the first structure in the other projects before

2 Answers

You can't! torch.save is saving the objects state_dict() only.

When you use the following:

checkpoint = {'model': Net(), 'state_dict': model.state_dict(),'optimizer' :optimizer.state_dict()}
torch.save(checkpoint, 'Checkpoint.pth')

You are trying to save the model itself, but this data is saved in the model.state_dict() and when loading a model with the state_dict you should first initiate a model object.

This is exactly the reason why the second method works properly:

checkpoint = {'state_dict': model.state_dict(),'optimizer' :optimizer.state_dict()}
torch.save(checkpoint, 'Checkpoint.pth')

I would suggest reading the pytorch docs of how to properly save\load a model in the following link: https://pytorch.org/tutorials/beginner/saving_loading_models.html

Do the usual proper way to save and load models https://pytorch.org/tutorials/beginner/saving_loading_models.html and if you have args or dicts you want to save and perhaps a lambda function sometimes I use dill and the errors go away. e.g.

def save_for_meta_learning(args, ckpt_filename='ckpt.pt'):
    if is_lead_worker(args.rank):
        import dill
        args.logger.save_current_plots_and_stats()
        # - ckpt
        assert uutils.xor(args.training_mode == 'epochs', args.training_mode == 'iterations')
        args_pickable = uutils.make_args_pickable(args)
        # args.meta_learner.args = args_pickable
        f: nn.Module = get_model_from_ddp(args.base_model)
        # pickle vs torch_uu.save https://discuss.pytorch.org/t/advantages-disadvantages-of-using-pickle-module-to-save-models-vs-torch-save/79016
        torch.save({'training_mode': args.training_mode,  # its or epochs
                    'it': args.it,
                    'epoch_num': args.epoch_num,
                    # 'args': args_pickable,
                    'args_pickable': args_pickable,
                    # 'meta_learner': args.meta_learner,
                    'meta_learner_str': str(args.meta_learner),
                    # 'f': f,
                    'f_state_dict': f.state_dict(),
                    'f_str': str(f),
                    # 'f_modules': f._modules,
                    # 'f_modules_str': str(f._modules),
                    'outer_opt_state_dict': args.outer_opt.state_dict()
                    },
                   pickle_module=dill,
                   f=args.log_root / ckpt_filename)
Related