File size: 2,980 Bytes
e935244
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
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}%")