import torch from torch import nn from torch.utils.data import DataLoader, TensorDataset def evaluate(model, data_loader, loss_fn, device): model.eval() loss_total = 0.0 correct = 0 sample_count = 0 with torch.inference_mode(): for inputs, targets in data_loader: inputs = inputs.to(device) targets = targets.to(device) logits = model(inputs) batch_size = targets.size(0) loss_total += loss_fn(logits, targets).item() * batch_size correct += (logits.argmax(dim=1) == targets).sum().item() sample_count += batch_size return { "loss": loss_total / sample_count, "accuracy": correct / sample_count, "correct": correct, "total": sample_count, } device = torch.device("cpu") features = torch.tensor( [ [0.0, 0.1], [0.2, 0.0], [1.0, 1.1], [1.2, 0.9], [0.1, 0.2], ], dtype=torch.float32, ) targets = torch.tensor([0, 0, 1, 1, 0]) validation_loader = DataLoader( TensorDataset(features, targets), batch_size=2, shuffle=False, ) model = nn.Sequential( nn.Linear(2, 2), nn.Dropout(p=0.75), ).to(device) with torch.no_grad(): model[0].weight.copy_(torch.tensor([[1.0, -1.0], [1.0, 1.0]])) model[0].bias.copy_(torch.tensor([0.25, -1.0])) loss_fn = nn.CrossEntropyLoss() metrics = evaluate(model, validation_loader, loss_fn, device) with torch.inference_mode(): full_logits = model(features.to(device)) full_loss = loss_fn(full_logits, targets.to(device)).item() full_correct = (full_logits.argmax(dim=1) == targets.to(device)).sum().item() metrics_match = ( abs(metrics["loss"] - full_loss) < 1e-7 and metrics["correct"] == full_correct ) print(f"model training mode: {model.training}") print(f"evaluated samples: {metrics['total']}") print(f"validation loss: {metrics['loss']:.4f}") print( f"validation accuracy: {metrics['correct']}/{metrics['total']} " f"({metrics['accuracy']:.0%})" ) print(f"full-batch metrics match: {metrics_match}")