Is it possible to execute from the point where the neural network model is interrupted?

Viewed 49

Assume that I am training a neural network model. I am storing the tensor file of the neural network model for every 15 epochs in .pth format.

I need to run 1000 epochs in total. Suppose I stopped my program during the 501st epoch, then I have the following files

15.pth, 30.pth, 45.pth, 60.pth, 75.pth,.... 420.pth, 435.pth, 450.pth, 465.pth, 480.pth, 495.pth

Then my doubt is

Is it possible to use the last stored model 495.pth and continue execution as it generally happens if done without any interruption? In short, I am asking for something similar to the "resumption" of the training phase with a few modifications to the existing code. I am just asking for such a possibility.

I am asking for general practice and not particular to any code. If such a method exists, I will be free to stop any program under execution and can resume later. Currently, I cannot use resources for shorter programs if longer programs are in execution and hence I am asking this question.

2 Answers

I order to resume training from a checkpoint, you need to save the entire state of your training process. This includes:

  1. Current weights of the model.
  2. State of the optimizer: most optimizers keep track of different statistics of the updates, e.g., momentum, variance etc.
  3. State of the learning rate scheduler.
  4. Additional "state" variables unique to your code.

If you saved all this information, you should be able to fully restore the "state" of your training process and resume from that point.

So what I do is the following: After each epoch I save my models weights into a .pt file and each time I run my program in gerneral I check if the resume argument is set to True. If so, I initialize the model using the weights in the .pt file as just continue training, if not I initialize random weights as normal. This could look like this:

def train(resume: bool=False):
    model = Model()
    if resume:
        model.load_state_dict(torch.load("weights.pt"))
   
    criterion = Loss()
    optimizer = Optimizer()

    for epoch in range(100):
        for data, targets in dataloader:
            optimizer.zero_grad()

            predictions = model.train()(data)
            loss = criterion(predicitions, targets)

            loss.backward()
            optimizer.step()

        torch.save(model.state_dict(), "weights.pt")

So if I interrupt the training, I can still continue after my last epoch that I saved. Normally you are logging more stuff than only the weights, for example the learning-rate scheduler or simply the loss and accuracy history. For that you could save the training history into a json file and read it out if resume is True.

Related