PaySim-Fraud / src /models /evaluate.py
Nickanas's picture
Upload 24 files
71c85a4 verified
Raw History Blame Contribute Delete
3.11 kB
import torch
from sklearn.metrics import (
average_precision_score, precision_score, recall_score,
f1_score, precision_recall_curve,
)
def predict_one_epoch(model, dataloader):
"""Evaluate a trained model on a dataloader. Self-contained:
derives device from the model and builds a pos_weight-matched loss
from the data it's given.
"""
device = next(model.parameters()).device # take device from the model itself
model.eval()
# pos_weight from THIS dataloader, so eval loss matches the weighted train loss
y_all = torch.cat([y for _, y in dataloader])
n_pos = y_all.sum().item()
pos_weight = torch.tensor([(y_all.numel() - n_pos) / max(n_pos, 1)], device=device)
loss_fn = torch.nn.BCEWithLogitsLoss(pos_weight=pos_weight)
test_loss, n_batches = 0.0, 0
all_probs, all_y = [], []
with torch.inference_mode():
for X, y in dataloader:
X, y = X.to(device), y.to(device)
logits = model(X)
test_loss += loss_fn(logits, y).item()
n_batches += 1
all_probs.append(torch.sigmoid(logits).cpu())
all_y.append(y.cpu())
probs = torch.cat(all_probs).numpy().ravel()
y_true = torch.cat(all_y).numpy().ravel()
preds = (probs > 0.5).astype(int)
return {
"loss": test_loss / n_batches,
"acc": (preds == y_true).mean(),
"pr_auc": average_precision_score(y_true, probs),
"precision": precision_score(y_true, preds, zero_division=0),
"recall": recall_score(y_true, preds, zero_division=0),
}
def get_probs(model, dataloader):
"""Run the model over a dataloader, return raw (probs, y_true) arrays.
Threshold-free.This is used to score any cutoff afterwards"""
device = next(model.parameters()).device
model.eval()
all_probs, all_y = [], []
with torch.inference_mode():
for X, y in dataloader:
logits = model(X.to(device))
all_probs.append(torch.sigmoid(logits).cpu())
all_y.append(y.cpu())
probs = torch.cat(all_probs).numpy().ravel()
y_true = torch.cat(all_y).numpy().ravel()
return probs, y_true
def best_threshold(probs, y_true):
"""Scan all cutoffs, return the one that maximises F1."""
precisions, recalls, thresholds = precision_recall_curve(y_true, probs)
# curve returns 1 more precision/recall than thresholds -> drop last
f1s = 2 * precisions[:-1] * recalls[:-1] / (precisions[:-1] + recalls[:-1] + 1e-9)
best_idx = f1s.argmax()
return thresholds[best_idx]
def metrics_at(probs, y_true, threshold):
"""Compute metrics using a specific probability cutoff."""
preds = (probs >= threshold).astype(int)
return {
"pr_auc": average_precision_score(y_true, probs), # threshold-free
"precision": precision_score(y_true, preds, zero_division=0),
"recall": recall_score(y_true, preds, zero_division=0),
"f1": f1_score(y_true, preds, zero_division=0),
"threshold": threshold,
}