import torch as pt from torch.utils.data import DataLoader from transformers import AutoModelForImageClassification from augmentations import resize_down_up, decimate, checkerboard_alias_attack from 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): 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}%")