nula-cifar10-robust-v0 / evaluate.py
MamaPearl's picture
Create evaluate.py
e935244 verified
Raw
History Blame
2.98 kB
import torch as pt
import torch.nn.functional as F
from torch.utils.data import DataLoader
from transformers import AutoModelForImageClassification
from src.augmentations import resize_down_up, decimate, checkerboard_alias_attack
from src.dataset import get_device
import torchvision.transforms as T
from PIL import Image
test_transform = T.Compose([
T.ToTensor(),
T.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
])
def nula_collate_fn(batch):
pixel_values = pt.stack([test_transform(item["img"].convert("RGB")) for item in batch])
labels = pt.tensor([item["label"] for item in batch])
return {"pixel_values": pixel_values, "labels": labels}
@pt.no_grad()
def evaluate_clean(model, loader, device):
model.eval()
total_correct, total_examples = 0, 0
for batch in loader:
x = batch["pixel_values"].to(device, non_blocking=True)
y = batch["labels"].to(device, non_blocking=True)
preds = model(pixel_values=x).logits.argmax(dim=1)
total_correct += (preds == y).sum().item()
total_examples += y.size(0)
return total_correct / total_examples
@pt.no_grad()
def evaluate_under_transform(model, loader, device, transform_fn):
model.eval()
total_correct, total_examples = 0, 0
for batch in loader:
x = batch["pixel_values"].to(device, non_blocking=True)
y = batch["labels"].to(device, non_blocking=True)
preds = model(pixel_values=transform_fn(x)).logits.argmax(dim=1)
total_correct += (preds == y).sum().item()
total_examples += y.size(0)
return total_correct / total_examples
def report_stress_suite(model, loader, device):
blur = pt.nn.Sequential() # placeholder — BlurPool handled inside augmentations
return {
"clean": evaluate_clean(model, loader, device),
"resize_0.5": evaluate_under_transform(model, loader, device, lambda x: resize_down_up(x, scale=0.5)),
"resize_0.25": evaluate_under_transform(model, loader, device, lambda x: resize_down_up(x, scale=0.25)),
"decimate_x2": evaluate_under_transform(model, loader, device, lambda x: decimate(x, factor=2)),
"checker_0.03": evaluate_under_transform(model, loader, device, lambda x: checkerboard_alias_attack(x, epsilon=0.03)),
"checker_0.05": evaluate_under_transform(model, loader, device, lambda x: checkerboard_alias_attack(x, epsilon=0.05)),
}
if __name__ == "__main__":
from datasets import load_dataset
DEVICE = get_device()
dataset = load_dataset("uoft-cs/cifar10")
test_loader = DataLoader(dataset["test"], batch_size=128, shuffle=False,
num_workers=0, collate_fn=nula_collate_fn)
model = AutoModelForImageClassification.from_pretrained(
"./nula-best-model", trust_remote_code=True
).to(DEVICE)
results = report_stress_suite(model, test_loader, DEVICE)
for k, v in results.items():
print(f"{k:20} {100*v:.2f}%")