SpiceNet / code /eval_backbone_shortcut.py
Noushad999's picture
Upload folder using huggingface_hub
1ea7ba6 verified
Raw
History Blame Contribute Delete
3.13 kB
"""
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()