Code
pytorch
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
def train(model, loader, epochs=5, lr=1e-3, device="cpu"):
model.to(device).train()
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=lr)
for epoch in range(epochs):
total = 0.0
for xb, yb in loader:
xb, yb = xb.to(device), yb.to(device)
optimizer.zero_grad()
logits = model(xb)
loss = criterion(logits, yb)
loss.backward()
optimizer.step()
total += loss.item() * xb.size(0)
print(f"epoch {epoch} loss={total / len(loader.dataset):.4f}")
@torch.no_grad()
def evaluate(model, loader, device="cpu"):
model.to(device).eval()
correct = 0
for xb, yb in loader:
xb, yb = xb.to(device), yb.to(device)
preds = model(xb).argmax(1)
correct += (preds == yb).sum().item()
return correct / len(loader.dataset)