File size: 3,131 Bytes
1ea7ba6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
70
71
72
73
74
75
76
77
78
79
80
81
82
"""

Cross-source 2x2 matrix for a timm backbone (both source checkpoints).



    python eval_backbone_shortcut.py --model resnet50 \

        --ss_ckpt outputs/checkpoints/bench_resnet50_ss/best.pth \

        --in_ckpt outputs/checkpoints/bench_resnet50_indian/best.pth

"""
import sys, os, argparse, json
from pathlib import Path

_base = "/mnt/d/SpiceNet" if os.path.exists("/mnt/d/SpiceNet") else "D:/SpiceNet"
sys.path.insert(0, _base)

import numpy as np
import torch
from torch.utils.data import DataLoader
from sklearn.metrics import f1_score
import timm

import config
from src.dataset import SpiceDataset, get_val_transform, load_manifest_splits

ROOT = Path(_base)
SS_M = str(ROOT / "outputs" / "manifest_overlap_ss.json")
IN_M = str(ROOT / "outputs" / "manifest_overlap_indian.json")


def _loader(paths, labels):
    return DataLoader(SpiceDataset(paths, labels, get_val_transform(), multimodal=False),
                      batch_size=64, num_workers=2, pin_memory=True)


def _load(model_name, ckpt, n, device):
    m = timm.create_model(model_name, pretrained=False, num_classes=n).to(device)
    ck = torch.load(ckpt, map_location=device, weights_only=False)
    m.load_state_dict(ck["model_state"]); m.eval()
    return m


@torch.no_grad()
def _eval(model, loader, device, n):
    yt, yp = [], []
    for imgs, tex, col, labels in loader:
        yp.extend(model(imgs.to(device)).argmax(1).cpu().tolist()); yt.extend(labels.tolist())
    yt, yp = np.array(yt), np.array(yp)
    return float((yt == yp).mean()), float(f1_score(yt, yp, average="macro",
                                                    labels=list(range(n)), zero_division=0))


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", required=True)
    ap.add_argument("--ss_ckpt", required=True)
    ap.add_argument("--in_ckpt", required=True)
    args = ap.parse_args()

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    ss, classes = load_manifest_splits(SS_M)
    ins, _ = load_manifest_splits(IN_M)
    n = len(classes)
    ss_loader, in_loader = _loader(*ss["test"]), _loader(*ins["test"])
    ss_model = _load(args.model, args.ss_ckpt, n, device)
    in_model = _load(args.model, args.in_ckpt, n, device)

    res = {"model": args.model}
    for tag, model in (("ss", ss_model), ("in", in_model)):
        row = {}
        for ttag, loader in (("ss", ss_loader), ("in", in_loader)):
            acc, f1 = _eval(model, loader, device, n)
            row[ttag] = {"acc": acc, "macro_f1": f1}
        res[tag] = row
    res["ss_tax_pp"] = round((res["ss"]["ss"]["acc"] - res["ss"]["in"]["acc"]) * 100, 2)
    res["in_tax_pp"] = round((res["in"]["in"]["acc"] - res["in"]["ss"]["acc"]) * 100, 2)
    out = ROOT / "outputs" / f"bench_{args.model}.json"
    json.dump(res, open(out, "w"), indent=2)
    print(f"{args.model}: within SS {res['ss']['ss']['acc']*100:.1f} / IN {res['in']['in']['acc']*100:.1f} | "
          f"collapse IN->SS tax {res['in_tax_pp']} pp  ->  {out}")


if __name__ == "__main__":
    main()