def train_step(model: torch.nn.Module, data_loader: torch.utils.data.DataLoader, loss_fn: torch.nn.Module, optimizer: torch.optim.Optimizer, accuracy_fn): # Zero loss and acc train_loss, train_acc = 0, 0 model.train() for batch, (X, y) in enumerate(data_loader): y_pred = model(X) loss = loss_fn(y_pred, y) train_loss += loss.item() train_acc += accuracy_fn(y_true=y, y_pred=y_pred.argmax(dim=1)) optimizer.zero_grad() loss.backward() optimizer.step() train_loss /= len(data_loader) train_acc /= len(data_loader) print(f"Train loss: {train_loss:.2f} | Train accuracy: {train_acc:.2f}%") return train_loss def test_step(data_loader: torch.utils.data.DataLoader, model: torch.nn.Module, loss_fn: torch.nn.Module, accuracy_fn): test_loss, test_acc = 0, 0 model.eval() with torch.no_grad(): for X, y in data_loader: test_pred = model(X) test_loss += loss_fn(test_pred, y).item() test_acc += accuracy_fn(y_true=y, y_pred=test_pred.argmax(dim=1)) test_loss /= len(data_loader) test_acc /= len(data_loader) print(f"Test loss: {test_loss:.5f} | Test accuracy: {test_acc:.2f}%n") return test_loss