File size: 1,316 Bytes
e4ac40d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 | import tqdm
from src.process.metrics import *
import torch
@torch.no_grad()
def validate_one_epoch(model, loader, criterion, device, num_classes) -> dict[str, float]:
model.eval()
losses = []
accuracies = []
precisions = []
recalls = []
f1s = []
f1s_weighted = []
for images, labels in tqdm.tqdm(loader, leave=False):
images = images.to(device, non_blocking=True).float()
labels = labels.to(device, non_blocking=True).long()
logits = model(images)
loss = criterion(logits, labels)
m = metrics(logits, labels, num_classes)
losses.append(loss.detach())
accuracies.append(m["accuracy"])
precisions.append(m["precision"])
recalls.append(m["recall"])
f1s.append(m["f1"])
f1s_weighted.append(m["f1_w"])
losses = torch.stack(losses)
accuracies = torch.stack(accuracies)
precisions = torch.stack(precisions)
recalls = torch.stack(recalls)
f1s = torch.stack(f1s)
f1s_weighted = torch.stack(f1s_weighted)
return {
"loss": losses.mean().item(),
"accuracy": accuracies.mean().item(),
"precision": precisions.mean().item(),
"recall": recalls.mean().item(),
"f1": f1s.mean().item(),
"f1_w": f1s_weighted.mean().item()
}
|