glisicstefan/age-estimation-resnet50
0
1import torch2from tqdm.auto import tqdm3 4def train_step(model, dataloader, loss_fn, optimizer, mae, device):5 # Put model in train mode6 model.train()7 # Set train loss value8 train_loss = 09 # Loop through dataloader and batches10 for batch, (X, y) in enumerate(dataloader):11 # Send data to target device12 X, y = X.to(device), y.unsqueeze(1).to(device)13 # Forward pass14 y_pred = model(X)15 # Caluculate and accumulate loss16 loss = loss_fn(y_pred, y)17 train_loss += loss.item()18 # Calculate MAE19 mae(y_pred, y)20 # Optim zero grad21 optimizer.zero_grad()22 # Loss backward23 loss.backward()24 # Optim step25 optimizer.step()26 # Adjust metrics to get avg loss per batch 27 train_loss /= len(dataloader)28 train_mae = mae.compute()29 mae.reset()30 return train_loss, train_mae31 32 33def test_step(model, dataloader, loss_fn, mae, device):34 test_loss = 035 model.eval()36 with torch.inference_mode():37 for X, y in dataloader:38 X, y = X.to(device), y.unsqueeze(1).to(device)39 test_pred_logits = model(X)40 loss = loss_fn(test_pred_logits, y)41 test_loss += loss.item()42 mae(test_pred_logits, y)43 test_loss /= len(dataloader)44 test_mae = mae.compute()45 mae.reset()46 return test_loss, test_mae47 48 49def train(model, train_dataloader, test_dataloader, loss_fn, optimizer, mae, device, epochs=5):50 51 results = {"train_loss": [], "train_mae": [], "test_loss": [], "test_mae": []}52 53 for epoch in tqdm(range(epochs)):54 train_loss, train_mae = train_step(model=model,55 dataloader=train_dataloader,56 loss_fn=loss_fn,57 optimizer=optimizer,58 mae=mae,59 device=device)60 test_loss, test_mae = test_step(model=model,61 dataloader=test_dataloader,62 loss_fn=loss_fn,63 mae=mae,64 device=device)65 66 67 print(f"Epoch: {epoch+1} | "68 f"Train Loss: {train_loss:.4f} | Train MAE: {train_mae:.2f} | "69 f"Test Loss: {test_loss:.4f} | Test MAE: {test_mae:.2f}")70 71 72 results["train_loss"].append(train_loss)73 results["train_mae"].append(train_mae.item() if hasattr(train_mae, "item") else train_mae)74 results["test_loss"].append(test_loss)75 results["test_mae"].append(test_mae.item() if hasattr(test_mae, "item") else test_mae)76 77 return results