import torch from torch import nn torch.manual_seed(7) checkpoint_path = "training-checkpoint.tar" def build_model(): return nn.Sequential( nn.Linear(3, 4), nn.ReLU(), nn.Linear(4, 1), ) x = torch.tensor( [ [0.1, 0.2, 0.3], [0.4, 0.5, 0.6], ], dtype=torch.float32, ) y = torch.tensor( [ [0.6], [1.5], ], dtype=torch.float32, ) loss_fn = nn.MSELoss() def train_step(model, optimizer): model.train() optimizer.zero_grad() loss = loss_fn(model(x), y) loss.backward() optimizer.step() return loss.item()