loss_hist = {} loss_hist['train'] = {} loss_hist['test'] = {} epochs = 10 for epoch in tqdm(range(epochs)): print(f"Epoch: {epoch}n---------") train_loss = train_step(data_loader=train_loader, model=model_non_linear, loss_fn=loss_function, optimizer=optimizer, accuracy_fn=accuracy_fn ) loss_hist['train'][epoch] = train_loss test_loss = test_step(data_loader=test_loader, model=model_non_linear, loss_fn=loss_function, accuracy_fn=accuracy_fn ) loss_hist['test'][epoch] = test_loss