{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":""},"kaggle":{"accelerator":"gpu","isInternetEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"40a4e749","cell_type":"code","source":"%pip install -q --no-deps zennit==1.0.0 captum==0.9.0","metadata":{},"outputs":[],"execution_count":null},{"id":"ea59cf45","cell_type":"code","source":"import os\nimport glob\nimport math\nimport time\nimport types\nimport random\nimport warnings\nfrom contextlib import contextmanager\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nfrom torchvision import transforms\nfrom torchvision.models import vit_b_16, ViT_B_16_Weights, vision_transformer\nfrom skimage.segmentation import slic\nfrom captum.attr import KernelShap\nfrom zennit.composites import NameMapComposite\nfrom zennit.rules import Gamma\nfrom zennit.core import BasicHook, ParamMod, stabilize\n\nwarnings.filterwarnings(\"ignore\")\n\nSMOKE_TEST = False\n\nIMAGE_ROOT_CANDIDATES = [\n    \"/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val\",\n    \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val\",\n]\n\nCONFIG = {\n    \"image_root\": next((p for p in IMAGE_ROOT_CANDIDATES if os.path.isdir(p)), IMAGE_ROOT_CANDIDATES[0]),\n    \"num_samples\": 50 if SMOKE_TEST else 3200,\n    \"seed\": 0,\n    \"device\": \"cuda\" if torch.cuda.is_available() else \"cpu\",\n    \"autocast\": True,\n    \"num_steps\": 50,\n    \"flip_batch\": 51,\n    \"ig_steps\": 20,\n    \"sg_samples\": 20,\n    \"sg_sigma\": 0.01,\n    \"atman_p\": 1.0,\n    \"atman_t\": 0.1,\n    \"atman_batch\": 49,\n    \"roll_dt\": 0.90,\n    \"gxroll_dt\": 1.0,\n    \"shap_samples\": 2000,\n    \"shap_baseline\": 0.5,\n    \"shap_batch\": 250,\n    \"slic_segments\": 100,\n    \"slic_compactness\": 10,\n    \"gamma_conv\": 0.25,\n    \"gamma_linear\": 0.05,\n    \"methods\": [\n        \"random\", \"ixg\", \"ig\", \"smoothgrad\", \"gradcam\", \"attnroll\", \"gxattnroll\",\n        \"atman\", \"kernelshap\", \"cplrp\", \"gamma_cplrp\", \"attnlrp\",\n    ],\n    \"out_dir\": \"/kaggle/working\",\n    \"save_every\": 100,\n    \"time_budget_hours\": 11.0,\n    \"flip_to_normalized_zero\": True,\n    \"resume_root\": \"/kaggle/input\",\n}\n\nPAPER_VIT_B16 = {\n    \"random\": 0.01, \"ixg\": 0.80, \"ig\": 1.54, \"smoothgrad\": -0.04, \"gradcam\": 0.27,\n    \"attnroll\": 1.31, \"gxattnroll\": 2.60, \"atman\": 0.70, \"kernelshap\": 4.71,\n    \"cplrp\": 2.53, \"gamma_cplrp\": 6.06, \"attnlrp\": 6.19,\n}","metadata":{},"outputs":[],"execution_count":null},{"id":"c7375b92","cell_type":"code","source":"print(\"torch\", torch.__version__, \"| cuda available:\", torch.cuda.is_available())\nif torch.cuda.is_available():\n    print(\"gpu:\", torch.cuda.get_device_name(0), \"| capability:\", torch.cuda.get_device_capability(0))\n    torch.zeros(1, device=\"cuda\").sum().item()\nelse:\n    raise RuntimeError(\"No GPU: Settings -> Accelerator -> GPU T4 x2\")\nif not os.path.isdir(CONFIG[\"image_root\"]):\n    raise FileNotFoundError(\"ImageNet val not found. Add Input -> Competitions -> ImageNet Object Localization Challenge\")\nnum_images = len(glob.glob(os.path.join(CONFIG[\"image_root\"], \"*.JPEG\")))\nprint(\"image_root:\", CONFIG[\"image_root\"], \"|\", num_images, \"images\")\nif num_images < CONFIG[\"num_samples\"]:\n    raise FileNotFoundError(\"too few images found in image_root\")\ntry:\n    vit_b_16(weights=ViT_B_16_Weights.IMAGENET1K_V1)\nexcept Exception as e:\n    raise RuntimeError(\"Could not download ViT weights: Settings -> Internet -> On\") from e\nprint(\"SMOKE_TEST =\", SMOKE_TEST, \"| num_samples =\", CONFIG[\"num_samples\"], \"| preflight ok\")","metadata":{},"outputs":[],"execution_count":null},{"id":"73058b37","cell_type":"code","source":"class AttentionState:\n    def __init__(self):\n        self.mode = \"none\"\n        self.record = False\n        self.maps = []\n        self.scale = None\n\n\nSTATE = AttentionState()\nPATCH_TARGETS = [(nn.GELU, \"forward\"), (nn.LayerNorm, \"forward\")]\n\n\nclass DivideGradient(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, x, factor):\n        ctx.factor = factor\n        return x\n\n    @staticmethod\n    def backward(ctx, grad):\n        return grad / ctx.factor, None\n\n\nclass IdentityRule(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, fn, x):\n        out = fn(x)\n        if x.requires_grad:\n            ctx.save_for_backward(out / (x + 1e-10))\n        return out\n\n    @staticmethod\n    def backward(ctx, grad):\n        return None, ctx.saved_tensors[0] * grad\n\n\ndef divide_gradient(x, factor):\n    return DivideGradient.apply(x, factor)\n\n\ndef gelu_identity_forward(self, x):\n    return IdentityRule.apply(lambda t: F.gelu(t, approximate=self.approximate), x)\n\n\ndef layer_norm_identity_forward(self, x):\n    mean = x.mean(dim=-1, keepdim=True)\n    var = ((x - mean) ** 2).mean(dim=-1, keepdim=True)\n    y = (x - mean) / (var + self.eps).sqrt().detach()\n    if self.weight is not None:\n        y = y * self.weight\n    if self.bias is not None:\n        y = y + self.bias\n    return y\n\n\ndef zennit_forward_hook(self, module, input, output):\n    self.stored_tensors[\"input\"] = input\n    self.stored_tensors[\"output\"] = output\n\n\ndef zennit_backward_hook(self, module, grad_input, grad_output):\n    grad_output = (grad_output[0] * self.stored_tensors[\"output\"],)\n    grad_output[0].requires_grad = True\n    original_input = self.stored_tensors[\"input\"][0].clone()\n    inputs, outputs = [], []\n    for in_mod, param_mod, out_mod in zip(self.input_modifiers, self.param_modifiers, self.output_modifiers):\n        inp = in_mod(original_input).requires_grad_()\n        with ParamMod.ensure(param_mod)(module) as modified, torch.autograd.enable_grad():\n            out = out_mod(modified.forward(inp))\n        inputs.append(inp)\n        outputs.append(out)\n    grad_outputs = self.gradient_mapper(grad_output[0], outputs)\n    gradients = torch.autograd.grad(outputs, inputs, grad_outputs=grad_outputs, create_graph=grad_output[0].requires_grad)\n    relevance = self.reducer(inputs, gradients) / stabilize(original_input, epsilon=1e-10)\n    return tuple(relevance if original.shape == relevance.shape else None for original in grad_input)\n\n\nBasicHook.forward = zennit_forward_hook\nBasicHook.backward = zennit_backward_hook\n\n\ndef attention_forward(self, query, key, value, need_weights=False, **kwargs):\n    b, n, dim = query.shape\n    h = self.num_heads\n    d = dim // h\n    q, k, v = F.linear(query, self.in_proj_weight, self.in_proj_bias).chunk(3, dim=-1)\n    q, k, v = [t.reshape(b, n, h, d).transpose(1, 2) for t in (q, k, v)]\n    if STATE.mode == \"cplrp\":\n        q, k = q.detach(), k.detach()\n    elif STATE.mode == \"attnlrp\":\n        q, k, v = divide_gradient(q, 4), divide_gradient(k, 4), divide_gradient(v, 2)\n    scores = q @ k.transpose(-2, -1) / math.sqrt(d)\n    if STATE.scale is not None:\n        scores = scores * STATE.scale\n    attn = scores.softmax(dim=-1)\n    if STATE.record:\n        if attn.requires_grad:\n            attn.retain_grad()\n        STATE.maps.append(attn)\n    out = (attn @ v).transpose(1, 2).reshape(b, n, dim)\n    return self.out_proj(out), None\n\n\ndef install_attention(model):\n    for block in model.encoder.layers:\n        block.self_attention.forward = types.MethodType(attention_forward, block.self_attention)\n    return model\n\n\ndef snapshot():\n    return {(obj, name): getattr(obj, name) for obj, name in PATCH_TARGETS}\n\n\ndef restore(original):\n    for (obj, name), value in original.items():\n        setattr(obj, name, value)\n    STATE.mode, STATE.record, STATE.maps, STATE.scale = \"none\", False, [], None\n\n\ndef gamma_composite(model, cfg):\n    convs = tuple(n for n, m in model.named_modules() if isinstance(m, nn.Conv2d))\n    linears = tuple(\n        n for n, m in model.named_modules()\n        if isinstance(m, nn.Linear) and not n.endswith(\"out_proj\")\n    )\n    return NameMapComposite(name_map=[\n        (convs, Gamma(gamma=cfg[\"gamma_conv\"])),\n        (linears, Gamma(gamma=cfg[\"gamma_linear\"])),\n    ])\n\n\n@contextmanager\ndef lrp_mode(model, mode, use_gamma, original, cfg):\n    restore(original)\n    nn.GELU.forward = gelu_identity_forward\n    nn.LayerNorm.forward = layer_norm_identity_forward\n    STATE.mode = mode\n    composite = gamma_composite(model, cfg) if use_gamma else None\n    if composite is not None:\n        composite.register(model)\n    try:\n        yield\n    finally:\n        if composite is not None:\n            composite.remove()\n        restore(original)","metadata":{},"outputs":[],"execution_count":null},{"id":"408fca6a","cell_type":"code","source":"def set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n\n\nclass Normalizer:\n    def __init__(self, mean, std, device):\n        self.mean = torch.tensor(mean, device=device).view(1, 3, 1, 1)\n        self.std = torch.tensor(std, device=device).view(1, 3, 1, 1)\n\n    def __call__(self, x):\n        return (x - self.mean) / self.std\n\n\ndef load_model(cfg):\n    weights = ViT_B_16_Weights.IMAGENET1K_V1\n    preset = weights.transforms()\n    model = vit_b_16(weights=weights).to(cfg[\"device\"]).eval()\n    for p in model.parameters():\n        p.requires_grad_(False)\n    prep = transforms.Compose([\n        transforms.Resize(preset.resize_size, interpolation=preset.interpolation),\n        transforms.CenterCrop(preset.crop_size),\n        transforms.ToTensor(),\n    ])\n    return install_attention(model), prep, Normalizer(preset.mean, preset.std, cfg[\"device\"])\n\n\ndef load_paths(cfg):\n    paths = sorted(glob.glob(os.path.join(cfg[\"image_root\"], \"**\", \"*.JPEG\"), recursive=True))\n    if not paths:\n        raise FileNotFoundError(f\"no .JPEG files under {cfg['image_root']}\")\n    return random.Random(cfg[\"seed\"]).sample(paths, min(cfg[\"num_samples\"], len(paths)))\n\n\ndef load_previous_results(cfg):\n    found = []\n    for depth in range(1, 5):\n        found += glob.glob(os.path.join(cfg[\"resume_root\"], *[\"*\"] * depth, \"vit_per_image.csv\"))\n    if not found:\n        print(\"resume: no previous vit_per_image.csv found, starting fresh\")\n        return pd.DataFrame()\n    frames = [pd.read_csv(p) for p in sorted(set(found))]\n    previous = pd.concat(frames, ignore_index=True).drop_duplicates(subset=\"image\", keep=\"first\")\n    previous = previous.dropna(subset=cfg[\"methods\"])\n    for p, f in zip(sorted(set(found)), frames):\n        print(f\"resume: {len(f)} images from {p}\")\n    print(f\"resume: {len(previous)} unique finished images\")\n    return previous","metadata":{},"outputs":[],"execution_count":null},{"id":"a87c97be","cell_type":"code","source":"def patch_grid(model):\n    return model.image_size // model.patch_size\n\n\ndef to_pixels(patch_scores, model):\n    g = patch_grid(model)\n    grid = patch_scores.float().reshape(1, 1, g, g)\n    return F.interpolate(grid, size=(model.image_size, model.image_size), mode=\"bilinear\", align_corners=False)[0, 0]\n\n\ndef autocast(cfg):\n    return torch.autocast(\"cuda\", dtype=torch.float16, enabled=cfg[\"autocast\"] and cfg[\"device\"] == \"cuda\")\n\n\ndef input_times_grad(model, x, target):\n    x = x.clone().requires_grad_(True)\n    model(x)[0, target].backward()\n    return (x * x.grad).sum(1)[0].detach()\n\n\ndef attr_random(model, raw, x, target, norm, cfg, original):\n    return torch.randn(model.image_size, model.image_size)\n\n\ndef attr_ixg(model, raw, x, target, norm, cfg, original):\n    restore(original)\n    return input_times_grad(model, x, target)\n\n\ndef attr_ig(model, raw, x, target, norm, cfg, original):\n    restore(original)\n    alphas = torch.arange(1, cfg[\"ig_steps\"] + 1, device=x.device).float() / cfg[\"ig_steps\"]\n    xs = (alphas.view(-1, 1, 1, 1) * x).requires_grad_(True)\n    model(xs)[:, target].sum().backward()\n    return (x * xs.grad.mean(0, keepdim=True)).sum(1)[0].detach()\n\n\ndef attr_smoothgrad(model, raw, x, target, norm, cfg, original):\n    restore(original)\n    std = cfg[\"sg_sigma\"] * (x.max() - x.min())\n    xs = (x + std * torch.randn(cfg[\"sg_samples\"], *x.shape[1:], device=x.device)).requires_grad_(True)\n    model(xs)[:, target].sum().backward()\n    return xs.grad.mean(0).sum(0).detach()\n\n\ndef attention_maps(model, x, target, with_grad):\n    STATE.record, STATE.maps = True, []\n    try:\n        if with_grad:\n            xg = x.clone().requires_grad_(True)\n            model(xg)[0, target].backward()\n            maps = [(a.detach()[0], a.grad[0]) for a in STATE.maps]\n        else:\n            with torch.no_grad():\n                model(x)\n            maps = [(a[0], None) for a in STATE.maps]\n    finally:\n        STATE.record, STATE.maps = False, []\n    return maps\n\n\ndef rollout(layer_maps, dt):\n    n = layer_maps[0].shape[-1]\n    eye = torch.eye(n, device=layer_maps[0].device)\n    joint = eye.clone()\n    for a in layer_maps:\n        positive = a[a > 0].float()\n        if dt < 1.0 and positive.numel() > 0:\n            a = torch.where(a > torch.quantile(positive, dt), torch.zeros_like(a), a)\n        joint = (eye + a) @ joint\n    return joint[0, 1:]\n\n\ndef attr_attnroll(model, raw, x, target, norm, cfg, original):\n    restore(original)\n    maps = attention_maps(model, x, target, with_grad=False)\n    return to_pixels(rollout([a.mean(0) for a, _ in maps], cfg[\"roll_dt\"]), model).cpu()\n\n\ndef attr_gxattnroll(model, raw, x, target, norm, cfg, original):\n    restore(original)\n    maps = attention_maps(model, x, target, with_grad=True)\n    return to_pixels(rollout([(g * a).clamp(min=0).mean(0) for a, g in maps], cfg[\"gxroll_dt\"]), model).cpu()\n\n\ndef attr_gradcam(model, raw, x, target, norm, cfg, original):\n    restore(original)\n    a, g = attention_maps(model, x, target, with_grad=True)[-1]\n    cam = (a * g.mean(dim=(1, 2), keepdim=True)).mean(0).clamp(min=0)\n    return to_pixels(cam[0, 1:], model).cpu()\n\n\ndef attr_atman(model, raw, x, target, norm, cfg, original):\n    restore(original)\n    with torch.no_grad():\n        tokens = model._process_input(x)[0]\n        sim = F.cosine_similarity(tokens[:, None], tokens[None], dim=-1)\n        factor = torch.where(sim >= cfg[\"atman_t\"], 1 - cfg[\"atman_p\"] * sim, torch.ones_like(sim))\n        factor.fill_diagonal_(1 - cfg[\"atman_p\"])\n        scale = torch.ones(factor.shape[0], factor.shape[0] + 1, device=x.device)\n        scale[:, 1:] = factor\n        with autocast(cfg):\n            base = F.cross_entropy(model(x).float(), torch.tensor([target], device=x.device))\n            deltas = []\n            for chunk in scale.split(cfg[\"atman_batch\"]):\n                STATE.scale = chunk[:, None, None, :]\n                logits = model(x.expand(len(chunk), -1, -1, -1)).float()\n                labels = torch.full((len(chunk),), target, device=x.device)\n                deltas.append(F.cross_entropy(logits, labels, reduction=\"none\") - base)\n                STATE.scale = None\n    return to_pixels(torch.cat(deltas), model).cpu()\n\n\ndef attr_kernelshap(model, raw, x, target, norm, cfg, original):\n    restore(original)\n    image = raw[0].permute(1, 2, 0).cpu().numpy()\n    segments = slic(image, n_segments=cfg[\"slic_segments\"], compactness=cfg[\"slic_compactness\"], start_label=0)\n    if np.unique(segments).size < 3:\n        ps, (hh, ww) = model.patch_size, segments.shape\n        segments = (np.arange(hh)[:, None] // ps) * (ww // ps) + np.arange(ww)[None] // ps\n    mask = torch.from_numpy(segments).long().to(raw.device)[None, None]\n\n    def forward(inp):\n        with torch.no_grad(), autocast(cfg):\n            return model(norm(inp)).float()\n\n    attributions = KernelShap(forward).attribute(\n        raw, baselines=cfg[\"shap_baseline\"], target=target, feature_mask=mask,\n        n_samples=cfg[\"shap_samples\"], perturbations_per_eval=cfg[\"shap_batch\"],\n    )\n    return attributions.sum(1)[0].detach()\n\n\ndef make_lrp(mode, use_gamma):\n    def attribute(model, raw, x, target, norm, cfg, original):\n        with lrp_mode(model, mode, use_gamma, original, cfg):\n            return input_times_grad(model, x, target)\n    return attribute\n\n\nATTRIBUTORS = {\n    \"random\": attr_random,\n    \"ixg\": attr_ixg,\n    \"ig\": attr_ig,\n    \"smoothgrad\": attr_smoothgrad,\n    \"gradcam\": attr_gradcam,\n    \"attnroll\": attr_attnroll,\n    \"gxattnroll\": attr_gxattnroll,\n    \"atman\": attr_atman,\n    \"kernelshap\": attr_kernelshap,\n    \"cplrp\": make_lrp(\"cplrp\", False),\n    \"gamma_cplrp\": make_lrp(\"cplrp\", True),\n    \"attnlrp\": make_lrp(\"attnlrp\", True),\n}","metadata":{},"outputs":[],"execution_count":null},{"id":"93c77fc5","cell_type":"code","source":"def flip_curves(model, raw, norm, target, relevance, cfg):\n    h, w = relevance.shape\n    num_pixels, k = h * w, cfg[\"num_steps\"]\n    counts = torch.tensor([round(i * num_pixels / k) for i in range(k + 1)], device=raw.device)\n    curves = {}\n    for name, descending in ((\"morf\", True), (\"lerf\", False)):\n        order = torch.argsort(relevance.flatten().to(raw.device), descending=descending)\n        rank = torch.empty_like(order)\n        rank[order] = torch.arange(num_pixels, device=raw.device)\n        keep = (rank[None] >= counts[:, None]).view(k + 1, 1, h, w)\n        fill = norm.mean if cfg[\"flip_to_normalized_zero\"] else torch.zeros_like(norm.mean)\n        batch = torch.where(keep, raw, fill)\n        values = []\n        with torch.no_grad(), autocast(cfg):\n            for chunk in batch.split(cfg[\"flip_batch\"]):\n                values.append(model(norm(chunk))[:, target].float().cpu())\n        curves[name] = torch.cat(values).numpy()\n    return curves\n\n\ndef summarize(per_image, methods):\n    rows = []\n    for method in methods:\n        s = per_image[method].to_numpy()\n        rows.append({\n            \"method\": method,\n            \"delta_A\": s.mean(),\n            \"sem\": s.std(ddof=1) / np.sqrt(len(s)) if len(s) > 1 else np.nan,\n            \"paper\": PAPER_VIT_B16.get(method, np.nan),\n            \"n\": len(s),\n        })\n    return pd.DataFrame(rows).sort_values(\"delta_A\", ascending=False).reset_index(drop=True)\n\n\ndef plot_curves(curves, cfg):\n    xs = np.linspace(0, 1, cfg[\"num_steps\"] + 1)\n    fig, axes = plt.subplots(1, 3, figsize=(16, 4.5))\n    for method, c in curves.items():\n        axes[0].plot(xs, c[\"lerf\"] - c[\"morf\"], label=method)\n        axes[1].plot(xs, c[\"morf\"], label=method)\n        axes[2].plot(xs, c[\"lerf\"], label=method)\n    for ax, title in zip(axes, [\"LeRF - MoRF\", \"MoRF\", \"LeRF\"]):\n        ax.set_title(title)\n        ax.set_xlabel(\"fraction of pixels flipped\")\n        ax.set_ylabel(\"target logit\")\n    axes[0].legend(fontsize=8)\n    plt.tight_layout()\n    plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"a2c1a63b","cell_type":"code","source":"def main(cfg=CONFIG):\n    set_seed(cfg[\"seed\"])\n    original = snapshot()\n    model, prep, norm = load_model(cfg)\n    methods = cfg[\"methods\"]\n    previous = load_previous_results(cfg)\n    done = set(previous[\"image\"]) if len(previous) else set()\n    paths = [p for p in load_paths(cfg) if os.path.basename(p) not in done]\n    print(f\"{len(done)} images already done, {len(paths)} remaining\")\n    os.makedirs(cfg[\"out_dir\"], exist_ok=True)\n    per_image_csv = os.path.join(cfg[\"out_dir\"], \"vit_per_image.csv\")\n    sums = {m: {\"morf\": 0.0, \"lerf\": 0.0} for m in methods}\n    records = previous.to_dict(\"records\") if len(previous) else []\n    num_new = 0\n    budget = cfg[\"time_budget_hours\"] * 3600\n    start = time.time()\n    for i, path in enumerate(tqdm(paths, desc=\"images\")):\n        elapsed = time.time() - start\n        if i > 0 and elapsed + elapsed / i > budget:\n            print(f\"time budget reached after {i} images ({elapsed / 3600:.2f} h), stopping early\")\n            break\n        raw = prep(Image.open(path).convert(\"RGB\"))[None].to(cfg[\"device\"])\n        x = norm(raw)\n        with torch.no_grad():\n            target = int(model(x).argmax())\n        row = {\"image\": os.path.basename(path), \"target\": target}\n        for method in methods:\n            relevance = ATTRIBUTORS[method](model, raw, x, target, norm, cfg, original)\n            c = flip_curves(model, raw, norm, target, relevance, cfg)\n            row[method] = float(c[\"lerf\"].mean() - c[\"morf\"].mean())\n            sums[method][\"morf\"] = sums[method][\"morf\"] + c[\"morf\"]\n            sums[method][\"lerf\"] = sums[method][\"lerf\"] + c[\"lerf\"]\n        records.append(row)\n        num_new += 1\n        if i + 1 == 10:\n            projected = (time.time() - start) / 10 * len(paths) / 3600\n            print(f\"projected total: {projected:.2f} h for {len(paths)} images (budget {cfg['time_budget_hours']} h)\")\n        if (i + 1) % cfg[\"save_every\"] == 0:\n            pd.DataFrame(records).to_csv(per_image_csv, index=False)\n    per_image = pd.DataFrame(records)\n    per_image.to_csv(per_image_csv, index=False)\n    table = summarize(per_image, methods)\n    table.to_csv(os.path.join(cfg[\"out_dir\"], \"vit_results.csv\"), index=False)\n    curves = {m: {o: sums[m][o] / num_new for o in (\"morf\", \"lerf\")} for m in methods} if num_new else None\n    print(f\"{num_new} new images this run, {len(per_image)} total\")\n    return table, curves, per_image","metadata":{},"outputs":[],"execution_count":null},{"id":"1ad985e4","cell_type":"code","source":"table, curves, per_image = main()\nprint(table.to_string(index=False))\nif curves is not None:\n    plot_curves(curves, CONFIG)","metadata":{},"outputs":[],"execution_count":null}]}