{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":126777,"databundleVersionId":15314950},{"sourceType":"datasetVersion","sourceId":15049613,"datasetId":9634391,"databundleVersionId":15929655}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 0 — Optional allocator hint (run FIRST)\n# ============================================================\nimport os\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 1 — Install deps\n# ============================================================\n\n!pip -q install -U timm ImageHash scipy scikit-learn opencv-python","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T03:31:39.148140Z","iopub.execute_input":"2026-03-05T03:31:39.148736Z","iopub.status.idle":"2026-03-05T03:32:00.521608Z","shell.execute_reply.started":"2026-03-05T03:31:39.148709Z","shell.execute_reply":"2026-03-05T03:32:00.520576Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 2 — Imports\n# ============================================================\n\nimport math, random, os, shutil\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom PIL import Image\n\nimport cv2\nimport multiprocessing as mp\nfrom concurrent.futures import ProcessPoolExecutor\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, Sampler\nimport torchvision.transforms as transforms\nimport timm\n\nfrom sklearn.model_selection import StratifiedGroupKFold\n\nfrom scipy.spatial.distance import cdist\nfrom scipy.sparse import csr_matrix\nfrom scipy.sparse.csgraph import connected_components\n\nimport imagehash","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:27.711664Z","iopub.execute_input":"2026-03-05T04:24:27.712597Z","iopub.status.idle":"2026-03-05T04:24:33.448221Z","shell.execute_reply.started":"2026-03-05T04:24:27.712561Z","shell.execute_reply":"2026-03-05T04:24:33.447523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# REPLACE CELL 3 — Config (DINOv3-L + Gates + CLS-mask + FFT + PCGrad-ready)\n#   Adds knobs for:\n#     - CLS-guided masking pool\n#     - HardCircleLoss (top-k)\n#     - 3-group LRs (bb / arc / aux)\n#     - PCGrad logging\n#     - step-scheduler warmup fraction\n# ============================================================\n\nclass Config:\n    seed = 42\n\n    # -------- Backbone\n    model_name = \"vit_large_patch16_dinov3.lvd1689m\"\n    img_size   = 448\n\n    # -------- Train length\n    num_epochs = 15\n\n    # -------- Metric schedule\n    metric_loss = \"circle\"\n    metric_weight = 0.25\n    metric_warmup_epochs = 0\n    metric_ramp_epochs   =1\n    metric_every_steps   = 2\n    circle_on_feat_cls = False   # (DP-safe anchor for PCGrad)\n\n    # -------- HardCircleLoss knobs (if you switch criterion_metric)\n    circle_m = 0.20\n    circle_gamma = 8.0\n    hard_topk_pos = 2        # with K=3 => K-1=2 positives per anchor\n    hard_topk_neg = 24\n\n    # -------- Cache behavior\n    cache_crops = True\n    cache_root = \"/kaggle/working/jaguar_cache\"\n    use_crop = False\n    crop_pad_frac = 0.08\n    cache_resize = True\n    cache_letterbox = False\n    cache_pad_frac = 0.08\n    cache_workers = 8\n\n    # -------- CLS-guided masking (OUTSIDE THE BOX)\n    use_cls_attn_pool = True     # enable CLS->patch mask\n    cls_attn_temp  = 0.10        # try 0.08–0.15\n    cls_attn_clamp = 4.0         # prevents spike weights\n    cls_attn_mix   = 0.70        # masked GeM vs normal GeM mix (0..1)\n\n    # -------- Gates + FFT\n    gate_mode = \"sample\"      # \"sample\" (better) or \"scalar\" (faster)\n    max_cls_w = 0.50\n    max_fft_w = 0.25\n    init_cls_w = 0.20\n    init_fft_w = 0.10\n    gate_hidden = 128\n\n    use_fft_branch = True\n    fft_log = True\n    use_fuse_ln = False\n\n    # -------- Optim (base LRs)\n    weight_decay = 5e-4\n    lr_backbone = 2.6e-5\n    lr_head     = 1.6e-4\n    min_lr_backbone = 4e-6\n    min_lr_head     = 2.4e-5\n\n    # -------- 3-group LRs (bb / arc / aux) used by your get_param_groups + step scheduler\n    lr_arcface = 8.0e-5\n    lr_aux     = 3.0e-4\n    min_lr_arcface = min_lr_head * 0.75\n    min_lr_aux     = min_lr_head * 1.50\n\n    # -------- Step scheduler warmup fraction (used by StepWarmupCosine)\n    warmup_frac = 0.05\n\n    # (kept for compatibility if you ever revert to epoch-scheduler)\n    lr_Tmax = 15\n\n    # -------- PK sampling (batch=48)\n    P = 18\n    K = 6\n    grad_accum = 1\n    steps_mult = 5\n\n    # -------- ArcFace\n    arcface_s = 30.0   \n    arcface_m = 0.50\n    subcenter_k = 3\n\n    # -------- Head\n    neck_dropout = 0.10\n    embed_dim = None\n\n    # -------- Eval / inference\n    val_frac_per_id = 0.2\n    eval_every = 1\n    use_tta = False\n    use_qe = True\n    qe_topk = 3\n    use_rerank = False\n\n    # -------- DataLoader\n    num_workers = 4\n    pin_memory = True\n    persistent_workers = True\n    prefetch_factor = 2\n\n    # -------- Multi-GPU / memory\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    device_type = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    use_checkpointing = True\n\n    # -------- Checkpoints\n    save_topk = 4\n    ckpt_dir  = \"/kaggle/working/ckpts\"\n\nprint(\"torch.cuda.device_count():\", torch.cuda.device_count())\n!nvidia-smi -L\nprint(\"model:\", Config.model_name)\nprint(\"img_size:\", Config.img_size, \"| batch:\", Config.P * Config.K, \"| steps_mult:\", Config.steps_mult)\nprint(\"num_workers:\", Config.num_workers, \"| prefetch_factor:\", Config.prefetch_factor)\nprint(\"eval_every:\", Config.eval_every, \"| metric_every_steps:\", Config.metric_every_steps)\nprint(\"CLS-mask:\", Config.use_cls_attn_pool)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:33.449338Z","iopub.execute_input":"2026-03-05T04:24:33.450308Z","iopub.status.idle":"2026-03-05T04:24:33.916318Z","shell.execute_reply.started":"2026-03-05T04:24:33.450280Z","shell.execute_reply":"2026-03-05T04:24:33.915540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 4 — Seed + FAST flags\n# ============================================================\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nseed_everything(Config.seed)\n\ntorch.backends.cudnn.deterministic = False\ntorch.backends.cudnn.benchmark = True\n\ntorch.backends.cuda.matmul.allow_tf32 = True\ntorch.backends.cudnn.allow_tf32 = True\ntry:\n    torch.set_float32_matmul_precision(\"high\")\nexcept Exception:\n    pass\n\nprint(\"✅ Seed set + cudnn.benchmark=True (fast).\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:36.398875Z","iopub.execute_input":"2026-03-05T04:24:36.399795Z","iopub.status.idle":"2026-03-05T04:24:36.408630Z","shell.execute_reply.started":"2026-03-05T04:24:36.399758Z","shell.execute_reply":"2026-03-05T04:24:36.407797Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 5 — Preprocess helpers (CROP / NO-CROP)\n# ============================================================\n\ndef alpha_bbox_from_rgba(rgba: Image.Image, thr=1):\n    arr = np.array(rgba)\n    if arr.ndim != 3 or arr.shape[2] < 4:\n        w, h = rgba.size\n        return (0, 0, w, h), None\n\n    alpha = arr[:, :, 3]\n    ys, xs = np.where(alpha >= thr)\n    if len(xs) == 0 or len(ys) == 0:\n        w, h = rgba.size\n        return (0, 0, w, h), Image.fromarray(alpha)\n\n    x0, x1 = xs.min(), xs.max() + 1\n    y0, y1 = ys.min(), ys.max() + 1\n    return (x0, y0, x1, y1), Image.fromarray(alpha)\n\ndef _composite_rgba_on_gray(rgba: Image.Image, bg=128):\n    arr = np.array(rgba)  # HxWx4\n    rgb = arr[:, :, :3].astype(np.float32)\n    a   = (arr[:, :, 3:4].astype(np.float32) / 255.0)\n    bgc = np.full_like(rgb, float(bg))\n    out = (rgb * a + bgc * (1.0 - a)).astype(np.uint8)\n    return Image.fromarray(out)   # ✅ no mode=)\n\ndef crop_and_neutralize(img: Image.Image, pad_frac=0.08, bg=128):\n    rgba = img.convert(\"RGBA\") if img.mode != \"RGBA\" else img\n    bbox, _ = alpha_bbox_from_rgba(rgba)\n    comp = _composite_rgba_on_gray(rgba, bg=bg)\n\n    x0, y0, x1, y1 = bbox\n    w = x1 - x0\n    h = y1 - y0\n    pad = int(max(w, h) * float(pad_frac))\n\n    W, H = comp.size\n    x0p = max(0, x0 - pad)\n    y0p = max(0, y0 - pad)\n    x1p = min(W, x1 + pad)\n    y1p = min(H, y1 + pad)\n    return comp.crop((x0p, y0p, x1p, y1p)).convert(\"RGB\")\n\ndef no_crop_rgb(img: Image.Image, bg=128):\n    if img.mode == \"RGBA\":\n        return _composite_rgba_on_gray(img, bg=bg)\n    return img.convert(\"RGB\")\n\ndef preprocess_pil(img: Image.Image):\n    if getattr(Config, \"use_crop\", False):\n        return crop_and_neutralize(img, pad_frac=Config.crop_pad_frac, bg=128)\n    return no_crop_rgb(img, bg=128)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:37.174594Z","iopub.execute_input":"2026-03-05T04:24:37.175229Z","iopub.status.idle":"2026-03-05T04:24:37.185106Z","shell.execute_reply.started":"2026-03-05T04:24:37.175198Z","shell.execute_reply":"2026-03-05T04:24:37.184377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 6 — Cache builder (fast cv2)\n# ============================================================\n\nCACHE_SIDE = int(Config.img_size)\nCACHE_BG = 128\n\ndef _atomic_imwrite(path, bgr):\n    tmp = str(path) + \".tmp\"\n    ok = cv2.imwrite(tmp, bgr)\n    if ok:\n        os.replace(tmp, str(path))\n    return ok\n\ndef _letterbox_to_square_bgr(img, side=448, fill=128):\n    h, w = img.shape[:2]\n    scale = min(side / w, side / h)\n    nw, nh = int(round(w * scale)), int(round(h * scale))\n    resized = cv2.resize(img, (nw, nh), interpolation=cv2.INTER_AREA)\n    out = np.full((side, side, 3), fill, dtype=np.uint8)\n    x0 = (side - nw) // 2\n    y0 = (side - nh) // 2\n    out[y0:y0+nh, x0:x0+nw] = resized\n    return out\n\ndef _cache_one_cv2(task):\n    fn, src_dir, cache_dir, pad_frac = task\n    src = Path(src_dir) / fn\n    out_path = Path(cache_dir) / fn\n    if out_path.exists():\n        return 0\n\n    img = cv2.imread(str(src), cv2.IMREAD_UNCHANGED)\n    if img is None:\n        out = np.full((CACHE_SIDE, CACHE_SIDE, 3), CACHE_BG, np.uint8)\n        _atomic_imwrite(out_path, out)\n        return 1\n\n    if img.ndim == 2:\n        img = np.stack([img, img, img], axis=-1)\n\n    use_crop = bool(getattr(Config, \"use_crop\", False))\n\n    if img.shape[2] >= 4:\n        bgr = img[:, :, :3]\n        a   = img[:, :, 3]\n\n        if use_crop:\n            ys, xs = np.where(a > 0)\n            if len(xs) > 0:\n                x0, x1 = xs.min(), xs.max() + 1\n                y0, y1 = ys.min(), ys.max() + 1\n            else:\n                H, W = bgr.shape[:2]\n                x0, y0, x1, y1 = 0, 0, W, H\n\n            wbox = x1 - x0\n            hbox = y1 - y0\n            pad  = int(max(wbox, hbox) * float(pad_frac))\n\n            H, W = bgr.shape[:2]\n            x0p = max(0, x0 - pad); y0p = max(0, y0 - pad)\n            x1p = min(W, x1 + pad); y1p = min(H, y1 + pad)\n\n            bgr = bgr[y0p:y1p, x0p:x1p]\n            a   = a[y0p:y1p, x0p:x1p]\n\n        a_f = (a.astype(np.float32) / 255.0)[..., None]\n        bg  = np.full_like(bgr, CACHE_BG, dtype=np.uint8)\n        comp = (bgr.astype(np.float32) * a_f + bg.astype(np.float32) * (1.0 - a_f)).astype(np.uint8)\n    else:\n        comp = img[:, :, :3].copy()\n        if comp.dtype != np.uint8:\n            comp = np.clip(comp, 0, 255).astype(np.uint8)\n\n    if getattr(Config, \"cache_resize\", True):\n        if getattr(Config, \"cache_letterbox\", False):\n            comp = _letterbox_to_square_bgr(comp, side=CACHE_SIDE, fill=CACHE_BG)\n        else:\n            if comp.shape[0] != CACHE_SIDE or comp.shape[1] != CACHE_SIDE:\n                comp = cv2.resize(comp, (CACHE_SIDE, CACHE_SIDE), interpolation=cv2.INTER_AREA)\n\n    _atomic_imwrite(out_path, comp)\n    return 1\n\ndef build_crop_cache_fast(filenames, src_dir, cache_dir, pad_frac=0.08, max_workers=8):\n    src_dir   = Path(src_dir)\n    cache_dir = Path(cache_dir)\n    cache_dir.mkdir(parents=True, exist_ok=True)\n\n    total, used, free = shutil.disk_usage(\"/kaggle/working\")\n    if free < 500 * 1024 * 1024:\n        print(f\"❌ Not enough free space in /kaggle/working: {free/1e9:.3f} GB free.\")\n        return\n\n    filenames = list(dict.fromkeys(list(filenames)))\n    disk = {p.name for p in src_dir.glob(\"*\") if p.is_file()}\n    filenames = [fn for fn in filenames if fn in disk]\n\n    missing = [fn for fn in filenames if not (cache_dir / fn).exists()]\n    print(f\"Cache dir: {cache_dir} | need to build: {len(missing)}/{len(filenames)}\")\n    if len(missing) == 0:\n        print(\"✅ Cache already complete.\")\n        return\n\n    max_workers = int(min(max_workers, mp.cpu_count()))\n    tasks = [(fn, str(src_dir), str(cache_dir), float(pad_frac)) for fn in missing]\n\n    with ProcessPoolExecutor(max_workers=max_workers) as ex:\n        for _ in tqdm(ex.map(_cache_one_cv2, tasks), total=len(tasks), desc=f\"Cache {cache_dir.name}\"):\n            pass\n\n    print(\"✅ Done caching:\", cache_dir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:39.895244Z","iopub.execute_input":"2026-03-05T04:24:39.895999Z","iopub.status.idle":"2026-03-05T04:24:39.913123Z","shell.execute_reply.started":"2026-03-05T04:24:39.895968Z","shell.execute_reply":"2026-03-05T04:24:39.912291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 7 — Transforms (use ImageNet mean/std for DINO)\n# ============================================================\n\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\nclass ResizeIfNeeded:\n    def __init__(self, size):\n        self.size = int(size)\n    def __call__(self, img: Image.Image):\n        if img.size == (self.size, self.size):\n            return img\n        return img.resize((self.size, self.size))\n        \nclass CoarseBlockMask(torch.nn.Module):\n    \"\"\"\n    Structured occlusion on a tensor image (C,H,W).\n    Fills 1-2 rectangles with random noise (in normalized space).\n    \"\"\"\n    def __init__(self, p=0.22, n_holes=(1, 2), size_frac=(0.18, 0.45)):\n        super().__init__()\n        self.p = float(p)\n        self.n_holes = n_holes\n        self.size_frac = size_frac\n\n    def forward(self, x: torch.Tensor):\n        if (not torch.is_tensor(x)) or x.ndim != 3:\n            return x\n        if random.random() > self.p:\n            return x\n\n        C, H, W = x.shape\n        n = random.randint(int(self.n_holes[0]), int(self.n_holes[1]))\n\n        for _ in range(n):\n            fh = random.uniform(*self.size_frac)\n            fw = random.uniform(*self.size_frac)\n            hh = max(1, int(H * fh))\n            ww = max(1, int(W * fw))\n            y0 = random.randint(0, max(0, H - hh))\n            x0 = random.randint(0, max(0, W - ww))\n\n            # random noise in \"normalized\" space (roughly matches your tensor stats)\n            noise = torch.randn((C, hh, ww), device=x.device, dtype=x.dtype) * 0.5\n            x[:, y0:y0+hh, x0:x0+ww] = noise\n\n        return x\n\n\ntrain_transform = transforms.Compose([\n    # ✅ more close-ups + less extreme aspect distortion\n    transforms.RandomResizedCrop(\n        Config.img_size,\n        scale=(0.25, 1.00),      # was 0.35 -> allow more zoom-in close-ups\n        ratio=(0.75, 1.33)       # was (0.70,1.40) -> reduce weird stretching\n    ),\n    transforms.RandomHorizontalFlip(p=0.5),\n\n    # ✅ keep color jitter but slightly calmer (spots/contrast matter)\n    transforms.RandomApply([transforms.ColorJitter(0.18, 0.18, 0.10, 0.02)], p=0.8),\n    transforms.RandomGrayscale(p=0.05),\n\n    # ✅ mild pose/translation jitter (helps cutout mismatch)\n    transforms.RandomApply([\n        transforms.RandomAffine(\n            degrees=10,\n            translate=(0.06, 0.06),\n            scale=(0.85, 1.15)\n        )\n    ], p=0.7),\n\n    transforms.ToTensor(),\n    transforms.Normalize(MEAN, STD),\n\n    # ✅ structured occlusion (the big missing-piece failure mode)\n    CoarseBlockMask(p=0.22, n_holes=(1, 2), size_frac=(0.18, 0.45)),\n\n    # ✅ small erasing for minor artifacts (keep but slightly reduce)\n    transforms.RandomErasing(\n        p=0.25,\n        scale=(0.02, 0.12),\n        ratio=(0.3, 3.3),\n        value=\"random\"  # if your torchvision doesn't support this, remove value=...\n    ),\n])\n\n\ntest_transform = transforms.Compose([\n    ResizeIfNeeded(Config.img_size),\n    transforms.ToTensor(),\n    transforms.Normalize(MEAN, STD),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:41.383365Z","iopub.execute_input":"2026-03-05T04:24:41.383769Z","iopub.status.idle":"2026-03-05T04:24:41.396243Z","shell.execute_reply.started":"2026-03-05T04:24:41.383739Z","shell.execute_reply":"2026-03-05T04:24:41.395475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 8 — Dataset + Sampler\n# ============================================================\n\nclass JaguarDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None, is_test=False, label_map=None, cache_dir=None):\n        self.df = df.reset_index(drop=True).copy()\n        self.img_dir = Path(img_dir)\n        self.cache_dir = Path(cache_dir) if cache_dir is not None else None\n        self.transform = transform\n        self.is_test = is_test\n\n        if not is_test:\n            assert label_map is not None\n            self.df[\"label\"] = self.df[\"ground_truth\"].map(label_map).astype(int)\n\n    def __len__(self):\n        return len(self.df)\n\n    def _open(self, fname):\n        if self.cache_dir is not None:\n            p = self.cache_dir / fname\n            if p.exists():\n                return Image.open(p).convert(\"RGB\")\n\n        p = self.img_dir / fname\n        img = Image.open(p)\n        img = preprocess_pil(img)\n        return img.convert(\"RGB\")\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        fname = row[\"filename\"]\n\n        try:\n            img = self._open(fname)\n        except Exception:\n            img = Image.new(\"RGB\", (Config.img_size, Config.img_size), (128, 128, 128))\n\n        if self.transform:\n            img = self.transform(img)\n\n        if self.is_test:\n            return img, fname\n        return img, torch.tensor(int(row[\"label\"]), dtype=torch.long)\n\nclass PKBatchSampler(Sampler):\n    # No-repeat cycling + set_epoch (helps imbalance)\n    def __init__(self, labels, P, K, steps_per_epoch, seed=42):\n        self.labels = np.asarray(labels).astype(int)\n        self.P = int(P)\n        self.K = int(K)\n        self.steps_per_epoch = int(steps_per_epoch)\n        self.seed = int(seed)\n        self.epoch = 0\n\n        self.label_to_indices = {}\n        for i, y in enumerate(self.labels):\n            self.label_to_indices.setdefault(int(y), []).append(i)\n        self.unique_labels = np.array(sorted(self.label_to_indices.keys()))\n        self._reset()\n\n    def _reset(self):\n        self.rng = np.random.RandomState(self.seed + self.epoch)\n        self.pools = {}\n        self.ptrs  = {}\n        for y, idxs in self.label_to_indices.items():\n            idxs = np.array(idxs)\n            self.rng.shuffle(idxs)\n            self.pools[y] = idxs\n            self.ptrs[y] = 0\n\n    def set_epoch(self, epoch):\n        self.epoch = int(epoch)\n        self._reset()\n\n    def __len__(self):\n        return self.steps_per_epoch\n\n    def _take_k(self, y):\n        idxs = self.pools[y]\n        n = len(idxs)\n\n        if n <= self.K:\n            if n == 1:\n                return [int(idxs[0])] * self.K\n            pick = self.rng.choice(idxs, size=self.K, replace=True)\n            return [int(x) for x in pick]\n\n        out = []\n        while len(out) < self.K:\n            p = self.ptrs[y]\n            remain = n - p\n            need = self.K - len(out)\n            take = min(remain, need)\n            out.extend(idxs[p:p+take].tolist())\n            self.ptrs[y] += take\n\n            if self.ptrs[y] >= n:\n                self.rng.shuffle(idxs)\n                self.pools[y] = idxs\n                self.ptrs[y] = 0\n\n        return [int(x) for x in out]\n\n    def __iter__(self):\n        for _ in range(self.steps_per_epoch):\n            chosen = self.rng.choice(self.unique_labels, size=self.P, replace=(len(self.unique_labels) < self.P))\n            batch = []\n            for y in chosen:\n                batch.extend(self._take_k(int(y)))\n            yield batch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:44.231343Z","iopub.execute_input":"2026-03-05T04:24:44.232201Z","iopub.status.idle":"2026-03-05T04:24:44.246234Z","shell.execute_reply.started":"2026-03-05T04:24:44.232165Z","shell.execute_reply":"2026-03-05T04:24:44.245485Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 9 — Split (Near-duplicates) + StratifiedGroupKFold (BEST fold + force IDs)\n# ============================================================\n\nTRAIN_CSV = \"/kaggle/input/jaguar-re-id/train.csv\"\nTEST_CSV  = \"/kaggle/input/jaguar-re-id/test.csv\"\nTRAIN_DIR = \"/kaggle/input/jaguar-re-id/train/train\"\nTEST_DIR  = \"/kaggle/input/jaguar-re-id/test/test\"\n\ntrain_df = pd.read_csv(TRAIN_CSV)\ntest_df  = pd.read_csv(TEST_CSV)\n\nunique_ids = sorted(train_df[\"ground_truth\"].unique())\nlabel_map = {name: i for i, name in enumerate(unique_ids)}\nnum_classes = len(unique_ids)\nprint(\"num_classes:\", num_classes)\n\nPHASH_THRESH = 11\nVAL_FRAC = float(getattr(Config, \"val_frac_per_id\", 0.2))\nRNG_SEED = int(getattr(Config, \"seed\", 43))\nNPROC = min(8, mp.cpu_count())\nFORCE_INCLUDE_MISSING_IDS = True\n\n# (optional) stabilize order\ntrain_df = train_df.sort_values(\"filename\").reset_index(drop=True)\n\nTRAIN_DIR_P = Path(TRAIN_DIR)\n\ndef _calc_phash(args):\n    fn, img_dir = args\n    try:\n        p = Path(img_dir) / fn\n        with Image.open(p) as img:\n            img = img.convert(\"RGB\") if img.mode != \"RGB\" else img\n            h = imagehash.phash(img)  # 64-bit\n            return fn, h.hash.flatten().astype(np.uint8)\n    except Exception:\n        return fn, None\n\nprint(f\"Computing pHash with {NPROC} workers on {len(train_df)} images...\")\nargs_list = [(fn, str(TRAIN_DIR_P)) for fn in train_df[\"filename\"].tolist()]\n\nwith mp.Pool(processes=NPROC) as pool:\n    results = list(tqdm(pool.imap(_calc_phash, args_list), total=len(args_list)))\n\nvalid = [(fn, h) for (fn, h) in results if h is not None]\nfilenames = [x[0] for x in valid]\nH = np.stack([x[1] for x in valid], axis=0).astype(np.int8)  # (N,64)\n\nprint(\"Valid hashed images:\", len(filenames))\nprint(\"Computing Hamming distance matrix...\")\ndist = cdist(H, H, metric=\"hamming\") * H.shape[1]\n\nprint(\"Building duplicate groups with threshold:\", PHASH_THRESH)\nadj = (dist <= PHASH_THRESH)\nnp.fill_diagonal(adj, False)\n\ngraph = csr_matrix(adj)\nn_comp, comp_labels = connected_components(csgraph=graph, directed=False, return_labels=True)\nprint(\"Connected components:\", n_comp)\n\nimage_to_group = {}\ndup_groups = 0\nfor c in range(n_comp):\n    idxs = np.where(comp_labels == c)[0]\n    if len(idxs) > 1:\n        dup_groups += 1\n        gid = f\"group_{c}\"\n        for i in idxs:\n            image_to_group[filenames[i]] = gid\n\nall_files = set(train_df[\"filename\"].tolist())\nfor fn in (all_files - set(image_to_group.keys())):\n    image_to_group[fn] = fn\n\ntrain_df = train_df.copy()\ntrain_df[\"group_id\"] = train_df[\"filename\"].map(image_to_group)\ntrain_df[\"label\"] = train_df[\"ground_truth\"].map(label_map).astype(int)\nprint(\"Duplicate groups found:\", dup_groups)\n\n# --- choose best fold\nn_splits = max(2, int(round(1.0 / max(VAL_FRAC, 1e-6))))\nprint(f\"StratifiedGroupKFold: n_splits={n_splits} (val≈{1.0/n_splits:.3f}, requested {VAL_FRAC:.3f})\")\n\nsgkf = StratifiedGroupKFold(n_splits=n_splits, shuffle=True, random_state=RNG_SEED)\ny = train_df[\"label\"].values\ng = train_df[\"group_id\"].values\n\nbest = None\nfor fold, (tr_idx, va_idx) in enumerate(sgkf.split(train_df, y=y, groups=g)):\n    va_ids = train_df.iloc[va_idx][\"label\"].nunique()\n    missing = num_classes - va_ids\n    frac = len(va_idx) / len(train_df)\n    size_err = abs(frac - VAL_FRAC)\n    key = (missing, size_err)\n    if best is None or key < best[\"key\"]:\n        best = {\"fold\": fold, \"tr_idx\": tr_idx, \"va_idx\": va_idx, \"key\": key, \"va_ids\": va_ids, \"frac\": frac}\n\nprint(\"Chosen fold:\", best[\"fold\"], \"| val_frac:\", round(best[\"frac\"], 3), \"| IDs in VAL:\", best[\"va_ids\"], \"/\", num_classes)\n\nval_mask = np.zeros(len(train_df), dtype=bool)\nval_mask[best[\"va_idx\"]] = True\n\n# --- force missing IDs into VAL by moving 1 whole group (if possible)\nif FORCE_INCLUDE_MISSING_IDS:\n    all_ids = set(train_df[\"label\"].unique())\n    val_ids = set(train_df.loc[val_mask, \"label\"].unique())\n    missing_ids = sorted(list(all_ids - val_ids))\n\n    if len(missing_ids) > 0:\n        groups_per_id = train_df.groupby(\"label\")[\"group_id\"].nunique().to_dict()\n        moved = 0\n        impossible = []\n        for ymiss in missing_ids:\n            if groups_per_id.get(ymiss, 0) < 2:\n                impossible.append(ymiss)\n                continue\n            sub = train_df[(train_df[\"label\"] == ymiss) & (~val_mask)]\n            if len(sub) == 0:\n                impossible.append(ymiss)\n                continue\n            vc = sub[\"group_id\"].value_counts(ascending=True)\n            pick_group = vc.index[0]\n            move_idx = train_df.index[train_df[\"group_id\"] == pick_group].values\n            val_mask[move_idx] = True\n            moved += 1\n        print(\"Forced moved groups for missing IDs:\", moved)\n        if len(impossible) > 0:\n            print(\"⚠️ Impossible IDs:\", len(impossible))\n    else:\n        print(\"✅ No missing IDs in VAL.\")\n\ntr_df = train_df[~val_mask].reset_index(drop=True)\nva_df = train_df[val_mask].reset_index(drop=True)\n\nleak = set(tr_df[\"group_id\"]).intersection(set(va_df[\"group_id\"]))\nprint(\"train:\", len(tr_df), \"val:\", len(va_df))\nprint(\"Leaked groups (should be 0):\", len(leak))\nassert len(leak) == 0\n\nprint(f\"IDs in VAL: {va_df['label'].nunique()}/{num_classes}\")\nprint(f\"IDs in TR : {tr_df['label'].nunique()}/{num_classes}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T03:32:29.411298Z","iopub.execute_input":"2026-03-05T03:32:29.411827Z","iopub.status.idle":"2026-03-05T03:37:07.534875Z","shell.execute_reply.started":"2026-03-05T03:32:29.411799Z","shell.execute_reply":"2026-03-05T03:37:07.533989Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# REPLACE CELL 10 — Build caches + DataLoaders (FAST fallback split)\n#   + PSEUDO-LABEL SUPPORT (test images -> train only)\n#   - If tr_df/va_df exist: uses them\n#   - Else: tries to LOAD saved split CSVs (fast)\n#   - Else: creates a FAST per-ID stratified split (no pHash grouping)\n#   - Loads pseudo labels CSV (if exists) and appends to tr_df ONLY\n#   - Creates MIX_DIR with symlinks so JaguarDataset can load both train+test images\n#   - Skips cache build if already complete\n# ============================================================\n\nfrom pathlib import Path\nimport pandas as pd\nimport numpy as np\nimport os, time\nimport cv2\n\n# ----------------------------\n# FIX (ONLY): OpenCV imwrite error for \".tmp\" extension in _atomic_imwrite\n# ----------------------------\ndef _atomic_imwrite(out_path, bgr, jpg_quality=95, png_compression=3):\n    out_path = Path(out_path)\n    out_path.parent.mkdir(parents=True, exist_ok=True)\n\n    tmp_path = out_path.parent / f\"{out_path.stem}.__tmp_{os.getpid()}_{time.time_ns()}{out_path.suffix}\"\n\n    ext = out_path.suffix.lower()\n    params = []\n    if ext in [\".jpg\", \".jpeg\"]:\n        params = [int(cv2.IMWRITE_JPEG_QUALITY), int(jpg_quality)]\n    elif ext == \".png\":\n        params = [int(cv2.IMWRITE_PNG_COMPRESSION), int(png_compression)]\n\n    ok = cv2.imwrite(str(tmp_path), bgr, params)\n    if not ok:\n        raise RuntimeError(f\"cv2.imwrite failed for: {tmp_path} (ext={ext})\")\n\n    os.replace(str(tmp_path), str(out_path))\n    return str(out_path)\n\n# ---- paths\nTRAIN_CSV = \"/kaggle/input/jaguar-re-id/train.csv\"\nTEST_CSV  = \"/kaggle/input/jaguar-re-id/test.csv\"\nTRAIN_DIR_REAL = Path(\"/kaggle/input/jaguar-re-id/train/train\")\nTEST_DIR_REAL  = Path(\"/kaggle/input/jaguar-re-id/test/test\")\n\n# ---- optional saved splits (recommended)\nSAVED_TR = \"/kaggle/working/tr_df.csv\"\nSAVED_VA = \"/kaggle/working/va_df.csv\"\n\n# ---- pseudo labels (your path)\nPSEUDO_CSV = \"/kaggle/input/datasets/ddsdsfdsfsfa/bbbbbb/pseudo_labels_longtail.csv\"\nUSE_PSEUDO = True  # set False to disable\n\n# ensure dfs exist\nif \"train_df\" not in globals():\n    train_df = pd.read_csv(TRAIN_CSV)\nif \"test_df\" not in globals():\n    test_df = pd.read_csv(TEST_CSV)\n\n# ----------------------------\n# Build label_map / num_classes if missing\n# ----------------------------\nif \"label_map\" not in globals():\n    unique_ids = sorted(train_df[\"ground_truth\"].unique())\n    label_map = {name: i for i, name in enumerate(unique_ids)}\n    num_classes = len(unique_ids)\n    print(\"Built label_map | num_classes:\", num_classes)\nelse:\n    if \"num_classes\" not in globals():\n        num_classes = len(label_map)\n\n# ensure train_df has label\nif \"label\" not in train_df.columns:\n    train_df = train_df.copy()\n    train_df[\"label\"] = train_df[\"ground_truth\"].map(label_map).astype(int)\n\n# ----------------------------\n# Load pseudo labels (optional)\n# IMPORTANT: these filenames are from TEST_DIR, not TRAIN_DIR.\n# We'll symlink them into MIX_DIR later so JaguarDataset can read them.\n# ----------------------------\npseudo_df = None\nif USE_PSEUDO and Path(PSEUDO_CSV).exists():\n    pseudo_df = pd.read_csv(PSEUDO_CSV)\n    print(\"✅ Loaded pseudo CSV:\", PSEUDO_CSV, \"| rows:\", len(pseudo_df))\n\n    # must have filename\n    if \"filename\" not in pseudo_df.columns:\n        raise KeyError(f\"Pseudo CSV must contain 'filename'. cols={list(pseudo_df.columns)}\")\n\n    # ensure label exists\n    if \"label\" not in pseudo_df.columns:\n        if \"ground_truth\" in pseudo_df.columns:\n            pseudo_df = pseudo_df.copy()\n            pseudo_df[\"label\"] = pseudo_df[\"ground_truth\"].map(label_map).astype(int)\n        else:\n            raise KeyError(\"Pseudo CSV must contain either 'label' or 'ground_truth'.\")\n\n    pseudo_df[\"label\"] = pseudo_df[\"label\"].astype(int)\n    if \"ground_truth\" not in pseudo_df.columns:\n        inv_label_map = {v: k for k, v in label_map.items()}\n        pseudo_df[\"ground_truth\"] = pseudo_df[\"label\"].map(inv_label_map).astype(str)\n\n    pseudo_df[\"is_pseudo\"] = True\n    pseudo_df[\"filename\"] = pseudo_df[\"filename\"].astype(str)\n\nelse:\n    print(\"ℹ️ Pseudo CSV not found or disabled.\")\n    pseudo_df = None\n\n# ----------------------------\n# Get / build tr_df, va_df (ONLY from real train_df)\n# ----------------------------\ndef fast_per_id_split(df: pd.DataFrame, val_frac: float = 0.2, seed: int = 42):\n    rng = np.random.RandomState(seed)\n    tr_idx, va_idx = [], []\n\n    for y, sub in df.groupby(\"label\"):\n        idxs = sub.index.values.copy()\n        rng.shuffle(idxs)\n\n        n = len(idxs)\n        if n <= 1:\n            tr_idx.extend(idxs.tolist())\n            continue\n\n        n_val = int(round(n * val_frac))\n        n_val = max(1, n_val)\n        n_val = min(n_val, n - 1)  # keep >=1 in train\n\n        va_idx.extend(idxs[:n_val].tolist())\n        tr_idx.extend(idxs[n_val:].tolist())\n\n    tr_df_ = df.loc[tr_idx].reset_index(drop=True)\n    va_df_ = df.loc[va_idx].reset_index(drop=True)\n    return tr_df_, va_df_\n\nif (\"tr_df\" not in globals()) or (\"va_df\" not in globals()):\n    if Path(SAVED_TR).exists() and Path(SAVED_VA).exists():\n        tr_df = pd.read_csv(SAVED_TR)\n        va_df = pd.read_csv(SAVED_VA)\n        print(f\"✅ Loaded saved split: {SAVED_TR} + {SAVED_VA}\")\n    else:\n        print(\"⚠️ tr_df/va_df not found and no saved split CSVs.\")\n        print(\"   Creating FAST per-ID split (no pHash grouping).\")\n        tr_df, va_df = fast_per_id_split(\n            train_df,\n            val_frac=float(getattr(Config, \"val_frac_per_id\", 0.2)),\n            seed=int(getattr(Config, \"seed\", 42)),\n        )\n        tr_df.to_csv(SAVED_TR, index=False)\n        va_df.to_csv(SAVED_VA, index=False)\n        print(f\"✅ Saved FAST split to: {SAVED_TR} and {SAVED_VA}\")\n\n# ensure split dfs have label\nfor _df_name in [\"tr_df\", \"va_df\"]:\n    _df = globals()[_df_name]\n    if \"label\" not in _df.columns:\n        assert \"ground_truth\" in _df.columns, f\"{_df_name} missing label and ground_truth\"\n        _df = _df.copy()\n        _df[\"label\"] = _df[\"ground_truth\"].map(label_map).astype(int)\n        globals()[_df_name] = _df\n\nprint(\"Split sizes | train:\", len(tr_df), \"val:\", len(va_df),\n      \"| IDs in train:\", tr_df[\"label\"].nunique(), \"IDs in val:\", va_df[\"label\"].nunique())\n\n# ----------------------------\n# Append pseudo labels to TRAIN ONLY (never to val)\n# ----------------------------\nif pseudo_df is not None and len(pseudo_df) > 0:\n    # avoid dup filenames if any\n    tr_names = set(tr_df[\"filename\"].astype(str).tolist())\n    pseudo_df2 = pseudo_df[~pseudo_df[\"filename\"].isin(tr_names)].copy()\n\n    # keep columns consistent\n    for c in tr_df.columns:\n        if c not in pseudo_df2.columns:\n            pseudo_df2[c] = np.nan\n    for c in pseudo_df2.columns:\n        if c not in tr_df.columns:\n            tr_df[c] = np.nan\n\n    tr_df = pd.concat([tr_df, pseudo_df2[tr_df.columns]], ignore_index=True)\n    print(f\"✅ Added pseudo rows to tr_df: +{len(pseudo_df2)} | new tr_df size: {len(tr_df)}\")\nelse:\n    print(\"ℹ️ No pseudo rows appended.\")\n\n# ----------------------------\n# Create MIX_DIR so JaguarDataset can load both TRAIN + pseudo(TEST) images\n# (keeps your dataset code unchanged)\n# ----------------------------\nMIX_DIR = Path(\"/kaggle/working/mix_images\")\nMIX_DIR.mkdir(parents=True, exist_ok=True)\n\ndef _safe_link(src: Path, dst: Path):\n    if dst.exists():\n        return\n    try:\n        os.symlink(str(src), str(dst))\n    except FileExistsError:\n        return\n    except OSError:\n        # fallback: copy if symlink fails\n        import shutil\n        shutil.copy2(str(src), str(dst))\n\n# detect collisions (same filename exists in both train/test)\ntrain_disk = {p.name for p in TRAIN_DIR_REAL.glob(\"*\") if p.is_file()}\ntest_disk  = {p.name for p in TEST_DIR_REAL.glob(\"*\") if p.is_file()}\noverlap = sorted(list(train_disk & test_disk))\n\n# If overlap exists, rename TEST pseudo filenames to avoid collisions in MIX_DIR.\n# (This only affects training; submission still uses original filenames.)\nif overlap and (pseudo_df is not None) and len(pseudo_df) > 0:\n    print(f\"⚠️ Found {len(overlap)} filename collisions between train/test. Renaming pseudo filenames with 'test__' prefix.\")\n    # map only colliding pseudo files\n    ren = {}\n    for fn in overlap:\n        ren[fn] = \"test__\" + fn\n    # apply rename in tr_df ONLY for rows that are pseudo AND collide\n    if \"is_pseudo\" in tr_df.columns:\n        m = (tr_df[\"is_pseudo\"].fillna(False).astype(bool)) & (tr_df[\"filename\"].isin(list(ren.keys())))\n        tr_df.loc[m, \"filename\"] = tr_df.loc[m, \"filename\"].map(lambda x: ren.get(x, x))\n\n# Now link all files referenced by tr_df + va_df into MIX_DIR\nneed_files = set(tr_df[\"filename\"].astype(str).tolist()) | set(va_df[\"filename\"].astype(str).tolist())\n\nlinked = 0\nmissing = 0\nfor fn in sorted(need_files):\n    # if it's a renamed pseudo file, recover original test filename\n    src_candidates = []\n    if fn.startswith(\"test__\"):\n        src_candidates.append(TEST_DIR_REAL / fn.replace(\"test__\", \"\", 1))\n    else:\n        src_candidates.append(TRAIN_DIR_REAL / fn)\n        src_candidates.append(TEST_DIR_REAL / fn)\n\n    src = None\n    for cand in src_candidates:\n        if cand.exists():\n            src = cand\n            break\n\n    if src is None:\n        missing += 1\n        continue\n\n    _safe_link(src, MIX_DIR / fn)\n    linked += 1\n\nprint(f\"✅ MIX_DIR ready: {MIX_DIR} | linked={linked} | missing={missing}\")\n\n# ----------------------------\n# Cache dirs (NO-CROP mode = just resized copies, not bbox cropping)\n# ----------------------------\nMODE = \"crop\" if getattr(Config, \"use_crop\", False) else \"nocrop\"\nTRAIN_CACHE = Path(Config.cache_root) / MODE / \"train\"\nTEST_CACHE  = Path(Config.cache_root) / MODE / \"test\"\nTRAIN_CACHE.mkdir(parents=True, exist_ok=True)\nTEST_CACHE.mkdir(parents=True, exist_ok=True)\n\nprint(\"CACHE MODE:\", MODE)\nprint(\"TRAIN_CACHE:\", TRAIN_CACHE)\nprint(\"TEST_CACHE :\", TEST_CACHE)\n\ndef get_test_filenames(df: pd.DataFrame):\n    cols = set(df.columns)\n    if \"filename\" in cols:\n        return df[\"filename\"].astype(str).unique().tolist()\n    if (\"query_image\" in cols) and (\"gallery_image\" in cols):\n        return sorted(set(df[\"query_image\"].astype(str).tolist()) | set(df[\"gallery_image\"].astype(str).tolist()))\n    raise KeyError(f\"Cannot infer test filenames. Columns are: {list(df.columns)}\")\n\ntrain_files = tr_df[\"filename\"].astype(str).tolist()\ntest_files  = get_test_filenames(test_df)\n\ndef cache_status_fast(cache_dir: Path, filenames):\n    filenames = list(dict.fromkeys([str(x) for x in filenames]))\n    have = {p.name for p in cache_dir.glob(\"*\") if p.is_file()}\n    existing = sum((fn in have) for fn in filenames)\n    missing = len(filenames) - existing\n    return existing, missing, len(filenames)\n\n# ----------------------------\n# Build caches (only missing)\n# NOTE: train cache is built from MIX_DIR now (so pseudo files get cached too)\n# ----------------------------\nif bool(getattr(Config, \"cache_crops\", True)):\n    ex_tr, miss_tr, tot_tr = cache_status_fast(TRAIN_CACHE, train_files)\n    ex_te, miss_te, tot_te = cache_status_fast(TEST_CACHE,  test_files)\n\n    print(f\"TRAIN cache: {ex_tr}/{tot_tr} exist | missing {miss_tr}\")\n    print(f\"TEST  cache: {ex_te}/{tot_te} exist | missing {miss_te}\")\n\n    if miss_tr > 0:\n        build_crop_cache_fast(\n            train_files, str(MIX_DIR), str(TRAIN_CACHE),\n            pad_frac=float(getattr(Config, \"cache_pad_frac\", 0.08)),\n            max_workers=int(getattr(Config, \"cache_workers\", 8))\n        )\n    else:\n        print(\"✅ TRAIN cache complete. Skipping build.\")\n\n    if miss_te > 0:\n        build_crop_cache_fast(\n            test_files, str(TEST_DIR_REAL), str(TEST_CACHE),\n            pad_frac=float(getattr(Config, \"cache_pad_frac\", 0.08)),\n            max_workers=int(getattr(Config, \"cache_workers\", 8))\n        )\n    else:\n        print(\"✅ TEST cache complete. Skipping build.\")\n\n# ----------------------------\n# Datasets (train/val read from MIX_DIR; val uses only real train rows)\n# ----------------------------\ntrain_ds = JaguarDataset(\n    tr_df, str(MIX_DIR), transform=train_transform, is_test=False,\n    label_map=label_map, cache_dir=(TRAIN_CACHE if Config.cache_crops else None)\n)\nval_ds = JaguarDataset(\n    va_df, str(MIX_DIR), transform=test_transform, is_test=False,\n    label_map=label_map, cache_dir=(TRAIN_CACHE if Config.cache_crops else None)\n)\n\n# ----------------------------\n# steps_per_epoch + PK sampler\n# ----------------------------\nBATCH = int(Config.P) * int(Config.K)\nbase_steps = max(1, len(train_ds) // BATCH)\nsteps_per_epoch = int(base_steps * int(Config.steps_mult))\nprint(\"steps_per_epoch:\", steps_per_epoch, \"| batch:\", BATCH, \"| steps_mult:\", Config.steps_mult)\n\ntrain_sampler = PKBatchSampler(\n    train_ds.df[\"label\"].values,\n    P=int(Config.P), K=int(Config.K),\n    steps_per_epoch=int(steps_per_epoch),\n    seed=int(Config.seed)\n)\n\ntrain_loader = DataLoader(\n    train_ds,\n    batch_sampler=train_sampler,\n    num_workers=int(Config.num_workers),\n    pin_memory=bool(Config.pin_memory),\n    persistent_workers=bool(Config.persistent_workers) and int(Config.num_workers) > 0,\n    prefetch_factor=int(Config.prefetch_factor) if int(Config.num_workers) > 0 else None,\n)\n\nval_eval_loader = DataLoader(\n    val_ds,\n    batch_size=64,\n    shuffle=False,\n    num_workers=0,   # keep 0 for stable eval + debugging\n    pin_memory=True,\n    persistent_workers=False,\n    prefetch_factor=None,\n)\n\nval_labels = val_ds.df[\"label\"].values.astype(int)\nprint(\"✅ Loaders ready | train:\", len(train_ds), \"val:\", len(val_ds))\nprint(\"✅ Pseudo in train_ds:\", int(train_ds.df.get(\"is_pseudo\", False).fillna(False).sum()) if hasattr(train_ds, \"df\") else \"unknown\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:47.463907Z","iopub.execute_input":"2026-03-05T04:24:47.464670Z","iopub.status.idle":"2026-03-05T04:24:50.713976Z","shell.execute_reply.started":"2026-03-05T04:24:47.464636Z","shell.execute_reply":"2026-03-05T04:24:50.713351Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# REPLACE CELL 11 — ArcFace heads (AMP-safe)\n#   Fix: compute logits in FP32 so scatter_ never mismatches\n# ============================================================\n\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass DynamicSubCenterArcFace(nn.Module):\n    def __init__(self, in_features, out_features, margins, k=3, s=30.0, easy_margin=False):\n        super().__init__()\n        self.k = int(k)\n        self.s = float(s)\n        self.easy_margin = bool(easy_margin)\n\n        margins_t = torch.tensor(margins, dtype=torch.float32)\n        assert margins_t.numel() == int(out_features), \"margins length must equal out_features\"\n        self.register_buffer(\"margins\", margins_t)  \n\n        self.weight = nn.Parameter(torch.empty(out_features, self.k, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n    def forward(self, x, label=None):\n        # ✅ FIX: Force FP32 for numerical stability and to prevent scatter dtype mismatch\n        x = x.float()\n        w = self.weight.float()\n\n        x_norm = F.normalize(x, dim=1)\n\n        # sub-centers: [C,k,D] -> [C*k, D]\n        w2 = w.reshape(-1, w.size(-1))\n        w_norm = F.normalize(w2, dim=1)\n\n        cosine = F.linear(x_norm, w_norm)  # [B, C*k] \n        cosine = cosine.reshape(-1, w.size(0), self.k)  # [B,C,k]\n        cosine, _ = torch.max(cosine, dim=2)  # [B,C]\n        cosine = cosine.clamp(-1 + 1e-7, 1 - 1e-7)\n\n        if label is None:\n            return cosine * self.s\n\n        label = label.long().view(-1, 1)  # [B,1]\n\n        # --- target cosine only\n        cos_t = cosine.gather(1, label).squeeze(1)  # [B]\n        sin_t = torch.sqrt(torch.clamp(1.0 - cos_t * cos_t, min=0.0))  # [B]\n\n        # --- per-sample margin\n        m = self.margins[label.squeeze(1)]  # [B]\n\n        cos_m = torch.cos(m)\n        sin_m = torch.sin(m)\n\n        # phi_t = cos(theta + m)\n        phi_t = cos_t * cos_m - sin_t * sin_m\n\n        if self.easy_margin:\n            phi_t = torch.where(cos_t > 0, phi_t, cos_t)\n        else:\n            th = torch.cos(math.pi - m)\n            mm = torch.sin(math.pi - m) * m\n            phi_t = torch.where(cos_t > th, phi_t, cos_t - mm)\n\n        # --- replace only target logit\n        logits = cosine.clone()\n        logits.scatter_(1, label, phi_t.unsqueeze(1).to(logits.dtype)) # Extra safety cast\n        \n        return logits * self.s\n\n\nclass SubCenterArcFace(nn.Module):\n    \"\"\"\n    Standard SubCenter ArcFace with fixed margin m.\n    AMP-safe: head runs in FP32.\n    \"\"\"\n    def __init__(self, in_features, out_features, k=2, s=30.0, m=0.50, easy_margin=False):\n        super().__init__()\n        self.k = int(k)\n        self.s = float(s)\n        self.m = float(m)\n        self.easy_margin = bool(easy_margin)\n\n        self.weight = nn.Parameter(torch.empty(out_features, self.k, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.cos_m = math.cos(self.m)\n        self.sin_m = math.sin(self.m)\n        self.th = math.cos(math.pi - self.m)\n        self.mm = math.sin(math.pi - self.m) * self.m\n\n    def forward(self, x, label=None):\n        x = x.float()\n        w = self.weight.float()\n\n        x_norm = F.normalize(x, dim=1)\n        w2 = w.reshape(-1, w.size(-1))              # [C*k, D]\n        w2 = F.normalize(w2, dim=1)\n\n        cosine = F.linear(x_norm, w2)               # [B, C*k] FP32\n        cosine = cosine.reshape(-1, w.size(0), self.k)\n        cosine, _ = torch.max(cosine, dim=2)        # [B, C] FP32\n        cosine = cosine.clamp(-1 + 1e-7, 1 - 1e-7)\n\n        if label is None:\n            return cosine * self.s\n\n        label = label.long().view(-1, 1)\n\n        sine = torch.sqrt(torch.clamp(1.0 - cosine * cosine, min=0.0))\n        phi = cosine * self.cos_m - sine * self.sin_m\n\n        if self.easy_margin:\n            phi = torch.where(cosine > 0, phi, cosine)\n        else:\n            phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n\n        one_hot = torch.zeros_like(cosine)\n        one_hot.scatter_(1, label, 1.0)\n\n        logits = one_hot * phi + (1.0 - one_hot) * cosine\n        return logits * self.s","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:51.351034Z","iopub.execute_input":"2026-03-05T04:24:51.351702Z","iopub.status.idle":"2026-03-05T04:24:51.366566Z","shell.execute_reply.started":"2026-03-05T04:24:51.351673Z","shell.execute_reply":"2026-03-05T04:24:51.365815Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 11.1 — Build frequency-adaptive margins (FULL CELL) ✅ robust\n#   - frequent class => larger margin\n#   - rare class     => smaller margin\n#   - safe range: [m_min, m_max]\n#   - fixes KeyError: 'label' by creating tr_df['label'] if missing\n# ============================================================\n\nimport numpy as np\nimport torch\n\n# ---- must have split\nassert \"tr_df\" in globals(), \"Need tr_df first (run your split cell to create tr_df/va_df).\"\n\n# ---- ensure train_df + label_map exist\nTRAIN_CSV = globals().get(\"TRAIN_CSV\", \"/kaggle/input/jaguar-re-id/train.csv\")\nif \"train_df\" not in globals():\n    import pandas as pd\n    train_df = pd.read_csv(TRAIN_CSV)\n\nif \"label_map\" not in globals():\n    unique_ids = sorted(train_df[\"ground_truth\"].unique())\n    label_map = {name: i for i, name in enumerate(unique_ids)}\n\nif \"num_classes\" not in globals():\n    num_classes = len(label_map)\n\n# ---- ensure tr_df has integer labels\nif \"label\" not in tr_df.columns:\n    assert \"ground_truth\" in tr_df.columns, f\"tr_df missing both 'label' and 'ground_truth'. cols={list(tr_df.columns)}\"\n    tr_df = tr_df.copy()\n    tr_df[\"label\"] = tr_df[\"ground_truth\"].map(label_map).astype(int)\n\n# ----------------------------\n# counts per class\n# ----------------------------\ncounts = np.bincount(tr_df[\"label\"].values.astype(int), minlength=int(num_classes)).astype(np.float32)\n\n# ---- knobs (safe defaults)\nm_min = 0.35   # rare\nm_max = 0.55   # frequent\n\n# normalize by log-counts (stable)\nx = np.log(counts + 1.0)\nx = (x - x.min()) / (x.max() - x.min() + 1e-12)\n\nmargins = (m_min + (m_max - m_min) * x).astype(np.float32)\n\nprint(\"num_classes:\", int(num_classes))\nprint(\"counts min/max:\", int(counts.min()), int(counts.max()))\nprint(\"margins min/max:\", float(margins.min()), float(margins.max()))\n\n# keep this name to pass into the head\nCLASS_MARGINS = margins\nprint(\"✅ CLASS_MARGINS ready (len =\", len(CLASS_MARGINS), \")\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:54.033953Z","iopub.execute_input":"2026-03-05T04:24:54.034672Z","iopub.status.idle":"2026-03-05T04:24:54.043558Z","shell.execute_reply.started":"2026-03-05T04:24:54.034641Z","shell.execute_reply":"2026-03-05T04:24:54.042891Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# REPLACE CELL 12 — Model (v7: Dual-Head + CLS-mask + Toggleable shared L2-norm)\n#   FIX APPLIED:\n#     ✅ Arc branch now returns TWO representations:\n#        - feat_cls  = L2-normalized PRE-BN arc feature (retrieval-friendly)\n#        - arc_bn    = BNNeck(arc_prebn) used ONLY for ArcFace logits\n#     ✅ This keeps your pipeline (embed_pair returns (feat_tri, feat_cls))\n#        while stopping BNNeck from poisoning retrieval.\n#\n#   Summary:\n#     - embed_pair(x) -> (feat_tri, feat_cls)   # feat_cls is now pre-BN retrieval feature\n#     - forward(x,label) uses BNNeck output for head logits\n# ============================================================\n\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\n\n\nclass GeM(nn.Module):\n    def __init__(self, p=3.0, eps=1e-6):\n        super().__init__()\n        self.p = nn.Parameter(torch.ones(1) * float(p))\n        self.eps = float(eps)\n\n    def forward(self, x):\n        return F.avg_pool2d(\n            x.clamp(min=self.eps).pow(self.p),\n            (x.size(-2), x.size(-1))\n        ).pow(1.0 / self.p)\n\n\ndef _tokens_to_hw(patch_tokens: torch.Tensor):\n    \"\"\"\n    patch_tokens: [B, L, C] -> [B, C, H, W]\n    Handles non-square L by truncation to H*W.\n    \"\"\"\n    B, L, C = patch_tokens.shape\n    H = int(math.sqrt(L))\n    if H * H == L:\n        W = H\n        L2 = L\n    else:\n        H = max(1, int(math.sqrt(L)))\n        W = max(1, L // H)\n        L2 = H * W\n    patch_tokens = patch_tokens[:, :L2, :]\n    return patch_tokens.transpose(1, 2).reshape(B, C, H, W)\n\n\ndef _inv_sigmoid(p):\n    p = float(p)\n    p = min(max(p, 1e-6), 1.0 - 1e-6)\n    return math.log(p / (1.0 - p))\n\n\nclass ReIDBoss(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        print(\"v7-dualhead-cls-toggleL2 (arc preBN retrieval FIX)\")\n\n        # ----------------------------\n        # Backbone\n        # ----------------------------\n        self.backbone = timm.create_model(\n            Config.model_name,\n            pretrained=True,\n            num_classes=0,\n            cache_dir=\"/kaggle/working/hf\"\n        )\n\n        if getattr(Config, \"use_checkpointing\", False):\n            if hasattr(self.backbone, \"set_grad_checkpointing\"):\n                self.backbone.set_grad_checkpointing(True)\n            elif hasattr(self.backbone, \"gradient_checkpointing_enable\"):\n                self.backbone.gradient_checkpointing_enable()\n            print(\"Grad checkpointing ON\")\n\n        self.feat_dim = int(self.backbone.num_features)\n\n        # ----------------------------\n        # Pooling\n        # ----------------------------\n        self.gem = GeM(p=float(getattr(Config, \"gem_p\", 4.0)))\n\n        embed_dim = self.feat_dim if (getattr(Config, \"embed_dim\", None) is None) else int(Config.embed_dim)\n        self.embed_dim = int(embed_dim)\n\n        drop_p = float(getattr(Config, \"neck_dropout\", 0.0))\n        self.dropout = nn.Dropout(p=drop_p) if drop_p > 0 else nn.Identity()\n\n        # Map GeM + CLS to embed_dim\n        self.gem_to_embed = nn.Linear(self.feat_dim, embed_dim, bias=False) if embed_dim != self.feat_dim else nn.Identity()\n        self.cls_to_embed = nn.Linear(self.feat_dim, embed_dim, bias=False) if embed_dim != self.feat_dim else nn.Identity()\n\n        # Optional stabilizer (does NOT enforce L2)\n        self.shared_ln = nn.LayerNorm(embed_dim) if bool(getattr(Config, \"shared_use_ln\", False)) else nn.Identity()\n\n        # ----------------------------\n        # Learnable CLS gate\n        # ----------------------------\n        self.max_cls_w = float(getattr(Config, \"max_cls_w\", 0.50))\n        self.gate_mode = str(getattr(Config, \"gate_mode\", \"sample\")).lower()\n        init_cls = float(getattr(Config, \"init_cls_w\", 0.20))\n\n        if self.gate_mode == \"scalar\":\n            self.g_cls_logit = nn.Parameter(torch.tensor(_inv_sigmoid(init_cls / self.max_cls_w), dtype=torch.float32))\n            self.cls_gate_mlp = None\n        else:\n            hidden = int(getattr(Config, \"gate_hidden\", 128))\n            self.cls_gate_mlp = nn.Sequential(\n                nn.Linear(self.feat_dim, hidden),\n                nn.GELU(),\n                nn.Linear(hidden, 1),\n            )\n            with torch.no_grad():\n                self.cls_gate_mlp[-1].bias.fill_(_inv_sigmoid(init_cls / self.max_cls_w))\n\n        # ----------------------------\n        # ArcFace branch (own neck + BNNeck)\n        # ----------------------------\n        arc_neck = str(getattr(Config, \"arc_neck\", \"linear\")).lower()\n        if arc_neck == \"mlp\":\n            self.arc_neck = nn.Sequential(\n                nn.Linear(embed_dim, embed_dim, bias=False),\n                nn.GELU(),\n                nn.Linear(embed_dim, embed_dim, bias=False),\n            )\n        else:\n            self.arc_neck = nn.Linear(embed_dim, embed_dim, bias=False)\n\n        self.bnneck = nn.BatchNorm1d(embed_dim)\n        self.bnneck.bias.requires_grad_(False)\n\n        # ----------------------------\n        # Metric branch (Circle head)\n        # ----------------------------\n        met_drop = float(getattr(Config, \"metric_drop\", 0.10))\n        self.metric_mlp = nn.Sequential(\n            nn.Linear(embed_dim, embed_dim, bias=False),\n            nn.GELU(),\n            nn.Dropout(p=met_drop),\n            nn.Linear(embed_dim, embed_dim, bias=False),\n        )\n        self.metric_ln = nn.LayerNorm(embed_dim) if bool(getattr(Config, \"metric_use_ln\", True)) else nn.Identity()\n\n        # ----------------------------\n        # ArcFace head (dynamic margins if present)\n        # ----------------------------\n        if (\"DynamicSubCenterArcFace\" in globals()) and (\"CLASS_MARGINS\" in globals()):\n            self.head = DynamicSubCenterArcFace(\n                embed_dim, num_classes,\n                margins=CLASS_MARGINS,\n                k=int(getattr(Config, \"subcenter_k\", 3)),\n                s=float(getattr(Config, \"arcface_s\", 30.0)),\n                easy_margin=False\n            )\n        else:\n            self.head = SubCenterArcFace(\n                embed_dim, num_classes,\n                k=int(getattr(Config, \"subcenter_k\", 3)),\n                s=float(getattr(Config, \"arcface_s\", 30.0)),\n                m=float(getattr(Config, \"arcface_m\", 0.50)),\n                easy_margin=False\n            )\n\n    def forward_features_tokens(self, x):\n        \"\"\"\n        Returns:\n          cls_token: [B, C]\n          patch_map: [B, C, H, W] or None\n        \"\"\"\n        feat = self.backbone.forward_features(x)\n\n        if feat.dim() == 2:\n            return feat, None\n\n        if feat.dim() == 3:\n            B, N, C = feat.shape\n            num_prefix = int(getattr(self.backbone, \"num_prefix_tokens\", 1))\n            cls_token = feat[:, :num_prefix, :].mean(dim=1)  # [B,C]\n\n            if N > num_prefix:\n                patch_tokens = feat[:, num_prefix:, :]       # [B,L,C]\n                patch_map = _tokens_to_hw(patch_tokens)      # [B,C,H,W]\n                return cls_token, patch_map\n\n            return cls_token, None\n\n        return feat.flatten(1), None\n\n    def _get_w_cls(self, cls_token_fp32):\n        B = cls_token_fp32.size(0)\n        if self.gate_mode == \"scalar\":\n            return torch.sigmoid(self.g_cls_logit).view(1, 1).expand(B, 1) * self.max_cls_w\n        return torch.sigmoid(self.cls_gate_mlp(cls_token_fp32)) * self.max_cls_w\n\n    def embed_pair(self, x, return_bn=False):\n        \"\"\"\n        Returns:\n          feat_tri: metric embedding (L2)\n          feat_cls: arc embedding for retrieval (PRE-BN, L2)\n          (optional) arc_bn: BNNeck output for ArcFace logits\n        \"\"\"\n        cls_token, patch_map = self.forward_features_tokens(x)\n\n        # ----------------------------\n        # CLS-guided masking -> GeM\n        # ----------------------------\n        use_cls_attn = bool(getattr(Config, \"use_cls_attn_pool\", True))\n        attn_temp    = float(getattr(Config, \"cls_attn_temp\", 0.10))\n        attn_clamp   = float(getattr(Config, \"cls_attn_clamp\", 4.0))\n        attn_mix     = float(getattr(Config, \"cls_attn_mix\", 0.70))\n\n        if use_cls_attn and (patch_map is not None):\n            B, C, H, W = patch_map.shape\n            L = H * W\n\n            cls_norm = F.normalize(cls_token.float(), dim=1)\n            patch_flat = patch_map.float().view(B, C, L)\n            patch_norm = F.normalize(patch_flat, dim=1)\n\n            attn = torch.bmm(cls_norm.unsqueeze(1), patch_norm)          # [B,1,L]\n            attn = attn / max(attn_temp, 1e-6)\n            attn = attn - attn.max(dim=-1, keepdim=True).values\n            w = F.softmax(attn, dim=-1)                                  # [B,1,L]\n\n            # Mean-scaling + Clamp (Stable masking)\n            w = w / (w.mean(dim=-1, keepdim=True) + 1e-6)\n\n            attn_mask = w.view(B, 1, H, W)\n            if attn_clamp > 0:\n                attn_mask = attn_mask.clamp(0.0, attn_clamp)\n\n            masked = patch_map * attn_mask.to(patch_map.dtype)\n\n            gem_fg  = self.gem(masked).flatten(1)\n            gem_all = self.gem(patch_map).flatten(1)\n            gem_feat = attn_mix * gem_fg + (1.0 - attn_mix) * gem_all\n        else:\n            gem_feat = self.gem(patch_map).flatten(1) if (patch_map is not None) else cls_token\n\n        # ----------------------------\n        # Shared fusion: GeM + gated CLS\n        # ----------------------------\n        z_gem = self.gem_to_embed(self.dropout(gem_feat))\n        z_cls = self.cls_to_embed(self.dropout(cls_token))\n        w_cls = self._get_w_cls(cls_token.float())\n\n        shared = F.normalize(z_gem, dim=1) + w_cls.to(z_gem.dtype) * F.normalize(z_cls, dim=1)\n\n        # Toggleable shared L2 norm\n        if bool(getattr(Config, \"shared_l2norm\", True)):\n            shared = F.normalize(shared, dim=1)\n\n        shared = self.shared_ln(shared)\n\n        # ----------------------------\n        # Metric branch (Circle)\n        # ----------------------------\n        met = self.metric_ln(self.metric_mlp(shared))\n        feat_tri = F.normalize(met, dim=1)\n\n        # ----------------------------\n        # Arc branch (ArcFace)\n        #   FIX: retrieval uses PRE-BN arc feature (L2),\n        #        logits use BNNeck output only.\n        # ----------------------------\n        arc_prebn = self.arc_neck(shared)              # [B, D]\n        feat_cls  = F.normalize(arc_prebn, dim=1)      # ✅ retrieval-friendly \"FC\"\n\n        arc_bn = self.bnneck(arc_prebn)                # classifier feature\n\n        if return_bn:\n            return feat_tri, feat_cls, arc_bn\n        return feat_tri, feat_cls\n\n    def forward(self, x, label=None, return_emb=False):\n        # Inference / embedding mode\n        if label is None:\n            feat_tri, _ = self.embed_pair(x, return_bn=False)\n            return feat_tri\n\n        # Training: need BN feature for ArcFace logits\n        feat_tri, feat_cls, arc_bn = self.embed_pair(x, return_bn=True)\n\n        logits = self.head(arc_bn.float(), label.long())\n\n        if return_emb:\n            return feat_tri, feat_cls, logits\n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:54.711220Z","iopub.execute_input":"2026-03-05T04:24:54.711855Z","iopub.status.idle":"2026-03-05T04:24:54.738307Z","shell.execute_reply.started":"2026-03-05T04:24:54.711831Z","shell.execute_reply":"2026-03-05T04:24:54.737625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# REPLACE CELL 13 — Losses (Unweighted CE + FIXED CircleLoss)\n#   Fix: remove self-pairs from negatives properly\n# ============================================================\n\ncriterion_cls = nn.CrossEntropyLoss(label_smoothing=0.0)\nprint(\"✅ Using unweighted CE + label smoothing\")\n\nclass CircleLoss(nn.Module):\n    def __init__(self, m=0.25, gamma=16.0):\n        super().__init__()\n        self.m = float(m)\n        self.gamma = float(gamma)\n\n    def forward(self, emb, labels):\n        emb = F.normalize(emb, dim=1)\n        sim = emb @ emb.t()  # [B,B]\n        labels = labels.view(-1, 1)\n\n        same = (labels == labels.t())\n        eye = torch.eye(sim.size(0), device=sim.device, dtype=torch.bool)\n\n        # ✅ remove diagonal completely\n        pos_mask = same & (~eye)\n        neg_mask = (~same) & (~eye)\n\n        sp = sim[pos_mask]\n        sn = sim[neg_mask]\n\n        if sp.numel() == 0 or sn.numel() == 0:\n            return sim.new_tensor(0.0)\n\n        ap = torch.clamp_min(-sp.detach() + 1 + self.m, 0.)\n        an = torch.clamp_min(sn.detach() + self.m, 0.)\n\n        logit_p = - self.gamma * ap * (sp - (1 - self.m))\n        logit_n =   self.gamma * an * (sn - self.m)\n\n        loss_p = torch.logsumexp(logit_p, dim=0)\n        loss_n = torch.logsumexp(logit_n, dim=0)\n        return F.softplus(loss_p + loss_n)\n\nclass HardCircleLoss(nn.Module):\n    def __init__(self, m=0.25, gamma=16.0, topk_pos=4, topk_neg=60, island_thresh=0.20, debug=False):\n        super().__init__()\n        self.m = float(m)\n        self.gamma = float(gamma)\n        self.topk_pos = int(topk_pos)\n        self.topk_neg = int(topk_neg)\n        self.island_thresh = float(island_thresh)\n        self.debug = bool(debug)\n\n    def forward(self, emb, labels):\n        emb = F.normalize(emb, dim=1)\n        sim = emb @ emb.t()\n\n        labels = labels.view(-1, 1)\n        same = (labels == labels.t())\n        eye = torch.eye(sim.size(0), device=sim.device, dtype=torch.bool)\n\n        is_pos = same & (~eye)\n        is_neg = (~same) & (~eye)\n\n        # cap Kpos by PK batch structure\n        max_kpos = max(1, int(getattr(Config, \"K\", 3)) - 1)\n        k_pos = min(self.topk_pos, max_kpos)\n\n        max_actual_neg = int(is_neg.sum(dim=1).max().item())\n        k_neg = min(self.topk_neg, max_actual_neg)\n\n        if k_pos <= 0 or k_neg <= 0:\n            return sim.new_tensor(0.0)\n\n        # --- island mask\n        valid_pos = is_pos & (sim >= self.island_thresh)\n\n        # fallback per-anchor if nothing valid\n        valid_cnt = valid_pos.sum(dim=1)                       # [B]\n        fallback = (valid_cnt == 0)                            # [B]\n        if fallback.any():\n            valid_pos = torch.where(fallback[:, None], is_pos, valid_pos)\n\n        if self.debug:\n            with torch.no_grad():\n                tot_pos = is_pos.sum().item()\n                kept_pos = valid_pos.sum().item()\n                frac_kept = kept_pos / max(1.0, tot_pos)\n                frac_fallback = fallback.float().mean().item()\n                print(f\"[IslandCircle] keep_pos_frac={frac_kept:.3f} | fallback_frac={frac_fallback:.3f} | thresh={self.island_thresh}\")\n\n        big = torch.finfo(sim.dtype).max\n\n        # hardest positives: smallest sim among valid positives\n        sim_p = torch.where(valid_pos, sim, sim.new_full(sim.shape, big))\n        sp_hard = torch.topk(sim_p, k=k_pos, dim=1, largest=False).values\n\n        # hardest negatives: largest sim among negatives\n        sim_n = torch.where(is_neg, sim, sim.new_full(sim.shape, -big))\n        sn_hard = torch.topk(sim_n, k=k_neg, dim=1, largest=True).values\n\n        ap = torch.clamp_min(-sp_hard.detach() + 1 + self.m, 0.)\n        an = torch.clamp_min(sn_hard.detach() + self.m, 0.)\n\n        logit_p = -self.gamma * ap * (sp_hard - (1 - self.m))\n        logit_n =  self.gamma * an * (sn_hard - self.m)\n\n        loss_p = torch.logsumexp(logit_p, dim=1)\n        loss_n = torch.logsumexp(logit_n, dim=1)\n\n        return F.softplus(loss_p + loss_n).mean()\n\n\ncriterion_metric = HardCircleLoss(m=0.25, gamma=16.0, topk_pos=4, topk_neg=60)\n#criterion_metric = CircleLoss(m=0.25, gamma=16.0)\nprint(\"✅ Using CircleLoss(m=0.25, gamma=16.0) [fixed self-pairs]\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:24:58.623721Z","iopub.execute_input":"2026-03-05T04:24:58.624566Z","iopub.status.idle":"2026-03-05T04:24:58.639861Z","shell.execute_reply.started":"2026-03-05T04:24:58.624535Z","shell.execute_reply":"2026-03-05T04:24:58.639125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 14 — Identity-balanced mAP (macro over identities)\n# ============================================================\n\ndef average_precision_from_sorted_rel(rel_sorted: np.ndarray):\n    npos = rel_sorted.sum()\n    if npos == 0:\n        return 0.0\n    cumsum = np.cumsum(rel_sorted)\n    precision_at_k = cumsum / (np.arange(len(rel_sorted)) + 1)\n    return float(precision_at_k[rel_sorted].sum() / npos)\n\ndef identity_balanced_map(emb: np.ndarray, labels: np.ndarray):\n    sims = emb @ emb.T\n    N = sims.shape[0]\n    APs = np.zeros(N, dtype=np.float32)\n\n    for i in range(N):\n        sims[i, i] = -1e9\n        order = np.argsort(-sims[i])\n        rel = (labels[order] == labels[i])\n        APs[i] = average_precision_from_sorted_rel(rel)\n\n    per_id = []\n    for y in np.unique(labels):\n        per_id.append(APs[labels == y].mean())\n    return float(np.mean(per_id))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:25:01.127729Z","iopub.execute_input":"2026-03-05T04:25:01.128556Z","iopub.status.idle":"2026-03-05T04:25:01.134802Z","shell.execute_reply.started":"2026-03-05T04:25:01.128522Z","shell.execute_reply":"2026-03-05T04:25:01.133939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 15 — Train/Eval helpers (Metric schedule, NO PCGrad)\n#   - Removes all PCGrad logic completely\n#   - Uses standard: total_loss = CE + mw * metric\n#   - Keeps metric_every_steps + warmup/ramp schedule\n# ============================================================\n\nimport time\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom tqdm.auto import tqdm\n\n# ✅ ensure PCGrad is OFF\n\n\ndef _metric_weight_for_epoch(epoch_idx: int):\n    warm = int(getattr(Config, \"metric_warmup_epochs\", 0))\n    ramp = int(getattr(Config, \"metric_ramp_epochs\", 0))\n    base = float(getattr(Config, \"metric_weight\", 1.0))\n\n    if epoch_idx < warm:\n        return 0.0\n    if ramp <= 0:\n        return base\n\n    t = (epoch_idx - warm + 1) / float(ramp)\n    t = max(0.0, min(1.0, t))\n    return base * t\n\ndef train_epoch(model, loader, optimizer, scaler, epoch_idx: int):\n    model.train()\n    optimizer.zero_grad(set_to_none=True)\n\n    mw = _metric_weight_for_epoch(epoch_idx)\n    metric_every = int(getattr(Config, \"metric_every_steps\", 1))\n    inv_acc = 1.0 / float(getattr(Config, \"grad_accum\", 1))\n\n    total_meter = 0.0\n    cls_meter   = 0.0\n    met_meter   = 0.0\n\n    for step, (imgs, labels) in enumerate(tqdm(loader, leave=False, desc=\"Training\")):\n        imgs = imgs.to(Config.device, non_blocking=True).contiguous()\n        labels = labels.to(Config.device, non_blocking=True)\n\n        # ----------------------------\n        # Forward + CE/ArcFace (AMP)\n        # ----------------------------\n        with torch.amp.autocast(device_type=Config.device_type):\n            feat_tri, feat_cls, logits = model(imgs, labels, return_emb=True)\n            loss_cls = criterion_cls(logits, labels)\n\n        do_metric = (mw > 0) and (metric_every <= 1 or (step % metric_every == 0))\n\n        # ----------------------------\n        # Metric loss (FP32)\n        # ----------------------------\n        with torch.amp.autocast(device_type=Config.device_type, enabled=False):\n            if do_metric:\n                kind = str(getattr(Config, \"metric_loss\", \"circle\")).lower()\n                if kind == \"circle\":\n                    # choose where Circle is applied\n                    if bool(getattr(Config, \"circle_on_feat_cls\", True)):\n                        z = feat_cls.float()\n                    else:\n                        z = feat_tri.float()\n                    loss_met = criterion_metric(z, labels)\n                else:\n                    loss_met = criterion_metric(feat_tri.float(), labels)\n\n                if not torch.isfinite(loss_met):\n                    loss_met = loss_cls.float().new_tensor(0.0)\n            else:\n                loss_met = loss_cls.float().new_tensor(0.0)\n\n            total_loss = loss_cls.float() + mw * loss_met\n\n        # ----------------------------\n        # Normal backward\n        # ----------------------------\n        scaler.scale(total_loss * inv_acc).backward()\n\n        if (step + 1) % int(getattr(Config, \"grad_accum\", 1)) == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n\n        total_meter += float(total_loss.item())\n        cls_meter   += float(loss_cls.item())\n        met_meter   += float(loss_met.item())\n\n    # flush remainder\n    if len(loader) % int(getattr(Config, \"grad_accum\", 1)) != 0:\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad(set_to_none=True)\n\n    denom = max(1, len(loader))\n    return (total_meter / denom), (cls_meter / denom), (met_meter / denom), mw\n\n\n@torch.no_grad()\ndef extract_embeddings(model, loader, use_dp_eval=True, show_batch0_time=True):\n    model.eval()\n    feats, names = [], []\n\n    base = model.module if isinstance(model, nn.DataParallel) else model\n\n    dp_embed = None\n    if use_dp_eval and isinstance(model, nn.DataParallel) and torch.cuda.device_count() >= 2:\n        class _EmbedWrap(nn.Module):\n            def __init__(self, m):\n                super().__init__()\n                self.m = m\n            def forward(self, x):\n                ft, fc = self.m.embed_pair(x)\n                return ft, fc\n        dp_embed = nn.DataParallel(_EmbedWrap(base), device_ids=list(range(torch.cuda.device_count())))\n\n    it = tqdm(loader, desc=\"Extract\")\n    for bi, batch in enumerate(it):\n        imgs, second = batch\n        t0 = time.time()\n\n        imgs = imgs.to(Config.device, non_blocking=True).contiguous()\n\n        with torch.amp.autocast(device_type=Config.device_type, enabled=(Config.device_type == \"cuda\")):\n            if dp_embed is not None:\n                ft, fc = dp_embed(imgs)\n            else:\n                ft, fc = base.embed_pair(imgs)\n\n            if getattr(Config, \"use_tta\", False):\n                imgs_f = torch.flip(imgs, dims=[3])\n                if dp_embed is not None:\n                    ft2, fc2 = dp_embed(imgs_f)\n                else:\n                    ft2, fc2 = base.embed_pair(imgs_f)\n                ft = 0.5 * (ft + ft2)\n                fc = 0.5 * (fc + fc2)\n\n        z = torch.cat([F.normalize(ft, dim=1), F.normalize(fc, dim=1)], dim=1)\n        z = F.normalize(z, dim=1).float().cpu().numpy()\n\n        feats.append(z)\n        try:\n            names.extend(list(second))\n        except Exception:\n            pass\n\n        if show_batch0_time and bi == 0:\n            if Config.device_type == \"cuda\":\n                torch.cuda.synchronize()\n            it.set_postfix({\"batch0_sec\": f\"{(time.time()-t0):.1f}\"})\n\n    return np.concatenate(feats, axis=0), names","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:25:01.543942Z","iopub.execute_input":"2026-03-05T04:25:01.544855Z","iopub.status.idle":"2026-03-05T04:25:01.562282Z","shell.execute_reply.started":"2026-03-05T04:25:01.544822Z","shell.execute_reply":"2026-03-05T04:25:01.561541Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# NEW CELL — Pre-load Validation Set into RAM (pinned, FP16)\n#   Run AFTER creating val_eval_loader but BEFORE the train loop.\n#   Result: extract_embeddings(model, RAM_VAL_BATCHES) is FAST.\n# ============================================================\n\nimport torch\nfrom tqdm.auto import tqdm\n\nprint(\"⏳ Pre-loading validation set into RAM...\")\n\nRAM_VAL_BATCHES = []\nfor imgs, second in tqdm(val_eval_loader, desc=\"Pre-loading\"):\n    # imgs is CPU tensor (already transformed). Convert to FP16 to save RAM.\n    # (Safe for eval embeddings; GPU autocast will handle it too.)\n    imgs = imgs.contiguous().half()\n\n    # pin memory for fast non_blocking GPU transfer\n    imgs = imgs.pin_memory()\n\n    # keep second as-is (labels or filenames)\n    RAM_VAL_BATCHES.append((imgs, second))\n\nprint(f\"✅ Pre-loaded {len(RAM_VAL_BATCHES)} val batches into RAM (FP16 + pinned).\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:25:03.808563Z","iopub.execute_input":"2026-03-05T04:25:03.809245Z","iopub.status.idle":"2026-03-05T04:26:13.638918Z","shell.execute_reply.started":"2026-03-05T04:25:03.809214Z","shell.execute_reply":"2026-03-05T04:26:13.638383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# REPLACE CELL — STEP-based LR schedule (warmup + cosine) ✅\n#   Why: makes LR behavior consistent even if steps_mult / steps_per_epoch change\n#   Works for 2-group or 3-group optimizers automatically\n#   Usage:\n#     - Run this cell after you create optimizer + train_loader + Config.num_epochs\n#     - In the train loop: call step_scheduler.step() AFTER each optimizer.step()\n#     - Remove/ignore the old epoch-based LambdaLR\n# ============================================================\n\nimport math\n\nimport math\n\nclass StepWarmupCosine:\n    def __init__(\n        self,\n        optimizer,\n        train_loader_len,\n        num_epochs,\n        base_lrs,\n        min_lrs,\n        warmup_frac=0.05,\n        warmup_start_factor=0.10,  # ✅ start warmup at 10% of base_lr\n    ):\n        self.opt = optimizer\n        self.total_steps = int(train_loader_len) * int(num_epochs)\n        self.total_steps = max(1, self.total_steps)\n\n        self.warmup_steps = int(self.total_steps * float(warmup_frac))\n        self.warmup_steps = max(0, self.warmup_steps)\n\n        self.base_lrs = [float(x) for x in base_lrs]\n        self.min_lrs  = [float(x) for x in min_lrs]\n        self.warmup_start_factor = float(warmup_start_factor)\n\n        assert len(self.base_lrs) == len(self.opt.param_groups), \"base_lrs must match param_groups\"\n        assert len(self.min_lrs)  == len(self.opt.param_groups), \"min_lrs must match param_groups\"\n\n        self.step_idx = 0\n        self._apply_lrs(0)\n\n    def _lr_at(self, s, base_lr, min_lr):\n        # warmup start LR is a fraction of base LR (not min_lr)\n        warm_start = base_lr * self.warmup_start_factor\n\n        if self.warmup_steps > 0 and s < self.warmup_steps:\n            t = s / float(self.warmup_steps)\n            return warm_start + (base_lr - warm_start) * t\n\n        # cosine decay base -> min\n        t = (s - self.warmup_steps) / float(max(1, self.total_steps - self.warmup_steps))\n        t = max(0.0, min(1.0, t))\n        cos = 0.5 * (1.0 + math.cos(math.pi * t))\n        return min_lr + (base_lr - min_lr) * cos\n\n    def _apply_lrs(self, s):\n        for i, g in enumerate(self.opt.param_groups):\n            g[\"lr\"] = self._lr_at(s, self.base_lrs[i], self.min_lrs[i])\n\n    def step(self):\n        self.step_idx += 1\n        self._apply_lrs(self.step_idx)\n\n    def get_last_lr(self):\n        return [g[\"lr\"] for g in self.opt.param_groups]\n# add Codeadd Markdown\n# ============================================================\n# REPLACE CELL — Optimizer (AdamW) with per-group WEIGHT DECAY + no_decay\n#   - bb / arc / aux each get their own LR + weight_decay\n#   - bias + norm params get weight_decay=0\n# ============================================================\n\nimport torch\nimport torch.nn as nn\n\n# --- recommended starting decays (tune)\nConfig.wd_backbone = float(getattr(Config, \"wd_backbone\", 0.01))\nConfig.wd_arc      = float(getattr(Config, \"wd_arc\",      0.005))\nConfig.wd_aux      = float(getattr(Config, \"wd_aux\",      0.02))\n\ndef _is_norm_or_bias(name: str, p: torch.nn.Parameter):\n    if name.endswith(\".bias\"):\n        return True\n    # catch common norm params by module name in the parameter path\n    lname = name.lower()\n    if (\"bn\" in lname) or (\"batchnorm\" in lname) or (\"layernorm\" in lname) or (\".ln\" in lname):\n        return True\n    # also catch timm norms naming\n    if \"norm\" in lname:\n        return True\n    return False\n\ndef build_adamw_param_groups(model):\n    m = model.module if isinstance(model, nn.DataParallel) else model\n\n    lr_bb   = float(getattr(Config, \"lr_backbone\", 2.6e-5))\n    lr_head = float(getattr(Config, \"lr_head\", 1.6e-4))\n    lr_arc  = float(getattr(Config, \"lr_arcface\", lr_head * 0.75))\n    lr_aux  = float(getattr(Config, \"lr_aux\",     lr_head * 1.50))\n\n    wd_bb   = float(getattr(Config, \"wd_backbone\", 0.01))\n    wd_arc  = float(getattr(Config, \"wd_arc\",      0.005))\n    wd_aux  = float(getattr(Config, \"wd_aux\",      0.02))\n\n    # buckets: (decay / no_decay) x (bb / arc / aux)\n    bb_decay,  bb_nodecay  = [], []\n    arc_decay, arc_nodecay = [], []\n    aux_decay, aux_nodecay = [], []\n\n    for name, p in m.named_parameters():\n        if not p.requires_grad:\n            continue\n\n        no_decay = _is_norm_or_bias(name, p)\n\n        # --- group routing (matches your v5 intent)\n        if name.startswith(\"backbone.\"):\n            (bb_nodecay if no_decay else bb_decay).append(p)\n\n        elif name.startswith((\"arc_neck.\", \"bnneck.\", \"head.\")):\n            (arc_nodecay if no_decay else arc_decay).append(p)\n\n        else:\n            (aux_nodecay if no_decay else aux_decay).append(p)\n\n    # de-dupe by id\n    def uniq(ps):\n        seen, out = set(), []\n        for p in ps:\n            if id(p) not in seen:\n                out.append(p); seen.add(id(p))\n        return out\n\n    groups = []\n    # backbone\n    groups += [{\"params\": uniq(bb_decay),    \"lr\": lr_bb,  \"weight_decay\": wd_bb}] if bb_decay else []\n    groups += [{\"params\": uniq(bb_nodecay),  \"lr\": lr_bb,  \"weight_decay\": 0.0}]  if bb_nodecay else []\n\n    # arc branch + head\n    groups += [{\"params\": uniq(arc_decay),   \"lr\": lr_arc, \"weight_decay\": wd_arc}] if arc_decay else []\n    groups += [{\"params\": uniq(arc_nodecay), \"lr\": lr_arc, \"weight_decay\": 0.0}]    if arc_nodecay else []\n\n    # aux\n    groups += [{\"params\": uniq(aux_decay),   \"lr\": lr_aux, \"weight_decay\": wd_aux}] if aux_decay else []\n    groups += [{\"params\": uniq(aux_nodecay), \"lr\": lr_aux, \"weight_decay\": 0.0}]    if aux_nodecay else []\n\n    print(f\"[AdamW groups] bb(dec)={len(bb_decay)} bb(nodec)={len(bb_nodecay)} | \"\n          f\"arc(dec)={len(arc_decay)} arc(nodec)={len(arc_nodecay)} | \"\n          f\"aux(dec)={len(aux_decay)} aux(nodec)={len(aux_nodecay)}\")\n    print(f\"[WD] bb={wd_bb} arc={wd_arc} aux={wd_aux}\")\n\n    return groups\n\n\n# ----------------------------\n# Build base/min LR lists from your Config\n# (Supports 2 or 3 groups)\n# ----------------------------\n\n\n# ----------------------------\n# HOW TO USE (copy these 2 lines into your train loop)\n# ----------------------------\n# After every optimizer.step() (i.e., when you actually step, not every microbatch):\n#   step_scheduler.step()\n#\n# To print current LRs:\n#   lrs = step_scheduler.get_last_lr()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:26:13.639888Z","iopub.execute_input":"2026-03-05T04:26:13.640121Z","iopub.status.idle":"2026-03-05T04:26:13.659275Z","shell.execute_reply.started":"2026-03-05T04:26:13.640098Z","shell.execute_reply":"2026-03-05T04:26:13.658375Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 16 — Train loop (DataParallel) + STEP Scheduler +\n#   - Assumes StepWarmupCosine class is defined in another cell\n#   - Assumes train_epoch(...) is defined in another cell\n#   - Assumes extract_embeddings(...) + identity_balanced_map(...) exist\n#   - Uses monkey-patch optimizer.step() -> step_scheduler.step()\n# ============================================================\n\nimport os, math, time\nfrom pathlib import Path\nimport numpy as np\nimport torch\nimport torch.nn as nn\n\nos.makedirs(Config.ckpt_dir, exist_ok=True)\n\n# ----------------------------\n# Model\n# ----------------------------\nbase_model = ReIDBoss(num_classes=31).to(Config.device)\nif torch.cuda.device_count() >= 2:\n    model = nn.DataParallel(base_model, device_ids=list(range(torch.cuda.device_count())))\nelse:\n    model = base_model\n\n# ----------------------------\n# Param groups (3 groups: bb / arc / aux)\n# ---------------------------\noptimizer = torch.optim.AdamW(build_adamw_param_groups(model))\nn_groups = len(optimizer.param_groups)\nprint(\"optimizer.param_groups:\", n_groups)\n# ----------------------------\n# STEP scheduler (warmup + cosine)\n# ----------------------------\n# defaults for min lrs for the extra groups\nif n_groups == 6:\n    base_lrs = [\n        float(Config.lr_backbone), float(Config.lr_backbone),\n        float(Config.lr_arcface),  float(Config.lr_arcface),\n        float(Config.lr_aux),      float(Config.lr_aux),\n    ]\n    min_lrs  = [\n        float(Config.min_lr_backbone), float(Config.min_lr_backbone),\n        float(Config.min_lr_arcface),  float(Config.min_lr_arcface),\n        float(Config.min_lr_aux),      float(Config.min_lr_aux),\n    ]\nelif n_groups == 3:\n    base_lrs = [float(Config.lr_backbone), float(Config.lr_arcface), float(Config.lr_aux)]\n    min_lrs  = [float(Config.min_lr_backbone), float(Config.min_lr_arcface), float(Config.min_lr_aux)]\nelse:\n    # fallback: keep whatever optimizer already has\n    base_lrs = [g[\"lr\"] for g in optimizer.param_groups]\n    # simple safe min: 15% of base (you can tune)\n    min_lrs  = [0.15 * lr for lr in base_lrs]\n\nstep_scheduler = StepWarmupCosine(\n    optimizer=optimizer,\n    train_loader_len=len(train_loader),\n    num_epochs=int(Config.num_epochs),\n    base_lrs=base_lrs,\n    min_lrs=min_lrs,\n    warmup_frac=float(getattr(Config, \"warmup_frac\", 0.05)),\n)\n\n# ----------------------------\n# AMP scaler\n# ----------------------------\nscaler = torch.amp.GradScaler(enabled=(Config.device_type == \"cuda\"))\n\n# ----------------------------\n# Monkey-patch optimizer.step to also step the LR schedule\n# (so you DON'T need to edit train_epoch in another cell)\n# ----------------------------\n_orig_step = optimizer.step\ndef _step_with_sched(*args, **kwargs):\n    out = _orig_step(*args, **kwargs)\n    step_scheduler.step()\n    return out\noptimizer.step = _step_with_sched\n\n\nprint(\"🔥 Training (DINOv3 + ArcFace + Metric, PK, DP) — step LR schedule ON\")\nprint(\"Initial LRs:\", step_scheduler.get_last_lr())\n\nbest_map = -1.0\n\nfor epoch in range(int(Config.num_epochs)):\n    if hasattr(train_loader.batch_sampler, \"set_epoch\"):\n        train_loader.batch_sampler.set_epoch(epoch)\n\n    # train_epoch is defined elsewhere\n    loss, loss_cls, loss_met, tw = train_epoch(model, train_loader, optimizer, scaler, epoch)\n\n    # ---- eval\n    if (epoch + 1) % int(getattr(Config, \"eval_every\", 1)) == 0:\n        val_emb, _ = extract_embeddings(model, RAM_VAL_BATCHES)\n        cv_map = identity_balanced_map(val_emb, val_labels)\n    else:\n        cv_map = float(\"nan\")\n\n    # ---- optional FFT A/B check every N epochs\n   \n    # ---- logging\n    lrs = step_scheduler.get_last_lr()\n    print(\n        f\"Epoch {epoch+1}/{Config.num_epochs} | \"\n        f\"Loss {loss:.4f} (cls {loss_cls:.4f}, met {loss_met:.4f}, tw {tw:.2f}) | \"\n        f\"CV id-mAP {cv_map:.4f} | \"\n        f\"LR(bb) {lrs[0]:.2e} | LR(arc) {lrs[1]:.2e} | LR(aux) {lrs[2]:.2e}\"\n    )\n\n    # ---- checkpoint\n    if np.isfinite(cv_map) and cv_map > best_map:\n        best_map = cv_map\n        path = Path(Config.ckpt_dir) / \"best_v3.pt\"\n        if isinstance(model, nn.DataParallel):\n            torch.save(model.module.state_dict(), path)\n        else:\n            torch.save(model.state_dict(), path)\n        print(\"✅ Saved:\", path)\n\n    torch.cuda.empty_cache()\n\nprint(\"Best CV id-mAP:\", best_map)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T04:26:13.660241Z","iopub.execute_input":"2026-03-05T04:26:13.660547Z","iopub.status.idle":"2026-03-05T07:42:47.464031Z","shell.execute_reply.started":"2026-03-05T04:26:13.660525Z","shell.execute_reply":"2026-03-05T07:42:47.463077Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"v5-dualhead-cls\nGrad checkpointing ON\n[AdamW groups] bb(dec)=147 bb(nodec)=171 | arc(dec)=2 arc(nodec)=1 | aux(dec)=6 aux(nodec)=3\n[WD] bb=0.01 arc=0.005 aux=0.02\noptimizer.param_groups: 6\n🔥 Training (DINOv3 + ArcFace + Metric, PK, DP) — step LR schedule ON\nInitial LRs: [2.6e-06, 2.6e-06, 8.000000000000001e-06, 8.000000000000001e-06, 2.9999999999999997e-05, 2.9999999999999997e-05]\nExtract: 100%\n 6/6 [00:26<00:00,  4.38s/it, batch0_sec=4.5]\nEpoch 1/15 | Loss 10.6093 (cls 10.6093, met 0.0000, tw 0.00) | CV id-mAP 0.4203 | LR(bb) 2.60e-05 | LR(arc) 2.60e-05 | LR(aux) 8.00e-05\n✅ Saved: /kaggle/working/ckpts/best.pt\nExtract: 100%\n 6/6 [00:26<00:00,  4.36s/it, batch0_sec=4.5]\nEpoch 2/15 | Loss 2.4288 (cls 2.4288, met 0.0000, tw 0.00) | CV id-mAP 0.6033 | LR(bb) 2.56e-05 | LR(arc) 2.56e-05 | LR(aux) 7.88e-05\n✅ Saved: /kaggle/working/ckpts/best.pt\nExtract: 100%\n 6/6 [00:26<00:00,  4.38s/it, batch0_sec=4.5]\nEpoch 3/15 | Loss 0.3762 (cls 0.3256, met 2.0241, tw 0.03) | CV id-mAP 0.7346 | LR(bb) 2.47e-05 | LR(arc) 2.47e-05 | LR(aux) 7.62e-05\n✅ Saved: /kaggle/working/ckpts/best.pt\nExtract: 100%\n 6/6 [00:26<00:00,  4.34s/it, batch0_sec=4.5]\nEpoch 4/15 | Loss 0.1596 (cls 0.0712, met 1.7685, tw 0.05) | CV id-mAP 0.7279 | LR(bb) 2.33e-05 | LR(arc) 2.33e-05 | LR(aux) 7.24e-05\nExtract: 100%\n 6/6 [00:26<00:00,  4.37s/it, batch0_sec=4.5]\nEpoch 5/15 | Loss 0.1078 (cls 0.0217, met 1.7218, tw 0.05) | CV id-mAP 0.7358 | LR(bb) 2.15e-05 | LR(arc) 2.15e-05 | LR(aux) 6.73e-05\n✅ Saved: /kaggle/working/ckpts/best.pt\nExtract: 100%\n 6/6 [00:26<00:00,  4.38s/it, batch0_sec=4.5]\nEpoch 6/15 | Loss 0.0981 (cls 0.0127, met 1.7063, tw 0.05) | CV id-mAP 0.7414 | LR(bb) 1.94e-05 | LR(arc) 1.94e-05 | LR(aux) 6.14e-05\n✅ Saved: /kaggle/working/ckpts/best.pt\nExtract: 100%\n 6/6 [00:26<00:00,  4.36s/it, batch0_sec=4.6]\nEpoch 7/15 | Loss 0.0938 (cls 0.0091, met 1.6952, tw 0.05) | CV id-mAP 0.7411 | LR(bb) 1.71e-05 | LR(arc) 1.71e-05 | LR(aux) 5.49e-05\nExtract: 100%\n 6/6 [00:26<00:00,  4.38s/it, batch0_sec=4.5]\nEpoch 8/15 | Loss 0.0921 (cls 0.0075, met 1.6909, tw 0.05) | CV id-mAP 0.7417 | LR(bb) 1.47e-05 | LR(arc) 1.47e-05 | LR(aux) 4.81e-05\n✅ Saved: /kaggle/working/ckpts/best.pt\nExtract: 100%\n 6/6 [00:26<00:00,  4.36s/it, batch0_sec=4.5]\nEpoch 9/15 | Loss 0.0906 (cls 0.0063, met 1.6867, tw 0.05) | CV id-mAP 0.7435 | LR(bb) 1.23e-05 | LR(arc) 1.23e-05 | LR(aux) 4.14e-05\n✅ Saved: /kaggle/working/ckpts/best.pt","metadata":{}},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL — LONG-TAIL PSEUDO-LABELING ENGINE (FIXED + SAFE)\n#   Goal: add \"more views\" for rare IDs by mining TEST images.\n#   Fixes:\n#     - uses UNIQUE test filenames (NOT pair rows)\n#     - adds safety gates: sim + margin + vote + arcface_conf agreement\n#     - caps per class\n#     - NO crop (uses your existing test_transform + nocrop caches)\n# ============================================================\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader\nfrom tqdm.auto import tqdm\nfrom pathlib import Path\n\nTRAIN_DIR = Path(\"/kaggle/input/jaguar-re-id/train/train\")\nTEST_DIR  = Path(\"/kaggle/input/jaguar-re-id/test/test\")\n# -------------------------\n# USER KNOBS (safe defaults)\n# -------------------------\nMAX_SAMPLES_RARE = 25      # classes with < this many TRAIN images are \"rare\"\nTOPK_VOTE        = 25      # vote neighbors from gallery\nSIM_THR          = 0.78    # start 0.75-0.82 (use other gates for safety)\nMARGIN_THR       = 0.06    # sim(top1 class) - sim(top2 class)\nVOTE_THR         = 0.60    # fraction of topk neighbors that match top1 class\nCONF_THR         = 0.60    # arcface softmax max prob\nCAP_PER_RARE_ID  = 20      # max pseudo images per rare ID\nUSE_DBA_GALLERY  = True\nDBA_TOPK         = 3\nDBA_ALPHA        = 0.5\n\n# -------------------------\nprint(\"🎯 INITIATING LONG-TAIL PSEUDO-LABELING (SAFE)...\")\n\n# 0) Load best weights\nBEST = Path(Config.ckpt_dir) / \"best_v3.pt\"\nbase = model.module if isinstance(model, torch.nn.DataParallel) else model\nbase.load_state_dict(torch.load(BEST, map_location=\"cpu\"), strict=False)\nbase.eval()\nprint(\"✅ Loaded:\", BEST)\n\n# 1) Build UNIQUE test filenames correctly\nassert \"test_df\" in globals(), \"Need test_df from competition.\"\nassert (\"query_image\" in test_df.columns) and (\"gallery_image\" in test_df.columns), \"test_df must have query_image/gallery_image.\"\n\nunique_test = sorted(set(test_df[\"query_image\"]) | set(test_df[\"gallery_image\"]))\ntest_df_unique = pd.DataFrame({\"filename\": unique_test})\nprint(\"unique test images:\", len(unique_test))\n\n# 2) Datasets / loaders (NO CROP: uses your test_transform + nocrop caches)\ntrain_ds_eval = JaguarDataset(\n    tr_df, TRAIN_DIR, transform=test_transform, is_test=False, label_map=label_map,\n    cache_dir=(TRAIN_CACHE if getattr(Config, \"cache_crops\", True) else None)\n)\ntest_ds_eval = JaguarDataset(\n    test_df_unique, TEST_DIR, transform=test_transform, is_test=True,\n    cache_dir=(TEST_CACHE if getattr(Config, \"cache_crops\", True) else None)\n)\n\n# If workers ever crash, set num_workers=0 temporarily.\ntrain_loader_eval = DataLoader(train_ds_eval, batch_size=96, shuffle=False, num_workers=4, pin_memory=True)\ntest_loader_eval  = DataLoader(test_ds_eval,  batch_size=96, shuffle=False, num_workers=4, pin_memory=True)\n\ntrain_labels = train_ds_eval.df[\"label\"].values.astype(int)\ntest_filenames = test_ds_eval.df[\"filename\"].values.astype(object)\n\n# 3) OOM-safe extractor: embeddings + ArcFace confidence (if return_bn exists)\n@torch.no_grad()\ndef extract_Z_and_conf(model, loader, chunk_size=16):\n    model.eval()\n    base = model.module if isinstance(model, torch.nn.DataParallel) else model\n\n    Z_list, P_list, C_list = [], [], []\n\n    has_return_bn = (\"return_bn\" in base.embed_pair.__code__.co_varnames)\n\n    for imgs, second in tqdm(loader, desc=\"Extract(Z+conf)\"):\n        micro = torch.split(imgs, chunk_size) if imgs.size(0) > chunk_size else (imgs,)\n        for mb in micro:\n            mb = mb.to(Config.device, non_blocking=True).contiguous()\n\n            with torch.amp.autocast(device_type=Config.device_type, enabled=(Config.device_type == \"cuda\")):\n                if has_return_bn:\n                    ft, fc, arc_bn = base.embed_pair(mb, return_bn=True)\n                    logits = base.head(arc_bn.float(), label=None)  # no-margin logits\n                else:\n                    ft, fc = base.embed_pair(mb)\n                    logits = base.head(fc.float(), label=None) if hasattr(base, \"head\") else None\n\n                z = torch.cat([F.normalize(ft.float(), dim=1), F.normalize(fc.float(), dim=1)], dim=1)\n                z = F.normalize(z, dim=1)\n\n            Z_list.append(z.float().cpu().numpy())\n\n            if logits is not None:\n                probs = F.softmax(logits.float(), dim=1)\n                conf = probs.max(dim=1).values\n                pred = probs.argmax(dim=1)\n                P_list.append(pred.cpu().numpy().astype(np.int32))\n                C_list.append(conf.cpu().numpy().astype(np.float32))\n            else:\n                P_list.append(np.full((mb.size(0),), -1, np.int32))\n                C_list.append(np.zeros((mb.size(0),), np.float32))\n\n    Z = np.concatenate(Z_list, axis=0).astype(np.float32)\n    P = np.concatenate(P_list, axis=0).astype(np.int32)\n    C = np.concatenate(C_list, axis=0).astype(np.float32)\n    return Z, P, C\n\nprint(\"\\n⏳ Extracting TRAIN gallery...\")\nZ_train, P_train, C_train = extract_Z_and_conf(model, train_loader_eval, chunk_size=16)\n\nprint(\"\\n⏳ Extracting TEST images...\")\nZ_test, P_test, C_test = extract_Z_and_conf(model, test_loader_eval, chunk_size=16)\n\n# 4) Identify rare classes\nclass_counts = pd.Series(train_labels).value_counts()\nrare_classes = set(class_counts[class_counts < MAX_SAMPLES_RARE].index.astype(int).tolist())\nprint(f\"\\n📊 Rare IDs: {len(rare_classes)}/{len(class_counts)} with < {MAX_SAMPLES_RARE} train images.\")\nprint(\"Example rare IDs:\", sorted(list(rare_classes))[:12])\n\n# 5) DBA on gallery (optional)\ndef dba_gallery(G, topk=3, alpha=0.5):\n    G = G.astype(np.float32)\n    G /= (np.linalg.norm(G, axis=1, keepdims=True) + 1e-12)\n    Sg = G @ G.T\n    np.fill_diagonal(Sg, -1e9)\n    nn = np.argsort(-Sg, axis=1)[:, :int(topk)]\n    G2 = G.copy()\n    for i in range(len(G)):\n        G2[i] = G[i] + float(alpha) * G[nn[i]].mean(axis=0)\n    G2 /= (np.linalg.norm(G2, axis=1, keepdims=True) + 1e-12)\n    return G2\n\nZg = Z_train.copy()\nif USE_DBA_GALLERY:\n    print(\"✨ Applying DBA to TRAIN gallery...\")\n    Zg = dba_gallery(Zg, topk=DBA_TOPK, alpha=DBA_ALPHA)\n\n# 6) Mine pseudo labels with strong gates\nprint(f\"\\n🧮 Mining TEST for pseudo-labels...\")\nprint(f\"   gates: sim>={SIM_THR} | margin>={MARGIN_THR} | vote>={VOTE_THR} | conf>={CONF_THR} | rare-only\")\n\n# similarity test->gallery\nS = Z_test @ Zg.T  # [Nt, Ng]\nnn_idx = np.argmax(S, axis=1)\nsim_nn = S[np.arange(len(Z_test)), nn_idx].astype(np.float32)\ny_nn = train_labels[nn_idx].astype(int)\n\n# topk vote fraction\nk = int(min(TOPK_VOTE, Zg.shape[0]))\ntopk_idx = np.argpartition(-S, kth=k-1, axis=1)[:, :k]\ntopk_lbl = train_labels[topk_idx]\nvote = (topk_lbl == y_nn[:, None]).mean(axis=1).astype(np.float32)\n\n# margin vs 2nd-best class (class-max trick)\nC = int(train_labels.max() + 1)\nper_class_max = np.full((len(Z_test), C), -1e9, dtype=np.float32)\nfor c in range(C):\n    m = (train_labels == c)\n    if m.any():\n        per_class_max[:, c] = S[:, m].max(axis=1)\n\ntop1 = per_class_max[np.arange(len(Z_test)), y_nn]\ntmp = per_class_max.copy()\ntmp[np.arange(len(Z_test)), y_nn] = -1e9\ntop2 = tmp.max(axis=1)\nmargin = (top1 - top2).astype(np.float32)\n\n# arcface agreement gate (predicted class from logits must match embedding NN class)\nagree = (P_test == y_nn)\n\nkeep = (\n    np.isin(y_nn, np.array(list(rare_classes), dtype=int)) &\n    agree &\n    (C_test >= CONF_THR) &\n    (sim_nn >= SIM_THR) &\n    (margin >= MARGIN_THR) &\n    (vote >= VOTE_THR)\n)\n\ndf_keep = pd.DataFrame({\n    \"filename\": test_filenames,\n    \"label\": y_nn,\n    \"pseudo_sim\": sim_nn,\n    \"pseudo_conf\": C_test,\n    \"pseudo_margin\": margin,\n    \"pseudo_vote\": vote,\n    \"agree\": agree.astype(np.int8),\n    \"keep\": keep.astype(np.int8),\n})\n\ndf_keep[\"score\"] = (0.55*df_keep[\"pseudo_sim\"] + 0.45*df_keep[\"pseudo_conf\"]) * df_keep[\"pseudo_vote\"]\n\nprint(\"Kept:\", int(df_keep[\"keep\"].sum()), \"/\", len(df_keep), \"coverage=\", float(df_keep[\"keep\"].mean()))\nif int(df_keep[\"keep\"].sum()) == 0:\n    print(\"⚠️ No pseudo labels passed. Lower SIM_THR to ~0.75 OR CONF_THR to 0.55 (keep margin/vote).\")\n\n# cap per rare ID\npicked = []\nfor lab, sub in df_keep[df_keep[\"keep\"].eq(1)].groupby(\"label\"):\n    sub = sub.sort_values(\"score\", ascending=False).head(CAP_PER_RARE_ID)\n    picked.append(sub)\n\npseudo_df = pd.concat(picked, axis=0).reset_index(drop=True) if len(picked) else df_keep.head(0)\n\n# add ground_truth string for your pipeline\ninv_label_map = {v: k for k, v in label_map.items()}\npseudo_df[\"ground_truth\"] = pseudo_df[\"label\"].map(inv_label_map).astype(str)\npseudo_df[\"is_pseudo\"] = True\n\nprint(\"\\n✅ Pseudo-labels added per ID (top):\")\ndisplay(pseudo_df[\"label\"].value_counts().head(15))\n\n# 7) Merge into training df (ONLY tr_df, never va_df)\ntr_df2 = tr_df.copy()\ntr_df2[\"is_pseudo\"] = False\n\n# keep only needed cols, but preserve your schema\ncols_base = list(tr_df2.columns)\nfor c in [\"pseudo_sim\",\"pseudo_conf\",\"pseudo_margin\",\"pseudo_vote\",\"score\"]:\n    if c not in cols_base:\n        tr_df2[c] = np.nan\n\npseudo_add = pseudo_df.copy()\n# match columns\nfor c in tr_df2.columns:\n    if c not in pseudo_add.columns:\n        pseudo_add[c] = np.nan\n\ntr_df_pseudo = pd.concat([tr_df2, pseudo_add[tr_df2.columns]], ignore_index=True)\n\nprint(f\"\\nOld train size: {len(tr_df)} | New train size: {len(tr_df_pseudo)} | Added: {len(pseudo_df)}\")\npseudo_df.to_csv(\"/kaggle/working/pseudo_labels_longtail.csv\", index=False)\ntr_df_pseudo.to_csv(\"/kaggle/working/tr_df_pseudo.csv\", index=False)\nprint(\"✅ Saved:\")\nprint(\"  /kaggle/working/pseudo_labels_longtail.csv\")\nprint(\"  /kaggle/working/tr_df_pseudo.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:44:59.866024Z","iopub.execute_input":"2026-03-05T07:44:59.866794Z","iopub.status.idle":"2026-03-05T07:46:52.024152Z","shell.execute_reply.started":"2026-03-05T07:44:59.866761Z","shell.execute_reply":"2026-03-05T07:46:52.023172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL — Remove / flag within-class outliers (label-noise guard)\n#   Finds samples whose best positive similarity is extremely low.\n#   These are exactly like idx316: likely mislabeled / extreme mismatch.\n# ============================================================\n\nimport numpy as np\nimport pandas as pd\n\n# Requires: tr_df, train_ds or a train-eval loader, and a working extractor that returns embeddings per row order\n# We'll do it on TRAIN (not VAL) using test_transform for stability.\n\nTRAIN_EVAL_DF = tr_df.copy().reset_index(drop=True)\n\ntrain_eval_ds = JaguarDataset(\n    TRAIN_EVAL_DF, TRAIN_DIR,\n    transform=test_transform, is_test=False,\n    label_map=label_map,\n    cache_dir=(TRAIN_CACHE if getattr(Config, \"cache_crops\", True) else None)\n)\n\ntrain_eval_loader = DataLoader(train_eval_ds, batch_size=64, shuffle=False, num_workers=0, pin_memory=True)\n\nZ_train, _ = extract_embeddings(model, train_eval_loader)   # your extract_embeddings returns L2-normalized concat\ny = TRAIN_EVAL_DF[\"label\"].values.astype(int)\n\nS = Z_train @ Z_train.T\nnp.fill_diagonal(S, -1e9)\n\nbest_pos = np.full(len(y), -1.0, dtype=np.float32)\nfor i in range(len(y)):\n    same = (y == y[i])\n    same[i] = False\n    if same.any():\n        best_pos[i] = float(S[i][same].max())\n\nTRAIN_EVAL_DF[\"best_pos_sim_train\"] = best_pos\n\n# Threshold: start conservative; tune (0.15~0.30)\nTHR = 0.20\noutliers = TRAIN_EVAL_DF[TRAIN_EVAL_DF[\"best_pos_sim_train\"] < THR].copy()\nprint(\"Outliers found:\", len(outliers), \" / \", len(TRAIN_EVAL_DF))\ndisplay(outliers[[\"filename\",\"label\",\"best_pos_sim_train\"]].sort_values(\"best_pos_sim_train\").head(25))\n\n# Option A: drop them\ntr_df_clean = TRAIN_EVAL_DF[TRAIN_EVAL_DF[\"best_pos_sim_train\"] >= THR].drop(columns=[\"best_pos_sim_train\"]).reset_index(drop=True)\nprint(\"Old tr_df:\", len(tr_df), \" -> Clean tr_df:\", len(tr_df_clean))\n\n# Option B: keep them but mark (for downweighted loss)\n# TRAIN_EVAL_DF[\"is_outlier\"] = (TRAIN_EVAL_DF[\"best_pos_sim_train\"] < THR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:04:41.566317Z","iopub.execute_input":"2026-03-05T08:04:41.567267Z","iopub.status.idle":"2026-03-05T08:05:47.933278Z","shell.execute_reply.started":"2026-03-05T08:04:41.567230Z","shell.execute_reply":"2026-03-05T08:05:47.929840Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL — THE ULTIMATE DEBUG PACK (MINIMAL & SAFE)\n#   Fixes:\n#     - OOM safe extractions (chunking)\n#     - Scaler-safe grad probe (manual unscaling avoids RuntimeError)\n#     - Checkpoint-safe grad probe (backward inside autocast)\n# ============================================================\n\nimport gc, math\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn.functional as F\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom torch.utils.data import DataLoader\n\n# ----------------------------\n# FLAGS (Turn sections on/off to save time)\n# ----------------------------\nDO_BRANCH_CHECK      = True\nDO_ALPHA_SWEEP       = True\nDO_QE_VALVAL         = True\nDO_ORPHANS_MARGIN    = True\nDO_DELTA_PER_ID      = True\nDO_REALISTIC_QG      = True   # query=VAL, gallery=TRAIN  (recommended)\nDO_QE_DBA_QG         = True   # on top of DO_REALISTIC_QG\nDO_GRAD_PROBE        = True   # requires criterion_metric\nPRINT_TOPK_TABLES    = True\n\n# ----------------------------\n# Safety / cleanup\n# ----------------------------\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n\n# ----------------------------\n# 0) Load best checkpoint\n# ----------------------------\nBEST = Path(getattr(Config, \"ckpt_dir\", \"/kaggle/working/ckpts\")) / \"best_v3.pt\"\nbase = model.module if isinstance(model, torch.nn.DataParallel) else model\nif BEST.exists():\n    sd = torch.load(BEST, map_location=\"cpu\")\n    base.load_state_dict(sd, strict=False)\n    print(\"✅ Loaded:\", BEST)\nelse:\n    print(\"⚠️ Not found:\", BEST, \"using current weights\")\n\n# ----------------------------\n# Helpers\n# ----------------------------\ndef l2norm(X):\n    X = X.astype(np.float32)\n    return X / (np.linalg.norm(X, axis=1, keepdims=True) + 1e-12)\n\ndef macro_id_map(emb, labels):\n    emb = l2norm(emb)\n    return identity_balanced_map(emb.astype(np.float32), labels.astype(int))\n\n@torch.no_grad()\ndef extract_ft_fc_safe(model, loader_or_batches, chunk_size=16, use_tta=None):\n    model.eval()\n    base = model.module if isinstance(model, torch.nn.DataParallel) else model\n    if use_tta is None:\n        use_tta = bool(getattr(Config, \"use_tta\", False))\n\n    ft_list, fc_list = [], []\n    for imgs, _ in tqdm(loader_or_batches, desc=\"Extract(ft/fc)\"):\n        micro = torch.split(imgs, chunk_size) if imgs.size(0) > chunk_size else (imgs,)\n        for mb in micro:\n            mb = mb.to(Config.device, non_blocking=True).contiguous()\n            with torch.amp.autocast(device_type=Config.device_type, enabled=(Config.device_type == \"cuda\")):\n                ft, fc = base.embed_pair(mb)\n                if use_tta:\n                    mb_f = torch.flip(mb, dims=[3])\n                    ft2, fc2 = base.embed_pair(mb_f)\n                    ft = 0.5 * (ft + ft2)\n                    fc = 0.5 * (fc + fc2)\n            ft_list.append(F.normalize(ft.float(), dim=1).cpu().numpy())\n            fc_list.append(F.normalize(fc.float(), dim=1).cpu().numpy())\n\n    FT = np.concatenate(ft_list, axis=0).astype(np.float32)\n    FC = np.concatenate(fc_list, axis=0).astype(np.float32)\n    return FT, FC\n\ndef qe_valval(emb, topk=3, alpha=0.5):\n    emb = l2norm(emb)\n    S = emb @ emb.T\n    np.fill_diagonal(S, -1e9)\n    nn = np.argsort(-S, axis=1)[:, :topk]\n    out = emb.copy()\n    for i in range(len(emb)):\n        out[i] = emb[i] + float(alpha) * emb[nn[i]].mean(axis=0)\n    return l2norm(out)\n\ndef find_orphans_and_margins(emb, labels, orphan_thr=0.15, topn=12):\n    emb = l2norm(emb)\n    labels = labels.astype(int)\n    S = emb @ emb.T\n    np.fill_diagonal(S, -1e9)\n\n    same = labels[:, None] == labels[None, :]\n    np.fill_diagonal(same, False)\n\n    max_pos = np.full(len(labels), -1.0, dtype=np.float32)\n    best_neg = np.full(len(labels), -1.0, dtype=np.float32)\n    for i in range(len(labels)):\n        if same[i].any():\n            max_pos[i] = float(S[i][same[i]].max())\n        best_neg[i] = float(S[i][~same[i]].max())\n\n    margin = max_pos - best_neg\n    orphans = np.where(max_pos < float(orphan_thr))[0]\n\n    df = pd.DataFrame({\n        \"idx\": np.arange(len(labels)),\n        \"label\": labels,\n        \"best_pos\": max_pos,\n        \"best_neg\": best_neg,\n        \"margin\": margin,\n    }).sort_values([\"margin\",\"best_pos\"], ascending=[True, True]).head(topn)\n\n    return orphans, max_pos, df\n\ndef per_id_delta(ft, fc, labels):\n    labels = labels.astype(int)\n    Zft = l2norm(ft)\n    Zcat = l2norm(np.concatenate([ft, fc], axis=1))\n    Sft = Zft @ Zft.T\n    Sca = Zcat @ Zcat.T\n    np.fill_diagonal(Sft, -1e9)\n    np.fill_diagonal(Sca, -1e9)\n\n    def ap_mat(S):\n        def ap_one(i):\n            order = np.argsort(-S[i])\n            rel = (labels[order] == labels[i])\n            npos = int(rel.sum())\n            if npos == 0: return 0.0\n            c = np.cumsum(rel)\n            p = c / (np.arange(len(rel)) + 1)\n            return float(p[rel].sum() / npos)\n        return np.array([ap_one(i) for i in range(len(labels))], dtype=np.float32)\n\n    ap_ft  = ap_mat(Sft)\n    ap_cat = ap_mat(Sca)\n\n    out = []\n    for y in np.unique(labels):\n        m = (labels == y)\n        out.append((int(y), int(m.sum()), float(ap_ft[m].mean()), float(ap_cat[m].mean())))\n    df = pd.DataFrame(out, columns=[\"label\",\"n\",\"map_ft\",\"map_cat\"])\n    df[\"cat_minus_ft\"] = df[\"map_cat\"] - df[\"map_ft\"]\n    return df.sort_values(\"cat_minus_ft\", ascending=False), df.sort_values(\"cat_minus_ft\", ascending=True)\n\n# --- Realistic query-gallery eval ---\ndef ap_from_sorted_rel(rel_sorted):\n    npos = int(rel_sorted.sum())\n    if npos == 0: return 0.0\n    c = np.cumsum(rel_sorted)\n    p = c / (np.arange(len(rel_sorted)) + 1)\n    return float(p[rel_sorted].sum() / npos)\n\ndef id_map_query_gallery(Q, q_labels, G, g_labels):\n    Q = l2norm(Q); G = l2norm(G)\n    S = Q @ G.T\n    AP = np.zeros(len(Q), dtype=np.float32)\n    for i in range(len(Q)):\n        order = np.argsort(-S[i])\n        rel = (g_labels[order] == q_labels[i])\n        AP[i] = ap_from_sorted_rel(rel)\n    per_id = []\n    for y in np.unique(q_labels):\n        per_id.append(AP[q_labels == y].mean())\n    return float(np.mean(per_id))\n\ndef qe_query_with_gallery(Q, G, topk=3, alpha=0.5):\n    Q = l2norm(Q); G = l2norm(G)\n    S = Q @ G.T\n    nn = np.argsort(-S, axis=1)[:, :topk]\n    Q2 = Q.copy()\n    for i in range(len(Q)):\n        Q2[i] = Q[i] + float(alpha) * G[nn[i]].mean(axis=0)\n    return l2norm(Q2)\n\ndef dba_gallery(G, topk=3, alpha=0.3):\n    G = l2norm(G)\n    S = G @ G.T\n    np.fill_diagonal(S, -1e9)\n    nn = np.argsort(-S, axis=1)[:, :topk]\n    G2 = G.copy()\n    for i in range(len(G)):\n        G2[i] = G[i] + float(alpha) * G[nn[i]].mean(axis=0)\n    return l2norm(G2)\n\n@torch.no_grad()\ndef extract_ft_fc_from_loader(model, loader, chunk_size=32):\n    model.eval()\n    base = model.module if isinstance(model, torch.nn.DataParallel) else model\n    ft_list, fc_list = [], []\n    for imgs, _ in tqdm(loader, desc=\"Extract(ft/fc)\"):\n        micro = torch.split(imgs, chunk_size) if imgs.size(0) > chunk_size else (imgs,)\n        for mb in micro:\n            mb = mb.to(Config.device, non_blocking=True).contiguous()\n            with torch.amp.autocast(device_type=Config.device_type, enabled=(Config.device_type == \"cuda\")):\n                ft, fc = base.embed_pair(mb)\n            ft_list.append(F.normalize(ft.float(), dim=1).cpu().numpy())\n            fc_list.append(F.normalize(fc.float(), dim=1).cpu().numpy())\n    return np.concatenate(ft_list, 0), np.concatenate(fc_list, 0)\n\n# ----------------------------\n# 1) Extract VAL FT/FC (once)\n# ----------------------------\nBATCH_SRC = RAM_VAL_BATCHES if \"RAM_VAL_BATCHES\" in globals() else val_eval_loader\nFT, FC = extract_ft_fc_safe(model, BATCH_SRC, chunk_size=16)\n\nlabels_val = val_labels.astype(int)\n\n# ----------------------------\n# 2) Branch check\n# ----------------------------\nif DO_BRANCH_CHECK:\n    m_ft  = macro_id_map(FT, labels_val)\n    m_fc  = macro_id_map(FC, labels_val)\n    m_cat = macro_id_map(np.concatenate([FT, FC], axis=1), labels_val)\n    print(f\"\\n✅ VAL-vs-VAL id-mAP | ft={m_ft:.4f} | fc={m_fc:.4f} | cat={m_cat:.4f}\")\n\n# ----------------------------\n# 3) Alpha sweep (FT + alpha*FC)\n# ----------------------------\nif DO_ALPHA_SWEEP:\n    alphas = [0.0, 0.05, 0.1, 0.2, 0.3, 0.5, 0.75, 1.0]\n    best = (-1, -1.0)\n    print(\"\\nAlpha sweep:\")\n    for a in alphas:\n        z = np.concatenate([FT, float(a)*FC], axis=1)\n        m = macro_id_map(z, labels_val)\n        print(f\"  alpha={a:>4} -> {m:.4f}\")\n        if m > best[1]:\n            best = (a, m)\n    print(\"✅ Best alpha:\", best[0], \"mAP:\", best[1])\n\n# ----------------------------\n# 4) QE on VAL-vs-VAL (fast, just to see trend)\n# ----------------------------\nif DO_QE_VALVAL:\n    Zcat = np.concatenate([FT, FC], axis=1)\n    base_m = macro_id_map(Zcat, labels_val)\n    print(\"\\nQE (VAL-vs-VAL) on CAT:\")\n    print(\"  baseline:\", base_m)\n    for k in [1,2,3,5]:\n        for a in [0.5, 1.0]:\n            m = macro_id_map(qe_valval(Zcat, topk=k, alpha=a), labels_val)\n            print(f\"  QE topk={k} alpha={a} -> {m:.4f}\")\n\n# ----------------------------\n# 5) Orphans + negative-margin samples\n# ----------------------------\nif DO_ORPHANS_MARGIN:\n    if \"VAL_META\" not in globals():\n        VAL_META = val_ds.df.copy().reset_index(drop=True)\n        VAL_META[\"idx\"] = np.arange(len(VAL_META))\n        if \"filename\" not in VAL_META.columns:\n            for c in [\"file\",\"image\",\"img\",\"path\",\"name\"]:\n                if c in VAL_META.columns:\n                    VAL_META[\"filename\"] = VAL_META[c].astype(str)\n                    break\n\n    Zcat = np.concatenate([FT, FC], axis=1)\n    orph, max_pos, worst_margin = find_orphans_and_margins(Zcat, labels_val, orphan_thr=0.15, topn=12)\n    VAL_META[\"max_pos_sim\"] = max_pos\n\n    print(\"\\nOrphans (max_pos < 0.15):\", len(orph), \"/\", len(labels_val))\n    if len(orph):\n        display(VAL_META.iloc[orph][[\"idx\",\"filename\",\"label\",\"max_pos_sim\"]].sort_values(\"max_pos_sim\").head(10))\n\n    print(\"\\nWorst negative-margin samples:\")\n    if PRINT_TOPK_TABLES:\n        dm = worst_margin.merge(VAL_META[[\"idx\",\"filename\"]], on=\"idx\", how=\"left\")\n        display(dm)\n\n# ----------------------------\n# 6) Per-ID delta: concat helps/hurts\n# ----------------------------\nif DO_DELTA_PER_ID:\n    print(\"\\nPer-ID delta (cat - ft):\")\n    top_gain, top_hurt = per_id_delta(FT, FC, labels_val)\n    print(\"Top gains:\")\n    if PRINT_TOPK_TABLES: display(top_gain.head(10))\n    print(\"Top hurts:\")\n    if PRINT_TOPK_TABLES: display(top_hurt.head(10))\n\n# ----------------------------\n# 7) Realistic evaluation (VAL query vs TRAIN gallery) + QE/DBA\n# ----------------------------\nif DO_REALISTIC_QG:\n    assert \"JaguarDataset\" in globals()\n    assert \"tr_df\" in globals() and \"va_df\" in globals()\n    assert \"label_map\" in globals()\n    assert \"TRAIN_DIR\" in globals()\n    assert \"test_transform\" in globals()\n    assert \"TRAIN_CACHE\" in globals()\n\n    train_eval_ds = JaguarDataset(\n        tr_df, TRAIN_DIR, transform=test_transform, is_test=False,\n        label_map=label_map, cache_dir=(TRAIN_CACHE if Config.cache_crops else None)\n    )\n    val_eval_ds2 = JaguarDataset(\n        va_df, TRAIN_DIR, transform=test_transform, is_test=False,\n        label_map=label_map, cache_dir=(TRAIN_CACHE if Config.cache_crops else None)\n    )\n\n    train_gallery_loader = DataLoader(train_eval_ds, batch_size=96, shuffle=False, num_workers=0, pin_memory=True)\n    val_query_loader     = DataLoader(val_eval_ds2, batch_size=96, shuffle=False, num_workers=0, pin_memory=True)\n\n    g_labels = train_eval_ds.df[\"label\"].values.astype(int)\n    q_labels = val_eval_ds2.df[\"label\"].values.astype(int)\n\n    FT_g, FC_g = extract_ft_fc_from_loader(model, train_gallery_loader, chunk_size=32)\n    FT_q, FC_q = extract_ft_fc_from_loader(model, val_query_loader, chunk_size=32)\n\n    Zg = np.concatenate([FT_g, FC_g], axis=1)\n    Zq = np.concatenate([FT_q, FC_q], axis=1)\n\n    base_qg = id_map_query_gallery(Zq, q_labels, Zg, g_labels)\n    print(f\"\\n✅ Realistic QG id-mAP (VAL query vs TRAIN gallery): {base_qg:.4f}\")\n\n    if DO_QE_DBA_QG:\n        # Best settings from your sweep\n        Qe = qe_query_with_gallery(Zq, Zg, topk=3, alpha=0.5)\n        m_qe = id_map_query_gallery(Qe, q_labels, Zg, g_labels)\n\n        Gd = dba_gallery(Zg, topk=3, alpha=0.5)\n        m_dba = id_map_query_gallery(Zq, q_labels, Gd, g_labels)\n\n        Qe2 = qe_query_with_gallery(Zq, Gd, topk=3, alpha=0.5)\n        m_both = id_map_query_gallery(Qe2, q_labels, Gd, g_labels)\n\n        print(f\"  +QE(query<-gallery) topk=3 alpha=0.5: {m_qe:.4f}\")\n        print(f\"  +DBA(gallery) topk=3 alpha=0.5:       {m_dba:.4f}\")\n        print(f\"  +DBA +QE (recommended):               {m_both:.4f}\")\n\n# ----------------------------\n# 8) Metric grad probe (SCALER-SAFE)\n# ----------------------------\nif DO_GRAD_PROBE:\n    assert \"criterion_metric\" in globals(), \"Need criterion_metric (HardCircleLoss) defined.\"\n    K = int(getattr(Config, \"K\", 6))\n    B = int(2*K)  # ensure 2 IDs\n\n    train_loader_dbg = DataLoader(train_ds, batch_sampler=train_sampler, num_workers=0, pin_memory=False)\n\n    torch.cuda.empty_cache()\n    optimizer.zero_grad(set_to_none=True)\n\n    imgs, labels = next(iter(train_loader_dbg))\n    imgs   = imgs[:B].to(Config.device, non_blocking=True).contiguous()\n    labels = labels[:B].to(Config.device, non_blocking=True).long()\n\n    print(f\"\\nGrad probe batch: B={B} unique IDs={int(labels.unique().numel())}\")\n\n    # keep autocast ON for backward (checkpoint-safe)\n    with torch.amp.autocast(device_type=Config.device_type, enabled=(Config.device_type==\"cuda\")):\n        ft, fc, arc_bn = base.embed_pair(imgs, return_bn=True)\n        logits = base.head(arc_bn.float(), labels)\n\n        loss_cls = F.cross_entropy(logits, labels, label_smoothing=0.0)\n        loss_met = criterion_metric(ft.float(), labels)\n        loss = loss_cls + 0.25 * loss_met\n\n        scaler.scale(loss).backward()\n\n    # 🔥 SCALER FIX: DO NOT call scaler.unscale_() here. \n    # Just manually divide the gradient readouts by the scaler multiplier.\n    inv_scale = 1.0 / scaler.get_scale()\n\n    def gsum(prefix):\n        s = 0.0\n        bad = 0\n        for n,p in base.named_parameters():\n            if n.startswith(prefix) and p.grad is not None:\n                # Multiply by inv_scale so we read the true gradient magnitudes\n                g = p.grad.detach() * inv_scale\n                if not torch.isfinite(g).all(): bad += 1\n                s += float(g.norm(2).item())\n        return s, bad\n\n    gb, bb = gsum(\"backbone.\")\n    gm, bm = gsum(\"metric_mlp.\")\n    ga, ba = gsum(\"arc_neck.\")\n    gh, bh = gsum(\"head.\")\n\n    print(f\"loss={float(loss.item()):.4f} | cls={float(loss_cls.item()):.4f} | met={float(loss_met.item()):.4f}\")\n    print(f\"grad backbone  : {gb:.4e} | bad={bb}\")\n    print(f\"grad metric_mlp: {gm:.4e} | bad={bm}\")\n    print(f\"grad arc_neck  : {ga:.4e} | bad={ba}\")\n    print(f\"grad head      : {gh:.4e} | bad={bh}\")\n\n    # Wipe the gradients so it doesn't affect actual training\n    optimizer.zero_grad(set_to_none=True)\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:47:18.106418Z","iopub.execute_input":"2026-03-05T07:47:18.106972Z","iopub.status.idle":"2026-03-05T07:51:30.619044Z","shell.execute_reply.started":"2026-03-05T07:47:18.106913Z","shell.execute_reply":"2026-03-05T07:51:30.618195Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL — ADVANCED GEOMETRIC & SPATIAL FORENSICS\n#   1. The Margin Waterfall (Decision Boundary Risk)\n#   2. Global Spatial Center-Bias Map\n#   3. Unsupervised Intra-Class Viewpoint Discovery\n# ============================================================\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn.functional as F\nfrom sklearn.cluster import KMeans\nfrom PIL import Image\nfrom pathlib import Path\n\n# Ensure embeddings exist\nassert \"FT\" in globals() and \"FC\" in globals(), \"Please run the extract cell first to generate FT and FC.\"\nZ = np.concatenate([FT, FC], axis=1).astype(np.float32)\nZ /= (np.linalg.norm(Z, axis=1, keepdims=True) + 1e-12)\nlabels = val_labels.astype(int)\n\n# ---------------------------------------------------------\n# 1. THE MARGIN WATERFALL (Decision Boundary Risk)\n# ---------------------------------------------------------\nprint(\"🌊 Generating Decision Boundary Waterfall...\")\nS = Z @ Z.T\nnp.fill_diagonal(S, -1e9)\nsame = labels[:, None] == labels[None, :]\nnp.fill_diagonal(same, False)\n\nmargins = []\nfor i in range(len(labels)):\n    if same[i].any():\n        best_pos = float(S[i][same[i]].max())\n        best_neg = float(S[i][~same[i]].max())\n        margins.append(best_pos - best_neg)\n\nmargins = np.sort(margins) # Sort from worst (negative) to best (positive)\n\nplt.figure(figsize=(12, 5))\nplt.plot(margins, linewidth=2, color='black')\nplt.fill_between(range(len(margins)), margins, 0, where=(margins < 0), color='red', alpha=0.5, label='Failures (Wrong Top-1)')\nplt.fill_between(range(len(margins)), margins, 0, where=((margins >= 0) & (margins < 0.10)), color='orange', alpha=0.5, label='Danger Zone (Margin < 0.10)')\nplt.fill_between(range(len(margins)), margins, 0, where=(margins >= 0.10), color='green', alpha=0.5, label='Safe Zone')\nplt.axhline(0, color='black', linestyle='--')\nplt.title(\"The Margin Waterfall: Pos_Sim vs Neg_Sim\", fontsize=14)\nplt.xlabel(\"Validation Queries (Sorted by Difficulty)\")\nplt.ylabel(\"Safety Margin (Best Pos - Best Neg)\")\nplt.legend()\nplt.grid(True, alpha=0.3)\nplt.show()\n\n# ---------------------------------------------------------\n# 2. GLOBAL SPATIAL BIAS MAP\n# ---------------------------------------------------------\nprint(\"🌍 Auditing Global Spatial Bias (Are edges being ignored?)...\")\n@torch.no_grad()\ndef generate_global_spatial_bias(model, loader, num_images=128):\n    model.eval()\n    base = model.module if isinstance(model, torch.nn.DataParallel) else model\n    \n    global_mask = None\n    count = 0\n    \n    for imgs, _ in loader:\n        imgs = imgs.to(Config.device, non_blocking=True).contiguous()\n        with torch.amp.autocast(device_type=Config.device_type, enabled=(Config.device_type == \"cuda\")):\n            cls_token, patch_map = base.forward_features_tokens(imgs)\n            \n        if patch_map is None: return None\n        \n        B, C, H, W = patch_map.shape\n        cls_norm = F.normalize(cls_token.float(), dim=1)\n        patch_flat = patch_map.float().view(B, C, H*W)\n        patch_norm = F.normalize(patch_flat, dim=1)\n        \n        attn = torch.bmm(cls_norm.unsqueeze(1), patch_norm) / 0.10\n        attn = attn - attn.max(dim=-1, keepdim=True).values\n        w = F.softmax(attn, dim=-1)\n        w = w.view(B, H, W).cpu().numpy()\n        \n        if global_mask is None:\n            global_mask = np.zeros((H, W), dtype=np.float32)\n            \n        for i in range(B):\n            if count >= num_images: break\n            # Normalize each mask to 0-1 so bright images don't dominate the average\n            w_i = w[i]\n            global_mask += (w_i - w_i.min()) / (w_i.max() - w_i.min() + 1e-8)\n            count += 1\n            \n        if count >= num_images: break\n        \n    return global_mask / count\n\n# Use val_eval_loader for deterministic results\nbias_map = generate_global_spatial_bias(model, val_eval_loader)\n\nif bias_map is not None:\n    plt.figure(figsize=(6, 6))\n    plt.imshow(bias_map, cmap='jet')\n    plt.colorbar(label='Average Attention Intensity')\n    plt.title(\"Global Spatial Bias Map (Averaged over 128 images)\")\n    plt.axis('off')\n    plt.show()\n\n# ---------------------------------------------------------\n# 3. UNSUPERVISED VIEWPOINT DISCOVERY\n# ---------------------------------------------------------\nprint(\"🔭 Running Intra-Class Viewpoint Discovery...\")\n\n# Find the ID with the most images in the validation set\ncounts = np.bincount(labels)\ntop_id = int(np.argmax(counts))\ntop_id_indices = np.where(labels == top_id)[0]\n\nprint(f\"Targeting Jaguar ID {top_id} (Contains {len(top_id_indices)} validation images)\")\n\n# Extract embeddings just for this ID\nid_embs = Z[top_id_indices]\n\n# Run K-Means to force the embeddings into 2 natural clusters\nkmeans = KMeans(n_clusters=2, random_state=42, n_init=10)\ncluster_assignments = kmeans.fit_predict(id_embs)\n\n# Visualize the split\ntry:\n    fig, axes = plt.subplots(2, 5, figsize=(15, 6))\n    fig.suptitle(f\"Unsupervised Viewpoint Discovery (ID {top_id}): Does the model cluster poses?\", fontsize=14)\n    \n    def _load_for_plot(idx):\n        fn = VAL_META.iloc[idx][\"filename\"]\n        # Try cache first, then raw dir\n        p1 = Path(TRAIN_CACHE) / fn if \"TRAIN_CACHE\" in globals() else None\n        p2 = Path(TRAIN_DIR) / fn if \"TRAIN_DIR\" in globals() else None\n        if p1 and p1.exists(): return Image.open(p1).convert(\"RGB\")\n        if p2 and p2.exists(): return Image.open(p2).convert(\"RGB\")\n        return Image.new(\"RGB\", (224,224), (128,128,128))\n\n    for c in range(2):\n        c_idx = top_id_indices[cluster_assignments == c]\n        for i in range(5):\n            ax = axes[c, i]\n            if i < len(c_idx):\n                ax.imshow(_load_for_plot(c_idx[i]))\n                ax.set_title(f\"Cluster {c}\")\n            ax.axis('off')\n            \n    plt.tight_layout()\n    plt.show()\nexcept Exception as e:\n    print(\"Could not load images for viewpoint discovery. Ensure VAL_META and directories are set.\")\n    print(\"Error:\", e)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:51:46.170282Z","iopub.execute_input":"2026-03-05T07:51:46.170954Z","iopub.status.idle":"2026-03-05T07:52:18.556693Z","shell.execute_reply.started":"2026-03-05T07:51:46.170921Z","shell.execute_reply":"2026-03-05T07:52:18.555743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL — LEVEL 3 GRANDMASTER FORENSICS (ALL-IN-ONE)\n#   1. Margin Waterfall (Decision Boundary Risk)\n#   2. The Tug-of-War (Metric vs ArcFace Disagreements)\n#   3. Manifold Black Holes (Finding the worst images)\n#   4. Dense Patch Correspondence (Biological Rosette Tracking)\n# ============================================================\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn.functional as F\nimport cv2\nfrom PIL import Image\nfrom pathlib import Path\n\n# ----------------------------\n# FLAGS\n# ----------------------------\nDO_WATERFALL    = True  # Show the safety margin curve\nDO_TUG_OF_WAR   = True  # Show where branches disagree\nDO_BLACK_HOLES  = True  # Show the most isolated outliers\nDO_PATCH_TRACK  = True  # Track dense rosettes between 2 images\n\n# ----------------------------\n# Initialization\n# ----------------------------\nprint(\"🔬 INITIATING LEVEL 3 FORENSICS...\\n\")\nassert \"FT\" in globals() and \"FC\" in globals(), \"Need FT and FC embeddings.\"\nassert \"VAL_META\" in globals(), \"Need VAL_META dataframe.\"\n\nlabels = val_labels.astype(int)\n\n# Normalize and concat\ndef l2norm(X): return X / (np.linalg.norm(X, axis=1, keepdims=True) + 1e-12)\nZ_ft = l2norm(FT)\nZ_fc = l2norm(FC)\nZ_cat = l2norm(np.concatenate([Z_ft, Z_fc], axis=1))\n\n# Global Similarity Matrix\nS = Z_cat @ Z_cat.T\nnp.fill_diagonal(S, -1e9)\nsame = labels[:, None] == labels[None, :]\nnp.fill_diagonal(same, False)\n\n# Image Loader Helper\ndef _load_img(idx):\n    fn = VAL_META.iloc[idx][\"filename\"]\n    p1 = Path(TRAIN_CACHE) / fn if \"TRAIN_CACHE\" in globals() else None\n    p2 = Path(TRAIN_DIR) / fn if \"TRAIN_DIR\" in globals() else None\n    if p1 and p1.exists(): return Image.open(p1).convert(\"RGB\")\n    if p2 and p2.exists(): return Image.open(p2).convert(\"RGB\")\n    return Image.new(\"RGB\", (224,224), (128,128,128))\n\n# ---------------------------------------------------------\n# 1. THE MARGIN WATERFALL\n# ---------------------------------------------------------\nif DO_WATERFALL:\n    print(\"🌊 1. Generating Decision Boundary Waterfall...\")\n    margins = []\n    for i in range(len(labels)):\n        if same[i].any():\n            best_pos = float(S[i][same[i]].max())\n            best_neg = float(S[i][~same[i]].max())\n            margins.append(best_pos - best_neg)\n            \n    margins = np.sort(margins)\n    plt.figure(figsize=(10, 4))\n    plt.plot(margins, linewidth=2, color='black')\n    plt.fill_between(range(len(margins)), margins, 0, where=(margins < 0), color='red', alpha=0.5, label='Fails (Margin < 0)')\n    plt.fill_between(range(len(margins)), margins, 0, where=((margins >= 0) & (margins < 0.10)), color='orange', alpha=0.5, label='Danger Zone (< 0.10)')\n    plt.fill_between(range(len(margins)), margins, 0, where=(margins >= 0.10), color='green', alpha=0.5, label='Safe Zone')\n    plt.axhline(0, color='black', linestyle='--')\n    plt.title(\"Margin Waterfall: Are your predictions safe?\")\n    plt.ylabel(\"Margin (Best Pos - Best Neg)\")\n    plt.legend()\n    plt.grid(True, alpha=0.3)\n    plt.show()\n\n# ---------------------------------------------------------\n# 2. THE TUG-OF-WAR (Branch Disagreements)\n# ---------------------------------------------------------\nif DO_TUG_OF_WAR:\n    print(\"\\n⚔️ 2. Finding Branch Disagreements...\")\n    def get_top1_correct(emb):\n        sims = emb @ emb.T\n        np.fill_diagonal(sims, -1e9)\n        return labels[np.argmax(sims, axis=1)] == labels\n\n    ft_correct = get_top1_correct(Z_ft)\n    fc_correct = get_top1_correct(Z_fc)\n\n    ft_wins = np.where(ft_correct & ~fc_correct)[0]\n    fc_wins = np.where(fc_correct & ~ft_correct)[0]\n\n    def plot_disagreements(indices, title, max_show=5):\n        if len(indices) == 0: return\n        cols = min(max_show, len(indices))\n        fig, axes = plt.subplots(1, cols, figsize=(3 * cols, 3))\n        fig.suptitle(title, fontsize=12)\n        if cols == 1: axes = [axes]\n        for i in range(cols):\n            idx = indices[i]\n            axes[i].imshow(_load_img(idx))\n            axes[i].set_title(f\"ID: {labels[idx]} | Idx: {idx}\")\n            axes[i].axis('off')\n        plt.show()\n\n    print(f\"   -> Metric (FT) won alone on {len(ft_wins)} images. ArcFace (FC) won alone on {len(fc_wins)} images.\")\n    plot_disagreements(ft_wins, \"Metric Branch Triumphs (ArcFace Failed)\")\n    plot_disagreements(fc_wins, \"ArcFace Branch Triumphs (Metric Failed)\")\n\n# ---------------------------------------------------------\n# 3. MANIFOLD BLACK HOLES\n# ---------------------------------------------------------\nif DO_BLACK_HOLES:\n    print(\"\\n🌌 3. Scanning for Manifold 'Black Holes' (Isolated Images)...\")\n    sorted_sims = np.sort(S, axis=1)[:, ::-1]\n    avg_k_sim = sorted_sims[:, :5].mean(axis=1) # Average of Top 5 neighbors\n    \n    # Lowest average similarity to anyone\n    black_holes = np.argsort(avg_k_sim)[:5] \n    \n    fig, axes = plt.subplots(1, 5, figsize=(15, 3))\n    fig.suptitle(\"The 5 Most Isolated Images in the Dataset (Lowest KNN Similarity)\", fontsize=12)\n    for i, idx in enumerate(black_holes):\n        axes[i].imshow(_load_img(idx))\n        axes[i].set_title(f\"Idx: {idx} | ID: {labels[idx]}\\nKNN Sim: {avg_k_sim[idx]:.3f}\")\n        axes[i].axis('off')\n    plt.show()\n\n# ---------------------------------------------------------\n# 4. DENSE PATCH CORRESPONDENCE (ROSETTE TRACKING)\n# ---------------------------------------------------------\nif DO_PATCH_TRACK:\n    print(\"\\n🐆 4. Biological Patch Tracking (1-to-1 Rosette Match)...\")\n    \n    # Find a class with at least 2 images to test (using the most frequent class)\n    top_id = int(np.bincount(labels).argmax())\n    idx_list = np.where(labels == top_id)[0]\n    \n    if len(idx_list) >= 2:\n        idx_A, idx_B = idx_list[0], idx_list[1]\n        img_A, img_B = _load_img(idx_A), _load_img(idx_B)\n        \n        base = model.module if isinstance(model, torch.nn.DataParallel) else model\n        base.eval()\n        \n        xA = test_transform(img_A).unsqueeze(0).to(Config.device)\n        xB = test_transform(img_B).unsqueeze(0).to(Config.device)\n        \n        with torch.no_grad(), torch.amp.autocast(device_type=Config.device_type, enabled=(Config.device_type==\"cuda\")):\n            _, patchA = base.forward_features_tokens(xA)\n            _, patchB = base.forward_features_tokens(xB)\n            \n        if patchA is not None:\n            B, C, H, W = patchA.shape\n            pA = F.normalize(patchA.view(C, -1).t(), dim=1) # [L, C]\n            pB = F.normalize(patchB.view(C, -1).t(), dim=1) # [L, C]\n            \n            sim_matrix = pA @ pB.t() # [L, L]\n            \n            # Find the strongest matched patch pair\n            best_patch_A = int(sim_matrix.max(dim=1).values.argmax())\n            best_patch_B = int(sim_matrix[best_patch_A].argmax())\n            \n            ay, ax = best_patch_A // W, best_patch_A % W\n            by, bx = best_patch_B // W, best_patch_B % W\n            \n            img_A_cv = np.array(img_A.resize((Config.img_size, Config.img_size)))\n            img_B_cv = np.array(img_B.resize((Config.img_size, Config.img_size)))\n            \n            scale_x, scale_y = Config.img_size / W, Config.img_size / H\n            pt_A = (int((ax+0.5)*scale_x), int((ay+0.5)*scale_y))\n            pt_B = (int((bx+0.5)*scale_x), int((by+0.5)*scale_y))\n            \n            cv2.circle(img_A_cv, pt_A, 12, (0, 255, 255), 4)\n            cv2.circle(img_B_cv, pt_B, 12, (0, 255, 255), 4)\n            \n            fig, axes = plt.subplots(1, 2, figsize=(10, 5))\n            fig.suptitle(f\"Dense Feature Matching (Tracking Anatomical Spot) - ID {top_id}\", fontsize=12)\n            axes[0].imshow(img_A_cv); axes[0].axis('off'); axes[0].set_title(\"Image A\")\n            axes[1].imshow(img_B_cv); axes[1].axis('off'); axes[1].set_title(\"Image B\")\n            plt.show()\n        else:\n            print(\"Model does not return patch tokens. Skipping Dense Tracking.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:52:18.558048Z","iopub.execute_input":"2026-03-05T07:52:18.558705Z","iopub.status.idle":"2026-03-05T07:52:22.192499Z","shell.execute_reply.started":"2026-03-05T07:52:18.558669Z","shell.execute_reply":"2026-03-05T07:52:22.191489Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL — LEVEL 4: PERTURBATION PHYSICS & MANIFOLD DYNAMICS\n#   1. Occlusion Sensitivity (True Pixel Importance)\n#   2. The Manifold Walk (Decision Boundary Sharpness)\n# ============================================================\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn.functional as F\nimport torchvision.transforms.functional as TF\nfrom PIL import Image\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\n\n# ----------------------------\n# FLAGS\n# ----------------------------\nDO_OCCLUSION_SENSITIVITY = True\nDO_MANIFOLD_WALK         = True\n\n# Ensure dependencies\nassert \"FT\" in globals() and \"FC\" in globals(), \"Need embeddings loaded.\"\nZ_cat = np.concatenate([FT, FC], axis=1).astype(np.float32)\nZ_cat /= (np.linalg.norm(Z_cat, axis=1, keepdims=True) + 1e-12)\nlabels = val_labels.astype(int)\n\ndef _load_img(idx):\n    fn = VAL_META.iloc[idx][\"filename\"]\n    p1 = Path(TRAIN_CACHE) / fn if \"TRAIN_CACHE\" in globals() else None\n    p2 = Path(TRAIN_DIR) / fn if \"TRAIN_DIR\" in globals() else None\n    if p1 and p1.exists(): return Image.open(p1).convert(\"RGB\")\n    if p2 and p2.exists(): return Image.open(p2).convert(\"RGB\")\n    return Image.new(\"RGB\", (448,448), (128,128,128))\n\n# ---------------------------------------------------------\n# 1. SLIDING WINDOW OCCLUSION SENSITIVITY\n# ---------------------------------------------------------\nif DO_OCCLUSION_SENSITIVITY:\n    print(\"⬛ 1. Running Occlusion Physics (Finding the True Identity Payload)...\")\n    \n    @torch.no_grad()\n    def occlusion_heatmap(val_idx, patch_size=64, stride=32):\n        base_mod = model.module if isinstance(model, torch.nn.DataParallel) else model\n        base_mod.eval()\n        \n        # Find best positive match in the dataset to act as the \"Anchor\"\n        S = Z_cat[val_idx] @ Z_cat.T\n        S[val_idx] = -1e9\n        S[labels != labels[val_idx]] = -1e9\n        best_pos_idx = int(np.argmax(S))\n        baseline_sim = float(S[best_pos_idx])\n        \n        img = _load_img(val_idx).resize((Config.img_size, Config.img_size))\n        img_t = test_transform(img) # [3, H, W]\n        _, H, W = img_t.shape\n        \n        heatmap = np.zeros((H, W), dtype=np.float32)\n        counts = np.zeros((H, W), dtype=np.float32)\n        \n        # Generate all occluded versions\n        occluded_tensors = []\n        coords = []\n        for y in range(0, H - patch_size + 1, stride):\n            for x in range(0, W - patch_size + 1, stride):\n                occ = img_t.clone()\n                # Erase the patch (simulate gray/mean pixel dropout)\n                occ[:, y:y+patch_size, x:x+patch_size] = 0.0 \n                occluded_tensors.append(occ)\n                coords.append((y, x))\n                \n        # Batch extract to save time\n        batch_size = 32\n        sim_drops = []\n        target_emb = torch.tensor(Z_cat[best_pos_idx], device=Config.device)\n        \n        for i in range(0, len(occluded_tensors), batch_size):\n            batch = torch.stack(occluded_tensors[i:i+batch_size]).to(Config.device)\n            with torch.amp.autocast(device_type=Config.device_type, enabled=(Config.device_type==\"cuda\")):\n                ft, fc = base_mod.embed_pair(batch)\n            \n            z = torch.cat([F.normalize(ft.float(), dim=1), F.normalize(fc.float(), dim=1)], dim=1)\n            z = F.normalize(z, dim=1)\n            \n            # Calculate similarity to the anchor\n            sims = (z @ target_emb).cpu().numpy()\n            sim_drops.extend(baseline_sim - sims)\n            \n        # Map drops back to spatial grid\n        for drop, (y, x) in zip(sim_drops, coords):\n            heatmap[y:y+patch_size, x:x+patch_size] += drop\n            counts[y:y+patch_size, x:x+patch_size] += 1\n            \n        heatmap = heatmap / (counts + 1e-8)\n        \n        # Plotting\n        fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n        fig.suptitle(f\"Occlusion Physics for Idx {val_idx} | Baseline Sim to Best Pos: {baseline_sim:.3f}\", fontsize=14)\n        \n        axes[0].imshow(img)\n        axes[0].set_title(\"Original Query\")\n        axes[0].axis('off')\n        \n        im = axes[1].imshow(heatmap, cmap='magma')\n        axes[1].set_title(\"Sim Drop Heatmap (Brighter = Critical Pixels)\")\n        axes[1].axis('off')\n        plt.colorbar(im, ax=axes[1], fraction=0.046, pad=0.04)\n        \n        axes[2].imshow(img)\n        axes[2].imshow(heatmap, cmap='magma', alpha=0.6)\n        axes[2].set_title(\"Overlay\")\n        axes[2].axis('off')\n        plt.show()\n\n    # Test the most frequent ID and the worst Orphan\n    top_id = int(np.bincount(labels).argmax())\n    occlusion_heatmap(np.where(labels == top_id)[0][0])\n    \n    # If you remember your orphan's index (e.g., 316), replace this with occlusion_heatmap(316)\n    if 316 in range(len(labels)):\n        occlusion_heatmap(316)\n\n\n# ---------------------------------------------------------\n# 2. THE MANIFOLD WALK (Interpolation Boundary Test)\n# ---------------------------------------------------------\nif DO_MANIFOLD_WALK:\n    print(\"\\n🚶 2. Walking the Manifold (Testing Boundary Sharpness)...\")\n    \n    # Pick two different Jaguars\n    id_A = np.unique(labels)[0]\n    id_B = np.unique(labels)[1]\n    \n    idx_A = np.where(labels == id_A)[0][0]\n    idx_B = np.where(labels == id_B)[0][0]\n    \n    z_A = Z_cat[idx_A]\n    z_B = Z_cat[idx_B]\n    \n    steps = 10\n    alphas = np.linspace(0, 1, steps)\n    \n    print(f\"Interpolating from Jaguar {id_A} to Jaguar {id_B}:\")\n    \n    for a in alphas:\n        # Linear Interpolation\n        z_t = (1.0 - a) * z_A + (a) * z_B\n        # L2 Normalize to keep it on the hypersphere (Spherical Interpolation approximation)\n        z_t = z_t / (np.linalg.norm(z_t) + 1e-12)\n        \n        # Search the entire validation gallery for the nearest neighbor\n        sims = z_t @ Z_cat.T\n        \n        # Exclude the exact start and end images so we see who else it hits\n        sims[idx_A] = -1e9\n        sims[idx_B] = -1e9\n        \n        top1_idx = int(np.argmax(sims))\n        top1_id = labels[top1_idx]\n        top1_sim = sims[top1_idx]\n        \n        # Create a visual bar for the transition\n        bar = \"█\" * int(a * 20) + \"-\" * (20 - int(a * 20))\n        \n        flag = \"✅\" if (top1_id == id_A or top1_id == id_B) else \"🚨\"\n        \n        print(f\"Step {a:.2f} [{bar}] -> Predicted ID: {top1_id:2d} (Sim: {top1_sim:.3f}) {flag}\")\n\n    print(\"\\nInterpretation:\")\n    print(\"If it snaps cleanly from ID A to ID B, your ArcFace margin created a beautiful, empty void between classes.\")\n    print(\"If you see '🚨' (it predicts Jaguar C, D, or E in the middle), your embedding space is tangled and dense.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:52:22.193653Z","iopub.execute_input":"2026-03-05T07:52:22.193990Z","iopub.status.idle":"2026-03-05T07:52:42.587035Z","shell.execute_reply.started":"2026-03-05T07:52:22.193965Z","shell.execute_reply":"2026-03-05T07:52:42.586430Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# NEW CELL — Advanced Metric Learning Diagnostics\n# ============================================================\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom collections import Counter\n\ndef advanced_reid_diagnostics(val_emb, val_labels):\n    print(\"🔍 Running Advanced Embedding Space Diagnostics...\")\n    \n    # 1. Calculate Similarity Matrix and Masks\n    sim_mat = val_emb @ val_emb.T\n    labels = np.asarray(val_labels)\n    N = len(labels)\n    \n    same_mask = labels[:, None] == labels[None, :]\n    np.fill_diagonal(same_mask, False)\n    diff_mask = ~same_mask\n    np.fill_diagonal(diff_mask, False)\n    \n    pos_sims = sim_mat[same_mask]\n    neg_sims = sim_mat[diff_mask]\n    \n    # ---------------------------------------------------------\n    # PLOT 1: Similarity Distributions\n    # ---------------------------------------------------------\n    plt.figure(figsize=(10, 5))\n    plt.hist(neg_sims, bins=100, alpha=0.6, density=True, label='Different Jaguar (Negatives)', color='red')\n    plt.hist(pos_sims, bins=100, alpha=0.6, density=True, label='Same Jaguar (Positives)', color='blue')\n    \n    # Calculate threshold where False Positives and False Negatives intersect\n    overlap_min = max(neg_sims.min(), pos_sims.min())\n    overlap_max = min(neg_sims.max(), pos_sims.max())\n    \n    plt.axvline(x=pos_sims.mean(), color='blue', linestyle='dashed', linewidth=1.5, label=f'Pos Mean: {pos_sims.mean():.2f}')\n    plt.axvline(x=neg_sims.mean(), color='red', linestyle='dashed', linewidth=1.5, label=f'Neg Mean: {neg_sims.mean():.2f}')\n    plt.axvspan(overlap_min, overlap_max, color='purple', alpha=0.2, label='The Overlap Zone (Confusion)')\n    \n    plt.title(\"Cosine Similarity Distribution: Are the classes separated?\", fontsize=14)\n    plt.xlabel(\"Cosine Similarity\")\n    plt.ylabel(\"Density\")\n    plt.legend()\n    plt.grid(True, alpha=0.3)\n    plt.show()\n\n    # ---------------------------------------------------------\n    # PLOT 2: Per-ID Autopsy (Worst Performing Jaguars)\n    # ---------------------------------------------------------\n    def average_precision_from_sorted_rel(rel_sorted: np.ndarray):\n        npos = rel_sorted.sum()\n        if npos == 0: return 0.0\n        cumsum = np.cumsum(rel_sorted)\n        precision_at_k = cumsum / (np.arange(len(rel_sorted)) + 1)\n        return float(precision_at_k[rel_sorted].sum() / npos)\n\n    APs = np.zeros(N, dtype=np.float32)\n    for i in range(N):\n        temp_sims = sim_mat[i].copy()\n        temp_sims[i] = -1e9 # ignore self\n        order = np.argsort(-temp_sims)\n        rel = (labels[order] == labels[i])\n        APs[i] = average_precision_from_sorted_rel(rel)\n\n    unique_ids = np.unique(labels)\n    id_maps = []\n    val_counts = Counter(labels)\n    \n    for y in unique_ids:\n        mean_ap = APs[labels == y].mean()\n        id_maps.append((y, mean_ap, val_counts[y]))\n        \n    # Sort by worst mAP\n    id_maps.sort(key=lambda x: x[1])\n    \n    print(\"\\n🚨 THE 'GHOST' JAGUARS (Bottom 5 Worst Performing IDs) 🚨\")\n    print(f\"{'True ID':<10} | {'Val mAP':<10} | {'Images in Val Set':<15}\")\n    print(\"-\" * 40)\n    for y, m_ap, count in id_maps[:5]:\n        print(f\"{y:<10} | {m_ap:<10.4f} | {count:<15}\")\n        \n    print(\"\\n🏆 THE 'EASY' JAGUARS (Top 3 Best Performing IDs) 🏆\")\n    print(f\"{'True ID':<10} | {'Val mAP':<10} | {'Images in Val Set':<15}\")\n    print(\"-\" * 40)\n    for y, m_ap, count in reversed(id_maps[-3:]):\n        print(f\"{y:<10} | {m_ap:<10.4f} | {count:<15}\")\n\n# Run it!\nadvanced_reid_diagnostics(val_emb, val_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:52:42.588316Z","iopub.execute_input":"2026-03-05T07:52:42.588574Z","iopub.status.idle":"2026-03-05T07:52:43.072750Z","shell.execute_reply.started":"2026-03-05T07:52:42.588547Z","shell.execute_reply":"2026-03-05T07:52:43.072106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# NEW CELL — Deep Structural Diagnostics (CMC & t-SNE)\n# ============================================================\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.manifold import TSNE\n\ndef deep_reid_diagnostics(val_emb, val_labels):\n    print(\"🔍 Running Deep Structural Diagnostics...\")\n    labels = np.asarray(val_labels)\n    N = len(labels)\n    sim_mat = val_emb @ val_emb.T\n    \n    # ---------------------------------------------------------\n    # PLOT 3: Cumulative Matching Characteristics (CMC)\n    # ---------------------------------------------------------\n    # CMC measures if the correct ID is within the Top-K closest embeddings\n    max_k = 10\n    cmc = np.zeros(max_k)\n    valid_queries = 0\n    \n    for i in range(N):\n        # Only evaluate if there's actually another image of this jaguar in the val set\n        if (labels == labels[i]).sum() > 1:\n            valid_queries += 1\n            sims = sim_mat[i].copy()\n            sims[i] = -1e9  # Ignore self\n            \n            # Get indices of top max_k matches\n            top_k_idx = np.argsort(-sims)[:max_k]\n            top_k_labels = labels[top_k_idx]\n            \n            # Check if true label is in top K\n            match_found = False\n            for k in range(max_k):\n                if top_k_labels[k] == labels[i]:\n                    match_found = True\n                if match_found:\n                    cmc[k] += 1\n                    \n    cmc = cmc / valid_queries\n    \n    plt.figure(figsize=(8, 5))\n    plt.plot(np.arange(1, max_k + 1), cmc, marker='o', linestyle='-', color='b')\n    plt.title(\"Cumulative Matching Characteristics (CMC)\", fontsize=14)\n    plt.xlabel(\"Rank (Top-K)\")\n    plt.ylabel(\"Matching Probability\")\n    plt.ylim(0, 1.05)\n    plt.xticks(np.arange(1, max_k + 1))\n    plt.grid(True, alpha=0.3)\n    for k in [0, 4, 9]:  # Annotate Rank-1, Rank-5, Rank-10\n        plt.text(k + 1, cmc[k] + 0.02, f\"{cmc[k]:.3f}\", ha='center', fontsize=10, fontweight='bold')\n    plt.show()\n\n    # ---------------------------------------------------------\n    # PLOT 4: The \"Island\" Test (t-SNE Projection)\n    # ---------------------------------------------------------\n    # Let's pick 2 \"Easy\" IDs and 2 \"Ghost\" IDs based on your previous logs\n    # Update these if your specific Ghost/Easy IDs changed\n    easy_ids = [19, 20]     # High mAP, many images\n    ghost_ids = [21, 8]     # Low mAP, Few-Shot or Viewpoint issues\n    \n    target_ids = easy_ids + ghost_ids\n    mask = np.isin(labels, target_ids)\n    \n    filtered_emb = val_emb[mask]\n    filtered_labels = labels[mask]\n    \n    print(\"⏳ Calculating t-SNE 2D Projection (this takes a few seconds)...\")\n    tsne = TSNE(n_components=2, perplexity=min(30, len(filtered_emb)-1), random_state=42)\n    emb_2d = tsne.fit_transform(filtered_emb)\n    \n    plt.figure(figsize=(10, 8))\n    colors = ['green', 'lime', 'red', 'darkred']\n    markers = ['o', 's', '^', 'X']\n    \n    for idx, (target_id, color, marker) in enumerate(zip(target_ids, colors, markers)):\n        id_mask = (filtered_labels == target_id)\n        status = \"EASY\" if target_id in easy_ids else \"GHOST\"\n        plt.scatter(emb_2d[id_mask, 0], emb_2d[id_mask, 1], \n                    c=color, marker=marker, label=f'ID {target_id} ({status})', \n                    s=80, edgecolors='k', alpha=0.8)\n        \n    plt.title(\"t-SNE Embedding Space: The 'Island' Test\", fontsize=14)\n    plt.legend(fontsize=12)\n    plt.grid(True, alpha=0.3)\n    plt.show()\n\n# Run it!\ndeep_reid_diagnostics(val_emb, val_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-04T22:48:31.590380Z","iopub.execute_input":"2026-03-04T22:48:31.590710Z","iopub.status.idle":"2026-03-04T22:48:32.775905Z","shell.execute_reply.started":"2026-03-04T22:48:31.590683Z","shell.execute_reply":"2026-03-04T22:48:32.775207Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# NEW CELL — Deep Core Forensics (v5 with Concat-Safe Math)\n# ============================================================\nimport torch\nimport torch.nn.functional as F\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom pathlib import Path\nimport cv2\n\ndef run_v5_core_forensics(model, val_df, img_dir, val_emb, val_labels, target_id=21):\n    print(f\"🔬 RUNNING V5 CORE FORENSICS ON ID {target_id}...\\n\")\n    model.eval()\n    base = model.module if hasattr(model, 'module') else model\n    \n    labels = np.asarray(val_labels)\n    id_indices = np.where(labels == target_id)[0]\n    \n    if len(id_indices) == 0:\n        print(f\"❌ ID {target_id} not found in validation set.\")\n        return\n        \n    print(f\"Found {len(id_indices)} images for ID {target_id}.\")\n    \n    # ---------------------------------------------------------\n    # TEST A: Sub-Center Routing\n    # ---------------------------------------------------------\n    print(\"\\n🗺️ TEST A: SUB-CENTER ROUTING\")\n    if hasattr(base, 'head') and hasattr(base.head, 'weight'):\n        w = base.head.weight.detach().cpu().float()\n        \n        if w.dim() == 2:\n            C_times_K, D = w.shape\n            K_centers = C_times_K // len(np.unique(labels)) \n            w = w.view(-1, K_centers, D)\n        else:\n            K_centers = w.shape[1]\n            \n        my_centers = F.normalize(w[target_id], dim=1) # [K, D]\n        my_embs = torch.tensor(val_emb[id_indices]).float() # [N, 2048]\n        \n        # 🔥 THE FIX: Safely slice concatenated embeddings\n        D_head = my_centers.shape[1]\n        if my_embs.shape[1] != D_head:\n            print(f\"  -> [Info] Your val_emb is {my_embs.shape[1]}D, but the ArcFace head is {D_head}D.\")\n            print(f\"  -> [Info] Slicing embedding to match ArcFace math...\")\n            # If using concat(feat_tri, feat_cls), ArcFace is mapped to feat_cls (the second half)\n            my_embs = my_embs[:, -D_head:] \n            my_embs = F.normalize(my_embs, dim=1) # Re-normalize after slicing\n        \n        routing_sims = my_embs @ my_centers.T # [N, K]\n        routed_centers = torch.argmax(routing_sims, dim=1).numpy()\n        \n        print(f\"  -> Out of {len(id_indices)} images for ID {target_id}:\")\n        for k in range(K_centers):\n            count = (routed_centers == k).sum()\n            avg_sim = routing_sims[routed_centers == k, k].mean().item() if count > 0 else 0.0\n            print(f\"     * Center {k}: Claims {count} images (Avg Sim: {avg_sim:.4f})\")\n            \n        if len(np.unique(routed_centers)) == 1:\n            print(\"  🚨 WARNING: Sub-Center Collapse! The model crams all viewpoints into one center.\")\n        else:\n            print(\"  ✅ HEALTHY: The model successfully routes different viewpoints to different centers!\")\n    else:\n        print(\"  -> Could not find ArcFace weights.\")\n\n    # ---------------------------------------------------------\n    # TEST B: DINOv3 Attention Heatmaps\n    # ---------------------------------------------------------\n    print(\"\\n🔥 TEST B: ATTENTION HEATMAPS (The 'Mean-Scale' check)\")\n    plot_indices = id_indices[:min(3, len(id_indices))]\n    fig, axes = plt.subplots(len(plot_indices), 2, figsize=(10, 4 * len(plot_indices)))\n    fig.suptitle(f\"v5 Attention Mask (ID {target_id})\", fontsize=16, y=1.02)\n    \n    from torchvision import transforms\n    MEAN = [0.485, 0.456, 0.406]\n    STD  = [0.229, 0.224, 0.225]\n    tfm = transforms.Compose([\n        transforms.Resize((448, 448)),\n        transforms.ToTensor(),\n        transforms.Normalize(MEAN, STD),\n    ])\n    \n    for row, idx in enumerate(plot_indices):\n        fname = val_df.iloc[idx]['filename']\n        p = Path(img_dir) / fname \n        try:\n            pil_img = Image.open(p).convert(\"RGB\")\n        except:\n            continue\n            \n        img_tensor = tfm(pil_img).unsqueeze(0).to(next(model.parameters()).device)\n        \n        with torch.inference_mode():\n            with torch.amp.autocast(device_type=\"cuda\" if torch.cuda.is_available() else \"cpu\"):\n                cls_token, patch_map = base.forward_features_tokens(img_tensor)\n                \n                B, C, H, W = patch_map.shape\n                L = H * W\n                cls_norm = F.normalize(cls_token.float(), dim=1)\n                patch_flat = patch_map.float().view(B, C, L)\n                patch_norm = F.normalize(patch_flat, dim=1)\n                \n                attn = torch.bmm(cls_norm.unsqueeze(1), patch_norm) / 0.10\n                attn = attn - attn.max(dim=-1, keepdim=True).values\n                w = F.softmax(attn, dim=-1)\n                \n                # V5 MEAN SCALING\n                w_v5 = w / (w.mean(dim=-1, keepdim=True) + 1e-6)\n                \n                max_multiplier = w_v5.max().item()\n                if row == 0:\n                    print(f\"  -> Image 0: The hottest patch is multiplied by {max_multiplier:.2f}x before GeM!\")\n                \n                attn_mask = w_v5.view(H, W).cpu().numpy()\n                attn_mask = (attn_mask - attn_mask.min()) / (attn_mask.max() - attn_mask.min() + 1e-8)\n        \n        heatmap = cv2.resize(attn_mask, (pil_img.size[0], pil_img.size[1]))\n        heatmap_color = cv2.applyColorMap(np.uint8(255 * heatmap), cv2.COLORMAP_JET)\n        heatmap_color = cv2.cvtColor(heatmap_color, cv2.COLOR_BGR2RGB)\n        \n        overlay = cv2.addWeighted(np.array(pil_img), 0.5, heatmap_color, 0.5, 0)\n        \n        ax1, ax2 = axes[row] if len(plot_indices) > 1 else axes\n        \n        ax1.imshow(pil_img)\n        ax1.set_title(f\"Original (Index {idx})\")\n        ax1.axis('off')\n        \n        ax2.imshow(overlay)\n        ax2.set_title(f\"v5 Mask (Peak Mult: {max_multiplier:.1f}x)\")\n        ax2.axis('off')\n        \n    plt.tight_layout()\n    plt.show()\n\n# Run it!\nrun_v5_core_forensics(model, va_df, TRAIN_DIR, val_emb, val_labels, target_id=21)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:55:24.330979Z","iopub.execute_input":"2026-03-05T07:55:24.331685Z","iopub.status.idle":"2026-03-05T07:55:26.495549Z","shell.execute_reply.started":"2026-03-05T07:55:24.331651Z","shell.execute_reply":"2026-03-05T07:55:26.494459Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# NEW CELL — Visualize examples from the smallest classes\n# ============================================================\n\nimport random\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom pathlib import Path\n\nTRAIN_DIR = Path(globals().get(\"TRAIN_DIR\", \"/kaggle/input/jaguar-re-id/train/train\"))\n\n# pick N smallest classes by TRAIN_SPLIT counts if available, else overall\nif \"tr_df\" in globals():\n    dfv = tr_df.copy()\n    if \"label\" not in dfv.columns:\n        dfv[\"label\"] = dfv[\"ground_truth\"].map(label_map).astype(int)\nelse:\n    dfv = train_df.copy()\n\ncounts = dfv.groupby(\"label\")[\"filename\"].count().sort_values()\nsmall_labels = counts.head(6).index.tolist()\n\nprint(\"Smallest labels:\", [(int(l), int(counts.loc[l])) for l in small_labels])\n\n# show M images per class\nM = 6\nrows = len(small_labels)\ncols = M\n\nplt.figure(figsize=(2.2*cols, 2.2*rows))\nplot_i = 1\n\nfor y in small_labels:\n    sub = dfv[dfv[\"label\"] == y][\"filename\"].tolist()\n    picks = random.sample(sub, k=min(M, len(sub)))\n\n    for fn in picks:\n        p = TRAIN_DIR / fn\n        try:\n            img = Image.open(p)\n            img = preprocess_pil(img) if \"preprocess_pil\" in globals() else img.convert(\"RGB\")\n        except Exception:\n            img = Image.new(\"RGB\", (256,256), (128,128,128))\n\n        ax = plt.subplot(rows, cols, plot_i)\n        ax.imshow(img)\n        ax.axis(\"off\")\n        if plot_i % cols == 1:\n            name = inv_label.get(int(y), str(int(y))) if \"inv_label\" in globals() else str(int(y))\n            ax.set_title(f\"label {int(y)} ({name})\", fontsize=10)\n        plot_i += 1\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL — THE ULTIMATE SUBMISSION ENGINE (v3)\n#   + CONF-gated FC concat (alpha=0.5, p=2)\n#   + Isolated Q/G spaces (Fixes the overlapping ID bug)\n#   + DBA(gallery) + QE(query <- gallery)\n#   + Self-match override safety\n# ============================================================\n\nimport gc, os\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\nfrom pathlib import Path\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader\n\n# ----------------------------\n# Helpers\n# ----------------------------\ndef l2norm(x, eps=1e-12):\n    n = np.linalg.norm(x, axis=1, keepdims=True)\n    return x / (n + eps)\n\ndef strip_module(sd):\n    if any(k.startswith(\"module.\") for k in sd.keys()):\n        return {k.replace(\"module.\", \"\", 1): v for k, v in sd.items()}\n    return sd\n\n@torch.no_grad()\ndef extract_ft_fc_conf_safe(model, loader, chunk_size=32):\n    model.eval()\n    base = model.module if isinstance(model, torch.nn.DataParallel) else model\n\n    FT_list, FC_list, CF_list, names = [], [], [], []\n\n    for imgs, fnames in tqdm(loader, desc=\"Extract(ft/fc/conf)\"):\n        for mb in torch.split(imgs, int(chunk_size), dim=0):\n            mb = mb.to(Config.device, non_blocking=True).contiguous()\n\n            with torch.amp.autocast(device_type=Config.device_type, enabled=(Config.device_type==\"cuda\")):\n                if \"return_bn\" in base.embed_pair.__code__.co_varnames:\n                    ft, fc, arc_bn = base.embed_pair(mb, return_bn=True)\n                    logits = base.head(arc_bn.float(), label=None)  \n                else:\n                    ft, fc = base.embed_pair(mb)\n                    logits = None\n\n                ft = F.normalize(ft.float(), dim=1)\n                fc = F.normalize(fc.float(), dim=1)\n\n                if logits is not None:\n                    probs = F.softmax(logits.float(), dim=1)\n                    conf = probs.max(dim=1).values\n                else:\n                    conf = torch.ones((ft.size(0),), device=ft.device, dtype=torch.float32)\n\n            FT_list.append(ft.cpu().numpy())\n            FC_list.append(fc.cpu().numpy())\n            CF_list.append(conf.float().cpu().numpy())\n\n        names.extend(list(fnames))\n\n    FT = np.concatenate(FT_list, axis=0).astype(np.float32)\n    FC = np.concatenate(FC_list, axis=0).astype(np.float32)\n    CF = np.concatenate(CF_list, axis=0).astype(np.float32)\n    return FT, FC, CF, names\n\ndef mix_conf(FT, FC, CF, alpha=0.5, p=2.0):\n    w = (CF ** float(p)).reshape(-1, 1).astype(np.float32)\n    Z = np.concatenate([FT, (float(alpha) * w * FC)], axis=1).astype(np.float32)\n    return l2norm(Z)\n\ndef dba_gallery(G, topk=3, alpha=0.5):\n    G = l2norm(G.astype(np.float32))\n    S = G @ G.T\n    np.fill_diagonal(S, -1e9)\n    nn = np.argsort(-S, axis=1)[:, :int(topk)]\n    G2 = G.copy()\n    for i in range(len(G)):\n        G2[i] = G[i] + float(alpha) * G[nn[i]].mean(axis=0)\n    return l2norm(G2)\n\ndef qe_query_with_gallery(Q, G, topk=3, alpha=0.5):\n    Q = l2norm(Q.astype(np.float32))\n    G = l2norm(G.astype(np.float32))\n    S = Q @ G.T\n    nn = np.argsort(-S, axis=1)[:, :int(topk)]\n    Q2 = Q.copy()\n    for i in range(len(Q)):\n        Q2[i] = Q[i] + float(alpha) * G[nn[i]].mean(axis=0)\n    return l2norm(Q2)\n\n# ----------------------------\n# 1) Load best checkpoint\n# ----------------------------\nprint(\"🚀 INITIATING GRANDMASTER SUBMISSION PIPELINE...\")\nCKPT = Path(getattr(Config, \"ckpt_dir\", \"/kaggle/working/ckpts\")) / \"best_v3.pt\"\nassert CKPT.exists(), f\"Missing checkpoint: {CKPT}\"\n\nif \"model\" not in globals():\n    base = ReIDBoss(num_classes=num_classes).to(Config.device)\n    model = nn.DataParallel(base) if torch.cuda.device_count() >= 2 else base\n\nstate = torch.load(CKPT, map_location=\"cpu\")\nif isinstance(state, dict) and \"state_dict\" in state:\n    state = state[\"state_dict\"]\nstate = strip_module(state)\n\nm = model.module if isinstance(model, nn.DataParallel) else model\nm.load_state_dict(state, strict=True)\nmodel.eval()\nprint(\"✅ Loaded:\", str(CKPT))\n\n# ----------------------------\n# 2) Build isolated sets (queries vs galleries)\n# ----------------------------\nassert \"test_df\" in globals(), \"Need test_df loaded\"\nassert {\"row_id\",\"query_image\",\"gallery_image\"}.issubset(set(test_df.columns))\n\nunique_queries = test_df[\"query_image\"].unique()\nunique_galleries = test_df[\"gallery_image\"].unique()\nunique_imgs = sorted(list(set(unique_queries) | set(unique_galleries)))\n\ntest_ds = JaguarDataset(\n    pd.DataFrame({\"filename\": unique_imgs}), TEST_DIR, transform=test_transform, is_test=True,\n    cache_dir=(TEST_CACHE if (\"TEST_CACHE\" in globals() and getattr(Config, \"cache_crops\", False)) else None)\n)\n\ntest_loader = DataLoader(\n    test_ds, batch_size=64, shuffle=False, num_workers=int(getattr(Config, \"num_workers\", 4)),\n    pin_memory=bool(getattr(Config, \"pin_memory\", True)),\n    persistent_workers=bool(getattr(Config, \"persistent_workers\", True)) and int(getattr(Config, \"num_workers\", 4)) > 0\n)\n\n# ----------------------------\n# 3) Extract & Build Base Embeddings\n# ----------------------------\nFT_all, FC_all, CF_all, names = extract_ft_fc_conf_safe(model, test_loader, chunk_size=32)\nname_to_idx = {n: i for i, n in enumerate(names)}\n\n# Confidence-gated feature fusion\nZ_all = mix_conf(FT_all, FC_all, CF_all, alpha=0.5, p=2.0)\n\n# ----------------------------\n# 4) ISOLATE SPACES & Apply DBA / QE\n# ----------------------------\n# Extract pure sets to prevent Gallery-smoothing from overwriting Query-expansion\nq_indices = np.array([name_to_idx[n] for n in unique_queries])\ng_indices = np.array([name_to_idx[n] for n in unique_galleries])\n\nZ_uq = Z_all[q_indices]\nZ_ug = Z_all[g_indices]\n\nDO_DBA_QE = True\nif DO_DBA_QE:\n    print(\"✨ Applying DBA to Gallery Space...\")\n    Z_ug = dba_gallery(Z_ug, topk=3, alpha=0.5)\n    \n    print(\"✨ Applying Query Expansion against Gallery...\")\n    Z_uq = qe_query_with_gallery(Z_uq, Z_ug, topk=3, alpha=0.5)\n\n# ----------------------------\n# 5) Memory-Safe Pair Scoring\n# ----------------------------\nprint(\"🧮 Scoring pairs...\")\n# Map df rows to their new isolated indices\nq_map = {n: i for i, n in enumerate(unique_queries)}\ng_map = {n: i for i, n in enumerate(unique_galleries)}\n\nq_score_idx = test_df[\"query_image\"].map(q_map).values\ng_score_idx = test_df[\"gallery_image\"].map(g_map).values\n\n# Dot product (Cosine Similarity)\nsims = np.sum(Z_uq[q_score_idx] * Z_ug[g_score_idx], axis=1).astype(np.float32)\n\n# ----------------------------\n# 6) Final Formatting & Overrides\n# ----------------------------\n# Map to [0,1]\npreds = (sims + 1.0) / 2.0\n\n# 🛡️ THE SAFETY OVERRIDE: If query_image == gallery_image, probability must be exactly 1.0\nsame_mask = (test_df[\"query_image\"] == test_df[\"gallery_image\"]).values\npreds[same_mask] = 1.0\n\npreds = np.clip(preds, 0.0, 1.0)\n\nsub = pd.DataFrame({\"row_id\": test_df[\"row_id\"].values, \"similarity\": preds})\nsub.to_csv(\"submission.csv\", index=False)\n\nprint(\"\\n🏆 SUCCESS! submission.csv saved.\")\nprint(f\"-> Mean: {float(preds.mean()):.4f} | Min: {float(preds.min()):.4f} | Max: {float(preds.max()):.4f}\")\n\n# Kaggle Commit Safety Check (Dummy test set is small, hidden set is large)\nexpected_len = len(test_df)\nassert len(sub) == expected_len, f\"Length mismatch: {len(sub)} vs {expected_len}\"\nassert np.isfinite(sub[\"similarity\"].values).all()\n\n# Cleanup\ndel FT_all, FC_all, CF_all, Z_all, Z_uq, Z_ug\ngc.collect()\ntorch.cuda.empty_cache()\ndisplay(sub.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:01:51.901493Z","iopub.execute_input":"2026-03-05T08:01:51.902000Z","iopub.status.idle":"2026-03-05T08:02:19.283862Z","shell.execute_reply.started":"2026-03-05T08:01:51.901965Z","shell.execute_reply":"2026-03-05T08:02:19.283149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 15 — Inference + (Optional) QE + Submission ✅ BETTER\n#   Improvements:\n#     - QE uses top-k neighbors EXCLUDING self, and mixes self + neighbor-mean (more stable)\n#     - QE is done separately for query and gallery (prevents \"query-query\" leakage)\n#     - Vectorized QE (fast)\n#     - Uses float32 throughout, safe normalization\n#     - Optional: compute similarity via cosine then map to [0,1]\n# ============================================================\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader\n\n# ----------------------------\n# QE helpers (vectorized)\n# ----------------------------\ndef l2norm(x: np.ndarray, eps=1e-12):\n    return x / (np.linalg.norm(x, axis=1, keepdims=True) + eps)\n\ndef qe_mix_topk(emb: np.ndarray, topk: int = 3, alpha: float = 0.5, exclude_self: bool = True):\n    \"\"\"\n    emb: NxD, assumed normalized\n    returns: NxD normalized\n    new = alpha * emb + (1-alpha) * mean(topk neighbors)\n    \"\"\"\n    emb = emb.astype(np.float32, copy=False)\n    N = emb.shape[0]\n    sims = emb @ emb.T  # cosine\n\n    if exclude_self:\n        np.fill_diagonal(sims, -1e9)\n\n    idx = np.argsort(-sims, axis=1)[:, :topk]     # NxK\n    neigh = emb[idx].mean(axis=1)                 # NxD\n    out = alpha * emb + (1.0 - alpha) * neigh\n    return l2norm(out)\n\n# ----------------------------\n# Load best checkpoint\n# ----------------------------\nstate = torch.load(\"/kaggle/working/ckpts/best.pt\", map_location=Config.device)\nif isinstance(model, nn.DataParallel):\n    model.module.load_state_dict(state)\nelse:\n    model.load_state_dict(state)\nmodel.eval()\n\n# ----------------------------\n# Build list of unique images (all query + all gallery)\n# ----------------------------\nunique_test = sorted(set(test_df[\"query_image\"]) | set(test_df[\"gallery_image\"]))\n\ntest_ds = JaguarDataset(\n    pd.DataFrame({\"filename\": unique_test}),\n    TEST_DIR,\n    transform=test_transform,\n    is_test=True,\n    cache_dir=(TEST_CACHE if Config.cache_crops else None),\n)\n\ntest_loader = DataLoader(\n    test_ds,\n    batch_size=32,  # ✅ a bit larger for faster inference (T4 usually ok)\n    shuffle=False,\n    num_workers=Config.num_workers,\n    pin_memory=Config.pin_memory,\n    persistent_workers=Config.persistent_workers,\n    prefetch_factor=Config.prefetch_factor,\n)\n\n# extract_embeddings already normalizes, but we re-normalize safely\nemb, names = extract_embeddings(model, test_loader)\nemb = emb.astype(np.float32, copy=False)\nemb = l2norm(emb)\n\nimg_map = {n: i for i, n in enumerate(names)}\n\n# ----------------------------\n# Separate QE for queries and galleries (recommended)\n# ----------------------------\nq_idx = test_df[\"query_image\"].map(img_map).values\ng_idx = test_df[\"gallery_image\"].map(img_map).values\n\n# unique indices for query and gallery sets\nuq = np.unique(q_idx)\nug = np.unique(g_idx)\n\nif getattr(Config, \"use_qe\", False):\n    # Use small topk (your earlier sweep showed big k hurts)\n    k = int(getattr(Config, \"qe_topk\", 3))\n    k = max(1, min(k, 20))\n    alpha = 0.5  # mix weight; 0.5 is stable\n\n    print(f\"Applying QE separately: k={k}, alpha={alpha} (exclude self)\")\n\n    # QE within each set only (query set and gallery set)\n    emb_q = emb[uq]\n    emb_g = emb[ug]\n\n    emb_q = qe_mix_topk(emb_q, topk=k, alpha=alpha, exclude_self=True)\n    emb_g = qe_mix_topk(emb_g, topk=k, alpha=alpha, exclude_self=True)\n\n    # write back\n    emb2 = emb.copy()\n    emb2[uq] = emb_q\n    emb2[ug] = emb_g\n    emb = emb2\n\n# ----------------------------\n# Compute similarity only for required pairs (fast & memory-safe)\n# ----------------------------\n# cosine similarity in [-1,1]\npreds = np.sum(emb[q_idx] * emb[g_idx], axis=1).astype(np.float32)\n\n# map to [0,1]\npreds = (preds + 1.0) / 2.0\npreds = np.clip(preds, 0.0, 1.0)\n\nsub = pd.DataFrame({\"row_id\": test_df[\"row_id\"].values, \"similarity\": preds})\nsub.to_csv(\"submission.csv\", index=False)\n\nprint(\"✅ submission.csv saved\")\nprint(\"Mean:\", float(preds.mean()), \"Min:\", float(preds.min()), \"Max:\", float(preds.max()))\n\n# Validate\nassert len(sub) == 137270\nassert sub[\"row_id\"].iloc[0] == 0 and sub[\"row_id\"].iloc[-1] == 137269\nassert np.isfinite(sub[\"similarity\"].values).all()\nassert (sub[\"similarity\"].values >= 0).all() and (sub[\"similarity\"].values <= 1).all()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Imports\n","metadata":{}},{"cell_type":"code","source":"import os\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport timm","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Config\n","metadata":{}},{"cell_type":"code","source":"class Config:\n    seed = 42\n    model_name = \"eva02_large_patch14_448.mim_m38m_ft_in22k_in1k\"\n\n    img_size = 448\n    embedding_dim = 1024\n    num_classes = 31\n\n    num_epochs = 10\n    batch_size = 4\n    grad_accum = 4\n\n    lr = 2e-5\n    weight_decay = 1e-3\n\n    arcface_s = 30.0\n    arcface_m = 0.50\n\n    use_tta = True\n    use_qe = True\n    use_rerank = True\n\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    device_type = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n\n\nseed_everything(Config.seed)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset\n","metadata":{}},{"cell_type":"code","source":"class JaguarDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None, is_test=False):\n        self.df = df\n        self.img_dir = Path(img_dir)\n        self.transform = transform\n        self.is_test = is_test\n        if not is_test:\n            unique_ids = sorted(df[\"ground_truth\"].unique())\n            self.label_map = {name: i for i, name in enumerate(unique_ids)}\n            self.df[\"label\"] = self.df[\"ground_truth\"].map(self.label_map)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_name = row[\"filename\"]\n        img_path = self.img_dir / img_name\n        try:\n            img = Image.open(img_path).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (Config.img_size, Config.img_size))\n\n        if self.transform:\n            img = self.transform(img)\n        if self.is_test:\n            return img, img_name\n        return img, torch.tensor(row[\"label\"], dtype=torch.long)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Transforms\n","metadata":{}},{"cell_type":"code","source":"train_transform = transforms.Compose(\n    [\n        transforms.Resize((Config.img_size, Config.img_size)),\n        transforms.RandomHorizontalFlip(),\n        # transforms.RandomAffine(degrees=15, translate=(0.1, 0.1), scale=(0.9, 1.1)),\n        transforms.RandomAffine(degrees=15, translate=(0.15, 0.15), scale=(0.85, 1.15)),\n        transforms.ColorJitter(brightness=0.2, contrast=0.2),\n        transforms.ToTensor(),\n        transforms.Normalize([0.481, 0.457, 0.408], [0.268, 0.261, 0.275]),\n        transforms.RandomErasing(p=0.25),\n    ]\n)\n\ntest_transform = transforms.Compose(\n    [\n        transforms.Resize((Config.img_size, Config.img_size)),\n        transforms.ToTensor(),\n        transforms.Normalize([0.481, 0.457, 0.408], [0.268, 0.261, 0.275]),\n    ]\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model\n","metadata":{}},{"cell_type":"code","source":"class GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1) * p)\n        self.eps = eps\n\n    def forward(self, x):\n        return F.avg_pool2d(\n            x.clamp(min=self.eps).pow(self.p), (x.size(-2), x.size(-1))\n        ).pow(1.0 / self.p)\n\n\nclass ArcFaceLayer(nn.Module):\n    def __init__(self, in_features, out_features, s=30.0, m=0.5):\n        super().__init__()\n        self.s = s\n        self.m = m\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n    def forward(self, input, label=None):\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        if label is None:\n            return cosine\n        phi = cosine - self.m\n        one_hot = torch.zeros_like(cosine)\n        one_hot.scatter_(1, label.view(-1, 1), 1)\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        return output * self.s\n\n\nclass EVABoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = timm.create_model(\n            Config.model_name, pretrained=True, num_classes=0\n        )\n        self.feat_dim = self.backbone.num_features\n        self.gem = GeM()\n        self.bn = nn.BatchNorm1d(self.feat_dim)\n        self.head = ArcFaceLayer(\n            self.feat_dim, Config.num_classes, s=Config.arcface_s, m=Config.arcface_m\n        )\n\n    def forward(self, x, label=None):\n        features = self.backbone.forward_features(x)\n        if features.dim() == 3:\n            B, N, C = features.shape\n            H = W = int(math.sqrt(N))\n            if H * W != N:\n                features = features[:, -H * W :, :]\n            features = features.permute(0, 2, 1).reshape(B, C, H, W)\n\n        emb = self.gem(features).flatten(1)\n        emb = self.bn(emb)\n        if label is not None:\n            return self.head(emb, label)\n        return emb","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Utils\n","metadata":{}},{"cell_type":"code","source":"def train_epoch(model, loader, optimizer, criterion, scaler):\n    model.train()\n    loss_meter = 0\n    for i, (imgs, labels) in enumerate(tqdm(loader, leave=False)):\n        imgs, labels = imgs.to(Config.device), labels.to(Config.device)\n\n        with torch.amp.autocast(Config.device_type):\n            loss = criterion(model(imgs, labels), labels)\n            loss = loss / Config.grad_accum\n\n        scaler.scale(loss).backward()\n\n        if (i + 1) % Config.grad_accum == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n\n        loss_meter += loss.item() * Config.grad_accum\n    return loss_meter / len(loader)\n\n\n@torch.no_grad()\ndef extract_features(model, loader):\n    model.eval()\n    feats, names = [], []\n    for imgs, fnames in tqdm(loader, desc=\"Inference\"):\n        imgs = imgs.to(Config.device)\n        f1 = model(imgs)\n        if Config.use_tta:\n            f2 = model(torch.flip(imgs, [3]))\n            f1 = (f1 + f2) / 2\n        feats.append(F.normalize(f1, dim=1).cpu())\n        names.extend(fnames)\n    return torch.cat(feats, dim=0).numpy(), names\n\n\ndef query_expansion(emb, top_k=3):\n    print(\"Applying QE...\")\n    sims = emb @ emb.T\n    indices = np.argsort(-sims, axis=1)[:, :top_k]\n    new_emb = np.zeros_like(emb)\n    for i in range(len(emb)):\n        new_emb[i] = np.mean(emb[indices[i]], axis=0)\n    return new_emb / np.linalg.norm(new_emb, axis=1, keepdims=True)\n\n\ndef k_reciprocal_rerank(prob, k1=20, k2=6, lambda_value=0.3):\n    print(\"Applying Re-ranking...\")\n    q_g_dist = 1 - prob\n    original_dist = q_g_dist.copy()\n    initial_rank = np.argsort(original_dist, axis=1)\n    nn_k1 = []\n    for i in range(prob.shape[0]):\n        forward_k1 = initial_rank[i, : k1 + 1]\n        backward_k1 = initial_rank[forward_k1, : k1 + 1]\n        fi = np.where(backward_k1 == i)[0]\n        nn_k1.append(forward_k1[fi])\n    jaccard_dist = np.zeros_like(original_dist)\n    for i in range(prob.shape[0]):\n        ind_non_zero = np.where(original_dist[i, :] < 0.6)[0]\n        ind_images = [\n            inv for inv in ind_non_zero if len(np.intersect1d(nn_k1[i], nn_k1[inv])) > 0\n        ]\n        for j in ind_images:\n            intersection = len(np.intersect1d(nn_k1[i], nn_k1[j]))\n            union = len(np.union1d(nn_k1[i], nn_k1[j]))\n            jaccard_dist[i, j] = 1 - intersection / union\n    return 1 - (jaccard_dist * lambda_value + original_dist * (1 - lambda_value))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Execution\n","metadata":{}},{"cell_type":"code","source":"TRAIN_CSV = \"/kaggle/input/jaguar-re-id/train.csv\"\nTEST_CSV = \"/kaggle/input/jaguar-re-id/test.csv\"\nTRAIN_DIR = \"/kaggle/input/jaguar-re-id/train/train\"\nTEST_DIR = \"/kaggle/input/jaguar-re-id/test/test\"\n\ntrain_df = pd.read_csv(TRAIN_CSV)\ntest_df = pd.read_csv(TEST_CSV)\n\n# ============================================================\n# ✅ ADD: label_map + fixed val split + mAP functions + val_eval_loader\n# (does NOT change your training loader/model/optimizer logic)\n# ============================================================\nseed = getattr(Config, \"seed\", 42)\nval_frac = getattr(Config, \"val_frac_per_id\", 0.2)\n\nunique_ids = sorted(train_df[\"ground_truth\"].astype(str).unique())\nlabel_map = {gid: i for i, gid in enumerate(unique_ids)}\n\nrng = np.random.RandomState(seed)\nval_mask = np.zeros(len(train_df), dtype=bool)\nfor gid, sub in train_df.groupby(\"ground_truth\"):\n    idxs = sub.index.values\n    n_val = max(1, int(len(idxs) * val_frac))\n    chosen = rng.choice(idxs, size=n_val, replace=False)\n    val_mask[chosen] = True\n\nva_df = train_df[val_mask].reset_index(drop=True)\nprint(\"mAP eval split | val:\", len(va_df))\n\n# val eval dataset must return (img, filename) → use is_test=True (same as your test_loader)\nval_eval_loader = DataLoader(\n    JaguarDataset(pd.DataFrame({\"filename\": va_df[\"filename\"].values}), TRAIN_DIR, test_transform, True),\n    batch_size=Config.batch_size,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=False,\n)\n\n# filename -> label for mAP\nfname_to_label = dict(zip(va_df[\"filename\"].astype(str).values,\n                          va_df[\"ground_truth\"].astype(str).map(label_map).astype(int).values))\n\ndef _l2norm_np(x, eps=1e-12):\n    n = np.linalg.norm(x, axis=1, keepdims=True)\n    return x / np.clip(n, eps, None)\n\ndef _average_precision_from_sorted_rel(rel_sorted: np.ndarray):\n    npos = rel_sorted.sum()\n    if npos == 0:\n        return 0.0\n    cumsum = np.cumsum(rel_sorted)\n    precision_at_k = cumsum / (np.arange(len(rel_sorted)) + 1)\n    return float(precision_at_k[rel_sorted].sum() / npos)\n\ndef identity_balanced_map(emb: np.ndarray, labels: np.ndarray):\n    emb = _l2norm_np(emb)\n    sims = emb @ emb.T\n    N = sims.shape[0]\n\n    APs = np.zeros(N, dtype=np.float32)\n    for i in range(N):\n        sims[i, i] = -1e9\n        order = np.argsort(-sims[i])\n        rel = (labels[order] == labels[i])\n        APs[i] = _average_precision_from_sorted_rel(rel)\n\n    per_id = []\n    for y in np.unique(labels):\n        per_id.append(APs[labels == y].mean())\n    return float(np.mean(per_id))\n# ============================================================\n\n\ntrain_loader = DataLoader(\n    JaguarDataset(train_df, TRAIN_DIR, train_transform),\n    batch_size=Config.batch_size,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=False,\n)\nmodel = EVABoss().to(Config.device)\noptimizer = torch.optim.AdamW(\n    model.parameters(), lr=Config.lr, weight_decay=Config.weight_decay\n)\nscaler = torch.amp.GradScaler(Config.device_type)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, T_max=Config.num_epochs\n)\n\nprint(f\"🔥 Training EVA-02 Large (448px)...\")\n\nfor epoch in range(Config.num_epochs):\n    loss = train_epoch(model, train_loader, optimizer, nn.CrossEntropyLoss(label_smoothing=0.05), scaler)\n    scheduler.step()\n\n    # ============================================================\n    # ✅ ADD: compute id-mAP on the fixed val split\n    # ============================================================\n    val_emb, val_names = extract_features(model, val_eval_loader)  # uses your existing function\n    val_labels = np.array([fname_to_label[str(n)] for n in val_names], dtype=np.int64)\n    cv_map = identity_balanced_map(val_emb, val_labels)\n    # ============================================================\n\n    print(\n        f\"Epoch {epoch+1}/{Config.num_epochs} | Loss: {loss:.4f} | CV id-mAP: {cv_map:.4f} | LR: {scheduler.get_last_lr()[0]:.2e}\"\n    )\n\nunique_test = sorted(set(test_df[\"query_image\"]) | set(test_df[\"gallery_image\"]))\ntest_loader = DataLoader(\n    JaguarDataset(\n        pd.DataFrame({\"filename\": unique_test}), TEST_DIR, test_transform, True\n    ),\n    batch_size=Config.batch_size,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=False,\n)\n\nemb, names = extract_features(model, test_loader)\nimg_map = {n: i for i, n in enumerate(names)}\n\nif Config.use_qe:\n    emb = query_expansion(emb)\nsim_matrix = emb @ emb.T\nif Config.use_rerank:\n    sim_matrix = k_reciprocal_rerank(sim_matrix)\n\npreds = []\nfor _, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"Mapping\"):\n    s = sim_matrix[img_map[row[\"query_image\"]], img_map[row[\"gallery_image\"]]]\n    preds.append(max(0.0, min(1.0, s)))\n\nsub = pd.DataFrame({\"row_id\": test_df[\"row_id\"], \"similarity\": preds})\nsub.to_csv(\"submission.csv\", index=False)\nprint(f\"✅ Done! Mean Sim: {np.mean(preds):.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}