CoolFace
Apppublic

glisicstefan/age-estimation-resnet50

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
engine.py77 linesDownload Raw Back to src
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