{"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}],"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,"execution":{"iopub.status.busy":"2026-02-27T00:36:34.772716Z","iopub.execute_input":"2026-02-27T00:36:34.773034Z","iopub.status.idle":"2026-02-27T00:36:34.780856Z","shell.execute_reply.started":"2026-02-27T00:36:34.773007Z","shell.execute_reply":"2026-02-27T00:36:34.779999Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 1 — Install\n# ============================================================\n!pip install -qU timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:36:36.771153Z","iopub.execute_input":"2026-02-27T00:36:36.771731Z","iopub.status.idle":"2026-02-27T00:36:43.573023Z","shell.execute_reply.started":"2026-02-27T00:36:36.771702Z","shell.execute_reply":"2026-02-27T00:36:43.572061Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 2 — Imports\n# ============================================================\nimport math, random\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom PIL import Image\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","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:36:58.223842Z","iopub.execute_input":"2026-02-27T00:36:58.224188Z","iopub.status.idle":"2026-02-27T00:37:10.726325Z","shell.execute_reply.started":"2026-02-27T00:36:58.224154Z","shell.execute_reply":"2026-02-27T00:37:10.725641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 3 — Config ✅ FULL REPLACE (mAP-focused)\n#   Changes:\n#     - More identities per batch at same batch size: P=8,K=3 (batch=24)\n#     - More optimizer updates per epoch: steps_mult=3\n#     - LR floor + slower decay controls: min_lr_* + lr_Tmax\n# ============================================================\n\nclass Config:\n    seed = 42\n\n    model_name = \"eva02_large_patch14_448.mim_m38m_ft_in22k_in1k\"\n    img_size = 448\n\n    num_epochs = 15\n\n    use_crop = False              # <-- toggle this\n    crop_pad_frac = 0.08         # used only if use_crop=True\n\n    # ✅ Cache behavior\n    cache_crops = True\n    cache_root = \"/kaggle/working/jaguar_cache\"  # NEW: base cache folder\n    cache_resize = True\n    cache_letterbox = False      # False=warp, True=letterbox\n\n    \n    lr_backbone = 2e-5\n    lr_head     = 1.2e-4\n    weight_decay = 5e-4\n\n    # ✅ LR schedule controls (cosine with floor, slower than epochs)\n    min_lr_backbone = 4e-6\n    min_lr_head     = 2.4e-5\n    lr_Tmax         = 15\n\n    # PK (batch stays 24)\n    P = 18\n    K = 2\n    grad_accum = 1\n \n    # ✅ more steps per epoch (more updates, better mAP)\n    steps_mult = 4\n\n    # ArcFace\n    arcface_s = 30.0\n    arcface_m = 0.50\n\n    # Triplet\n    triplet_margin = 0.40\n    triplet_weight = 0.3          # often better than 1.0\n    triplet_warmup_epochs = 0\n    triplet_ramp_epochs   = 1\n\n    # Split / eval\n    val_frac_per_id = 0.2\n    eval_every = 1\n\n    # Inference\n    use_tta = True\n    use_qe = True\n    qe_topk = 5\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    # Cache\n    cache_crops = True\n    cache_pad_frac = 0.08\n    cache_workers = 8\n\n    # Multi-GPU\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\n    # Memory\n    use_checkpointing = True\n\nprint(\"torch.cuda.device_count():\", torch.cuda.device_count())\n!nvidia-smi -L\nprint(\"img_size:\", Config.img_size, \"| batch:\", Config.P * Config.K, \"| steps_mult:\", Config.steps_mult)\nprint(\"LR backbone/head:\", Config.lr_backbone, Config.lr_head, \"| min LR:\", Config.min_lr_backbone, Config.min_lr_head, \"| T:\", Config.lr_Tmax)\nprint(\"triplet_weight:\", Config.triplet_weight, \"| warmup:\", Config.triplet_warmup_epochs, \"| ramp:\", Config.triplet_ramp_epochs)\nprint(\"checkpointing:\", Config.use_checkpointing)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:37:18.801130Z","iopub.execute_input":"2026-02-27T00:37:18.801450Z","iopub.status.idle":"2026-02-27T00:37:19.275615Z","shell.execute_reply.started":"2026-02-27T00:37:18.801423Z","shell.execute_reply":"2026-02-27T00:37:19.274241Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 4 — Seed + TF32 (speed)\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    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(Config.seed)\n\n# TF32 can speed up on Ampere+ (T4 supports tensor cores)\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\nprint(\"TF32 enabled (if supported).\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:37:20.735205Z","iopub.execute_input":"2026-02-27T00:37:20.735797Z","iopub.status.idle":"2026-02-27T00:37:20.750080Z","shell.execute_reply.started":"2026-02-27T00:37:20.735762Z","shell.execute_reply":"2026-02-27T00:37:20.749177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 5 — Preprocess helpers (CROP or NO-CROP controlled by Config.use_crop)\n#   - crop mode: alpha bbox crop + gray composite\n#   - no-crop  : full-frame gray composite (no bbox crop)\n#   Output is ALWAYS RGB\n# ============================================================\n\nimport numpy as np\nfrom PIL import Image\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, mode=\"RGB\")\n\ndef crop_and_neutralize(img: Image.Image, pad_frac=0.08, bg=128):\n    \"\"\"CROP MODE: alpha bbox crop + padding, composite to gray.\"\"\"\n    rgba = img.convert(\"RGBA\") if img.mode != \"RGBA\" else img\n    bbox, alpha = alpha_bbox_from_rgba(rgba)\n\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\n    return comp.crop((x0p, y0p, x1p, y1p)).convert(\"RGB\")\n\ndef no_crop_rgb(img: Image.Image, bg=128):\n    \"\"\"NO-CROP MODE: full-frame, composite RGBA->RGB if needed.\"\"\"\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    \"\"\"Single entry point controlled by Config.use_crop.\"\"\"\n    if getattr(Config, \"use_crop\", True):\n        return crop_and_neutralize(img, pad_frac=Config.crop_pad_frac, bg=128)\n    else:\n        return no_crop_rgb(img, bg=128)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:37:31.950749Z","iopub.execute_input":"2026-02-27T00:37:31.951059Z","iopub.status.idle":"2026-02-27T00:37:31.963760Z","shell.execute_reply.started":"2026-02-27T00:37:31.951018Z","shell.execute_reply":"2026-02-27T00:37:31.962896Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 6 — Cache builder (CROP/NO-CROP switch + mode-specific cache folders)\n#   - Writes RGB 448x448 cached images\n#   - If Config.cache_letterbox=True => keep aspect ratio + pad\n#     else => warp to (img_size,img_size)\n# ============================================================\n\nimport os, shutil\nfrom pathlib import Path\nimport numpy as np\nimport cv2\nimport multiprocessing as mp\nfrom concurrent.futures import ProcessPoolExecutor\nfrom tqdm import tqdm\n\n# ---- knobs\nCACHE_SIDE = int(getattr(Config, \"img_size\", 448))\nPNG_COMP   = 6\nCACHE_BG   = 128\n\ndef _atomic_imwrite(out_path, img_bgr):\n    out_path = Path(out_path)\n    out_path.parent.mkdir(parents=True, exist_ok=True)\n\n    ext = out_path.suffix.lower()\n    if ext == \".png\":\n        enc_ext = \".png\"\n        params = [int(cv2.IMWRITE_PNG_COMPRESSION), int(PNG_COMP)]\n    elif ext in [\".jpg\", \".jpeg\"]:\n        enc_ext = \".jpg\"\n        params = [int(cv2.IMWRITE_JPEG_QUALITY), 92]\n    else:\n        enc_ext = \".png\"\n        params = [int(cv2.IMWRITE_PNG_COMPRESSION), int(PNG_COMP)]\n        out_path = out_path.with_suffix(\".png\")\n\n    ok, buf = cv2.imencode(enc_ext, img_bgr, params)\n    if not ok:\n        return False\n\n    tmp = out_path.with_suffix(out_path.suffix + \".tmp\")\n    try:\n        with open(tmp, \"wb\") as f:\n            f.write(buf.tobytes())\n        os.replace(tmp, out_path)\n        return True\n    except Exception:\n        try:\n            if tmp.exists():\n                tmp.unlink()\n        except Exception:\n            pass\n        return False\n\ndef _letterbox_to_square_bgr(bgr, side=448, fill=128):\n    h, w = bgr.shape[:2]\n    if h <= 0 or w <= 0:\n        return np.full((side, side, 3), fill, np.uint8)\n\n    scale = min(side / w, side / h)\n    new_w = max(1, int(round(w * scale)))\n    new_h = max(1, int(round(h * scale)))\n\n    if (new_w, new_h) != (w, h):\n        bgr = cv2.resize(bgr, (new_w, new_h), interpolation=cv2.INTER_AREA)\n\n    canvas = np.full((side, side, 3), fill, np.uint8)\n    x0 = (side - new_w) // 2\n    y0 = (side - new_h) // 2\n    canvas[y0:y0+new_h, x0:x0+new_w] = bgr\n    return canvas\n\ndef _cache_one_cv2(args):\n    fn, src_dir, cache_dir, pad_frac = args\n    in_path  = os.path.join(src_dir, fn)\n    out_path = os.path.join(cache_dir, fn)\n\n    if os.path.exists(out_path):\n        return 0\n\n    img = cv2.imread(in_path, 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\", True))\n\n    # ---- If alpha exists, either crop bbox (crop-mode) or keep full frame (no-crop)\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        # composite alpha onto gray (now either cropped or full frame)\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    # ---- Resize to CACHE_SIDE\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    ok = _atomic_imwrite(out_path, comp)\n    return 1 if ok else 0\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        print(\"   Delete cache folder then rerun.\")\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-02-27T00:37:32.143094Z","iopub.execute_input":"2026-02-27T00:37:32.143438Z","iopub.status.idle":"2026-02-27T00:37:32.533401Z","shell.execute_reply.started":"2026-02-27T00:37:32.143410Z","shell.execute_reply":"2026-02-27T00:37:32.532598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 7 — Transforms (skip Resize cost if cache already 448x448)\n# ============================================================\n\nimport torchvision.transforms as transforms\nfrom PIL import Image\n\nclass ResizeIfNeeded:\n    def __init__(self, size):\n        self.size = int(size)\n\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\ntrain_transform = transforms.Compose([\n    ResizeIfNeeded(Config.img_size),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomAffine(degrees=10, translate=(0.08, 0.08), scale=(0.9, 1.1), shear=5),\n    transforms.ColorJitter(brightness=0.15, contrast=0.15, saturation=0.10, hue=0.02),\n    transforms.ToTensor(),\n    transforms.Normalize([0.481, 0.457, 0.408], [0.268, 0.261, 0.275]),\n    transforms.RandomErasing(p=0.10, scale=(0.02, 0.08), ratio=(0.3, 3.3)),\n])\n\ntest_transform = transforms.Compose([\n    ResizeIfNeeded(Config.img_size),\n    transforms.ToTensor(),\n    transforms.Normalize([0.481, 0.457, 0.408], [0.268, 0.261, 0.275]),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:37:32.534771Z","iopub.execute_input":"2026-02-27T00:37:32.535036Z","iopub.status.idle":"2026-02-27T00:37:32.543241Z","shell.execute_reply.started":"2026-02-27T00:37:32.535011Z","shell.execute_reply":"2026-02-27T00:37:32.542406Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 8 — Dataset (cache-first + crop/no-crop fallback via Config.use_crop)\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        # cache first\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        # raw fallback => preprocess controlled by Config.use_crop\n        p = self.img_dir / fname\n        img = Image.open(p)\n        img = preprocess_pil(img)  # <-- crop or no-crop decided by Config.use_crop\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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:37:32.544120Z","iopub.execute_input":"2026-02-27T00:37:32.544409Z","iopub.status.idle":"2026-02-27T00:37:32.560091Z","shell.execute_reply.started":"2026-02-27T00:37:32.544371Z","shell.execute_reply":"2026-02-27T00:37:32.559307Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 9 — PK Batch Sampler\n# ============================================================\n\nclass PKBatchSampler(Sampler):\n    def __init__(self, labels, P, K, steps_per_epoch, seed=42):\n        self.labels = np.array(labels)\n        self.P = P\n        self.K = K\n        self.steps_per_epoch = steps_per_epoch\n        self.rng = np.random.RandomState(seed)\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\n        self.unique_labels = np.array(sorted(list(self.label_to_indices.keys())))\n\n    def __len__(self):\n        return self.steps_per_epoch\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                idxs = self.label_to_indices[int(y)]\n                if len(idxs) >= self.K:\n                    pick = self.rng.choice(idxs, size=self.K, replace=False)\n                else:\n                    pick = self.rng.choice(idxs, size=self.K, replace=True)\n                batch.extend(pick.tolist())\n            yield batch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:37:32.573863Z","iopub.execute_input":"2026-02-27T00:37:32.574173Z","iopub.status.idle":"2026-02-27T00:37:32.582194Z","shell.execute_reply.started":"2026-02-27T00:37:32.574143Z","shell.execute_reply":"2026-02-27T00:37:32.581439Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 10 — ArcFace (proper cos(theta+m))\n# ============================================================\n\nclass ArcFace(nn.Module):\n    def __init__(self, in_features, out_features, s=30.0, m=0.50, easy_margin=False):\n        super().__init__()\n        self.s = s\n        self.m = m\n        self.easy_margin = easy_margin\n\n        self.weight = nn.Parameter(torch.empty(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.th = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\n\n    def forward(self, x, label=None):\n        cosine = F.linear(F.normalize(x), F.normalize(self.weight)).clamp(-1+1e-7, 1-1e-7)\n        if label is None:\n            return cosine * self.s\n\n        sine = torch.sqrt(1.0 - cosine * cosine)\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.view(-1, 1), 1.0)\n\n        logits = one_hot * phi + (1.0 - one_hot) * cosine\n        return logits * self.s\n\n\n\nclass SubCenterArcFace(nn.Module):\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 = 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_norm = F.normalize(x, dim=1)\n\n        w = self.weight.reshape(-1, self.weight.size(-1))          # [C*k, D]\n        w_norm = F.normalize(w, dim=1)\n\n        cosine = F.linear(x_norm, w_norm)                          # [B, C*k]\n        cosine = cosine.reshape(-1, self.weight.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        sine = torch.sqrt(1.0 - cosine * cosine)\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.view(-1, 1), 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-02-27T00:37:32.802140Z","iopub.execute_input":"2026-02-27T00:37:32.802444Z","iopub.status.idle":"2026-02-27T00:37:32.817556Z","shell.execute_reply.started":"2026-02-27T00:37:32.802419Z","shell.execute_reply":"2026-02-27T00:37:32.816806Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# REPLACE CELL 11 — Model (EVA + GeM + BNNeck + ArcFace)\n#   - Triplet uses pre-BN features\n#   - ArcFace uses post-BN features\n#   - Inference returns post-BN embedding\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) * p)\n        self.eps = eps\n\n    def forward(self, x):\n        return F.avg_pool2d(x.clamp(min=self.eps).pow(self.p), (x.size(-2), x.size(-1))).pow(1.0 / self.p)\n\nclass EVABoss(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        self.backbone = timm.create_model(Config.model_name, pretrained=True, num_classes=0)\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 = self.backbone.num_features\n        self.gem = GeM()\n\n        # ✅ BNNeck\n        self.bnneck = nn.BatchNorm1d(self.feat_dim)\n        # common trick: don't train BN bias\n        self.bnneck.bias.requires_grad_(False)\n\n        self.head = SubCenterArcFace(self.feat_dim, num_classes, k=3, s=Config.arcface_s, m=Config.arcface_m)\n\n    def forward_features_map(self, x):\n        feat = self.backbone.forward_features(x)\n\n        if feat.dim() == 2:\n            return feat.unsqueeze(-1).unsqueeze(-1)\n\n        if feat.dim() == 3:\n            B, N, C = feat.shape\n            num_prefix = getattr(self.backbone, \"num_prefix_tokens\", 1)\n            if N > num_prefix:\n                feat = feat[:, num_prefix:, :]\n                N = feat.shape[1]\n            H = W = int(math.sqrt(N))\n            if H * W != N:\n                W = int(math.sqrt(N))\n                H = N // W\n            return feat.transpose(1, 2).reshape(B, C, H, W)\n\n        return feat\n\n    def embed_pair(self, x):\n        fmap = self.forward_features_map(x)\n        feat = self.gem(fmap).flatten(1)         # pre-BN feature (for triplet)\n        feat_bn = self.bnneck(feat)              # post-BN feature (for ArcFace/infer)\n        return feat, feat_bn\n\n    def forward(self, x, label=None, return_emb=False):\n        feat_tri, feat_cls = self.embed_pair(x)\n\n        # inference embedding: use BNNeck output (standard BoT)\n        if label is None:\n            return feat_cls\n\n        logits = self.head(feat_cls, label)\n\n        if return_emb:\n            # return BOTH so you can pick which one for triplet\n            return feat_tri, feat_cls, logits\n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:37:33.198876Z","iopub.execute_input":"2026-02-27T00:37:33.199229Z","iopub.status.idle":"2026-02-27T00:37:33.213010Z","shell.execute_reply.started":"2026-02-27T00:37:33.199200Z","shell.execute_reply":"2026-02-27T00:37:33.212275Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 12 — 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\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))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:37:33.678939Z","iopub.execute_input":"2026-02-27T00:37:33.679832Z","iopub.status.idle":"2026-02-27T00:37:33.686521Z","shell.execute_reply.started":"2026-02-27T00:37:33.679800Z","shell.execute_reply":"2026-02-27T00:37:33.685609Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 13 — TripletLoss + train_epoch ✅ FULL REPLACE\n#   ✅ Logs cls/tri separately\n#   ✅ Triplet weight warmup/ramp by epoch\n# ============================================================\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom tqdm import tqdm\n\nclass TripletLoss(nn.Module):\n    def __init__(self, margin=0.3):\n        super().__init__()\n        self.margin = margin\n        self.ranking_loss = nn.MarginRankingLoss(margin=margin)\n\n    def forward(self, embeddings, targets):\n        n = embeddings.size(0)\n        if n <= 1:\n            return embeddings.new_tensor(0.0)\n\n        # pairwise euclidean distances\n        dist = torch.cdist(embeddings, embeddings, p=2)  # (N,N)\n        mask = targets.view(-1, 1).eq(targets.view(1, -1))  # positives mask\n\n        dist_ap, dist_an = [], []\n        for i in range(n):\n            pos = dist[i][mask[i]]      # includes self (0)\n            neg = dist[i][~mask[i]]\n\n            # need at least one other positive + at least one negative\n            if pos.numel() <= 1 or neg.numel() == 0:\n                continue\n\n            dist_ap.append(pos.max().unsqueeze(0))  # hardest positive\n            dist_an.append(neg.min().unsqueeze(0))  # hardest negative\n\n        if len(dist_ap) == 0:\n            return embeddings.new_tensor(0.0)\n\n        dist_ap = torch.cat(dist_ap)\n        dist_an = torch.cat(dist_an)\n        y = torch.ones_like(dist_an)\n        return self.ranking_loss(dist_an, dist_ap, y)\n\ncriterion_cls = nn.CrossEntropyLoss(label_smoothing=0.05)\ncriterion_tri = TripletLoss(margin=Config.triplet_margin)\n\ndef _triplet_weight_for_epoch(epoch_idx: int):\n    # epoch_idx is 0-based\n    warm = int(getattr(Config, \"triplet_warmup_epochs\", 0))\n    ramp = int(getattr(Config, \"triplet_ramp_epochs\", 0))\n    base = float(getattr(Config, \"triplet_weight\", 1.0))\n\n    if epoch_idx < warm:\n        return 0.0\n    if ramp <= 0:\n        return base\n\n    # ramp from 0 -> base over `ramp` epochs after warmup\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    tw = _triplet_weight_for_epoch(epoch_idx)\n\n    total_meter = 0.0\n    cls_meter   = 0.0\n    tri_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        with torch.amp.autocast(device_type=Config.device_type):\n            # ✅ BNNeck model returns 3 outputs\n            feat_tri, feat_cls, logits = model(imgs, labels, return_emb=True)\n\n            loss_cls = criterion_cls(logits, labels)\n\n            # ✅ Triplet on pre-BN features\n            feat_tri_n = F.normalize(feat_tri.float(), dim=1)\n            loss_tri = criterion_tri(feat_tri_n, labels)\n\n            total_loss = loss_cls + tw * loss_tri\n\n        scaler.scale(total_loss / Config.grad_accum).backward()\n\n        if (step + 1) % Config.grad_accum == 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        tri_meter   += float(loss_tri.item())\n\n    if len(loader) % Config.grad_accum != 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), (tri_meter / denom), tw\n\n\n@torch.no_grad()\ndef extract_embeddings(model, loader):\n    model.eval()\n    feats, names = [], []\n\n    for imgs, fnames in tqdm(loader, desc=\"Extract\"):\n        imgs = imgs.to(Config.device, non_blocking=True).contiguous()\n\n        # model(imgs) returns embeddings when label=None\n        f1 = model(imgs)\n\n        if getattr(Config, \"use_tta\", False):\n            f2 = model(torch.flip(imgs, dims=[3]))\n            f1 = 0.5 * (f1 + f2)\n\n        f1 = F.normalize(f1, dim=1).cpu().numpy()\n        feats.append(f1)\n        names.extend(list(fnames))\n\n    return np.concatenate(feats, axis=0), names","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:37:34.318929Z","iopub.execute_input":"2026-02-27T00:37:34.319670Z","iopub.status.idle":"2026-02-27T00:37:34.339083Z","shell.execute_reply.started":"2026-02-27T00:37:34.319636Z","shell.execute_reply":"2026-02-27T00:37:34.338215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 14A — Clean Split (Near-Duplicates) ✅ FULL CELL\n#   - Computes pHash for ALL train images\n#   - Builds near-duplicate \"session\" groups (connected components)\n#   - Splits PER ID using whole groups (no leakage)\n#   - Guarantees every ID keeps at least 1 group in TRAIN\n#   Output:\n#     - train_df with group_id\n#     - tr_df, va_df (use these in CELL 14B)\n# ============================================================\n\n!pip -q install ImageHash scipy\n\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom PIL import Image\nfrom multiprocessing import Pool, cpu_count\nfrom tqdm.auto import tqdm\nimport imagehash\n\nfrom scipy.spatial.distance import cdist\nfrom scipy.sparse import csr_matrix\nfrom scipy.sparse.csgraph import connected_components\n\n# ---- paths (must match your dataset)\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\n# ---- knobs\nPHASH_THRESH = 11  # from your earlier safe threshold\nVAL_FRAC_PER_ID = float(getattr(Config, \"val_frac_per_id\", 0.2))\nRNG_SEED = int(getattr(Config, \"seed\", 42))\nNPROC = min(8, cpu_count())\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 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\nN = len(filenames)\nprint(\"Valid hashed images:\", N)\n\nprint(\"Computing Hamming distance matrix...\")\ndist = cdist(H, H, metric=\"hamming\") * H.shape[1]  # -> [0..64]\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\n# filename -> group_id\nimage_to_group = {}\ndup_groups = 0\n\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\n# uniques as their own group\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)\n\nprint(\"Duplicate groups found:\", dup_groups)\n\n# Per-identity split by group (no leakage)\nrng = np.random.RandomState(RNG_SEED)\nval_mask = np.zeros(len(train_df), dtype=bool)\n\nfor gt, sub in train_df.groupby(\"ground_truth\"):\n    groups = sub[\"group_id\"].unique().tolist()\n    rng.shuffle(groups)\n\n    # if only one group, keep all in TRAIN (can't split without leakage)\n    if len(groups) == 1:\n        continue\n\n    idxs = sub.index.values\n    target = max(1, int(len(idxs) * VAL_FRAC_PER_ID))\n\n    chosen = []\n    count = 0\n\n    # choose groups for val but leave at least 1 group for train\n    for g in groups[:-1]:\n        g_count = int((sub[\"group_id\"] == g).sum())\n        chosen.append(g)\n        count += g_count\n        if count >= target:\n            break\n\n    val_mask[sub.index[sub[\"group_id\"].isin(chosen)]] = True\n\ntr_df = train_df[~val_mask].reset_index(drop=True)\nva_df = train_df[val_mask].reset_index(drop=True)\n\nprint(\"train:\", len(tr_df), \"val:\", len(va_df))\n\n# Leakage check\nleak = set(tr_df[\"group_id\"]).intersection(set(va_df[\"group_id\"]))\nprint(\"Leaked groups (should be 0):\", len(leak))\n\nprint(f\"IDs in VAL: {va_df['ground_truth'].nunique()}/{train_df['ground_truth'].nunique()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:37:34.798910Z","iopub.execute_input":"2026-02-27T00:37:34.799705Z","iopub.status.idle":"2026-02-27T00:42:52.949701Z","shell.execute_reply.started":"2026-02-27T00:37:34.799671Z","shell.execute_reply":"2026-02-27T00:42:52.948923Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# CELL 14B — Build caches + loaders + train (2×T4 DP) ✅ FULL CELL\n#   - Uses tr_df / va_df from CELL 14A (clean split)\n#   - Uses mode-specific cache folder based on Config.use_crop (crop vs nocrop)\n#   - PK sampler + steps_mult (CEIL)\n#   - Builds val_eval_loader ONCE\n#   - Cosine LR with floor (LambdaLR) for both param groups\n#   - Saves best.pt\n# ============================================================\n\nimport gc, math\nfrom pathlib import Path\nfrom torch.utils.data import DataLoader\n\n# ----------------------------\n# HARD RESET GPU OBJECTS\n# ----------------------------\nfor name in [\n    \"model\",\"base_model\",\"optimizer\",\"scheduler\",\"scaler\",\n    \"train_loader\",\"val_eval_loader\",\"train_ds\",\"val_ds\",\"pk_sampler\",\"val_eval_ds\",\n]:\n    if name in globals():\n        try:\n            del globals()[name]\n        except Exception:\n            pass\ngc.collect()\ntry:\n    torch.cuda.empty_cache()\nexcept Exception:\n    pass\n\n# ----------------------------\n# Mode-specific cache dirs (IMPORTANT)\n# ----------------------------\nmode = \"crop\" if getattr(Config, \"use_crop\", True) else \"nocrop\"\nCACHE_ROOT = Path(getattr(Config, \"cache_root\", \"/kaggle/working/jaguar_cache\")) / mode\nTRAIN_CACHE = CACHE_ROOT / \"train\"\nTEST_CACHE  = CACHE_ROOT / \"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\n# ----------------------------\n# Build caches (if enabled)\n#   - pad_frac only matters when use_crop=True\n# ----------------------------\nif getattr(Config, \"cache_crops\", False):\n    pad = float(getattr(Config, \"crop_pad_frac\", 0.08)) if getattr(Config, \"use_crop\", True) else 0.0\n\n    build_crop_cache_fast(\n        train_df[\"filename\"].values,\n        TRAIN_DIR,\n        TRAIN_CACHE,\n        pad_frac=pad,\n        max_workers=int(getattr(Config, \"cache_workers\", 8)),\n    )\n\n    unique_test = sorted(set(test_df[\"query_image\"]) | set(test_df[\"gallery_image\"]))\n    build_crop_cache_fast(\n        unique_test,\n        TEST_DIR,\n        TEST_CACHE,\n        pad_frac=pad,\n        max_workers=int(getattr(Config, \"cache_workers\", 8)),\n    )\n\n# ----------------------------\n# Datasets (use clean split tr_df / va_df)\n# ----------------------------\ntrain_ds = JaguarDataset(\n    tr_df, TRAIN_DIR, transform=train_transform, is_test=False,\n    label_map=label_map,\n    cache_dir=(TRAIN_CACHE if getattr(Config, \"cache_crops\", False) else None),\n)\nval_ds = JaguarDataset(\n    va_df, TRAIN_DIR, transform=test_transform, is_test=False,\n    label_map=label_map,\n    cache_dir=(TRAIN_CACHE if getattr(Config, \"cache_crops\", False) else None),\n)\n\ntrain_labels = train_ds.df[\"label\"].values\nval_labels   = val_ds.df[\"label\"].values  # aligned to va_df order (since val_ds built from va_df)\n\n# ----------------------------\n# Steps per epoch (CEIL) + PK sampler\n# ----------------------------\nbatch_size = int(Config.P) * int(Config.K)\nbase_steps = int(math.ceil(len(tr_df) / float(batch_size)))\nsteps_mult = int(getattr(Config, \"steps_mult\", 1))\nsteps_per_epoch = max(1, base_steps * steps_mult)\n\nprint(\"steps_per_epoch:\", steps_per_epoch, \"| batch:\", batch_size, \"| steps_mult:\", steps_mult)\n\npk_sampler = PKBatchSampler(\n    train_labels,\n    P=int(Config.P),\n    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=pk_sampler,\n    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    prefetch_factor=int(getattr(Config, \"prefetch_factor\", 2)) if int(getattr(Config, \"num_workers\", 4)) > 0 else None,\n)\n\n# val eval loader (returns filenames) — build ONCE\nval_names_df = va_df[[\"filename\"]].copy()\nval_eval_ds = JaguarDataset(\n    val_names_df, TRAIN_DIR, transform=test_transform, is_test=True,\n    cache_dir=(TRAIN_CACHE if getattr(Config, \"cache_crops\", False) else None),\n)\nval_eval_loader = DataLoader(\n    val_eval_ds,\n    batch_size=16,\n    shuffle=False,\n    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    prefetch_factor=int(getattr(Config, \"prefetch_factor\", 2)) if int(getattr(Config, \"num_workers\", 4)) > 0 else None,\n)\n\n# ----------------------------\n# Model + DataParallel\n# ----------------------------\nbase_model = EVABoss(num_classes=num_classes).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# Optimizer param groups (backbone vs head-ish)\n#   Works whether your model has neck OR bnneck\n# ----------------------------\ndef get_param_groups(model):\n    m = model.module if isinstance(model, nn.DataParallel) else model\n\n    head_parts = []\n    if hasattr(m, \"gem\"):    head_parts += list(m.gem.parameters())\n    if hasattr(m, \"neck\"):   head_parts += list(m.neck.parameters())\n    if hasattr(m, \"bnneck\"): head_parts += list(m.bnneck.parameters())\n    if hasattr(m, \"head\"):   head_parts += list(m.head.parameters())\n\n    return [\n        {\"params\": m.backbone.parameters(), \"lr\": float(Config.lr_backbone)},\n        {\"params\": head_parts, \"lr\": float(Config.lr_head)},\n    ]\n\noptimizer = torch.optim.AdamW(get_param_groups(model), weight_decay=float(Config.weight_decay))\n\n# ----------------------------\n# Cosine LR with floor (per param group)\n# ----------------------------\ndef cosine_floor(epoch, base_lr, min_lr, T):\n    t = min(epoch, T)\n    cos = 0.5 * (1.0 + math.cos(math.pi * t / T))\n    min_factor = float(min_lr) / float(base_lr)\n    return min_factor + (1.0 - min_factor) * cos\n\nscheduler = torch.optim.lr_scheduler.LambdaLR(\n    optimizer,\n    lr_lambda=[\n        lambda e: cosine_floor(e, Config.lr_backbone, Config.min_lr_backbone, Config.lr_Tmax),\n        lambda e: cosine_floor(e, Config.lr_head,     Config.min_lr_head,     Config.lr_Tmax),\n    ],\n)\n\nscaler = torch.amp.GradScaler(enabled=(Config.device_type == \"cuda\"))\n\nprint(\"🔥 Training (ArcFace + Triplet, PK + cache + 2GPU DataParallel)...\")\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    loss, loss_cls, loss_tri, tw = train_epoch(model, train_loader, optimizer, scaler, epoch)\n    scheduler.step()\n\n    if (epoch + 1) % int(getattr(Config, \"eval_every\", 1)) == 0:\n        val_emb, _ = extract_embeddings(model, val_eval_loader)\n        cv_map = identity_balanced_map(val_emb, val_labels)\n    else:\n        cv_map = float(\"nan\")\n\n    lrs = scheduler.get_last_lr()\n    print(\n        f\"Epoch {epoch+1}/{Config.num_epochs} | \"\n        f\"Loss {loss:.4f} (cls {loss_cls:.4f}, tri {loss_tri:.4f}, tw {tw:.2f}) | \"\n        f\"CV id-mAP {cv_map:.4f} | \"\n        f\"LR(backbone) {lrs[0]:.2e} | LR(head) {lrs[1]:.2e}\"\n    )\n\n    if np.isfinite(cv_map) and cv_map > best_map:\n        best_map = cv_map\n        if isinstance(model, nn.DataParallel):\n            torch.save(model.module.state_dict(), \"best.pt\")\n        else:\n            torch.save(model.state_dict(), \"best.pt\")\n\n    torch.cuda.empty_cache()\n2\nprint(\"Best CV id-mAP:\", best_map)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T00:43:10.819424Z","iopub.execute_input":"2026-02-27T00:43:10.820174Z","iopub.status.idle":"2026-02-27T03:06:53.580196Z","shell.execute_reply.started":"2026-02-27T00:43:10.820140Z","shell.execute_reply":"2026-02-27T03:06:53.577754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# FAIR steps: match old ~189 steps/epoch\n# ============================================================\n\nimport math\nTARGET = 189\nbase = math.ceil(len(train_df) / (Config.P * Config.K))\nConfig.steps_mult = max(1, int(round(TARGET / base)))\nprint(\"base:\", base, \"| new steps_mult:\", Config.steps_mult, \"| steps/epoch:\", base * Config.steps_mult)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# add Codeadd Markdown\n# ============================================================\n# QE k-selection on VAL (pick best k, then set Config.qe_topk)\n# ============================================================\n\nimport numpy as np\nimport torch\nfrom torch.utils.data import DataLoader\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    emb = emb.astype(np.float32, copy=False)\n    N = emb.shape[0]\n    sims = emb @ emb.T\n    if exclude_self:\n        np.fill_diagonal(sims, -1e9)\n    idx = np.argsort(-sims, axis=1)[:, :topk]\n    neigh = emb[idx].mean(axis=1)\n    out = alpha * emb + (1.0 - alpha) * neigh\n    return l2norm(out)\n\n# ---- load best\nstate = torch.load(\"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# ---- build val eval loader (img, filename)\nval_names_df = va_df[[\"filename\"]].copy()\nval_eval_ds = JaguarDataset(\n    val_names_df, TRAIN_DIR, transform=test_transform, is_test=True,\n    cache_dir=(TRAIN_CACHE if Config.cache_crops else None),\n)\nval_eval_loader = DataLoader(\n    val_eval_ds, batch_size=16, shuffle=False,\n    num_workers=Config.num_workers, pin_memory=Config.pin_memory,\n    persistent_workers=Config.persistent_workers, prefetch_factor=Config.prefetch_factor\n)\n\n# ---- labels aligned to va_df order\nname2label = dict(zip(val_ds.df[\"filename\"].values, val_ds.df[\"label\"].values))\nval_labels = np.array([name2label[f] for f in va_df[\"filename\"].values], dtype=np.int64)\n\n# ---- embeddings\nval_emb, _ = extract_embeddings(model, val_eval_loader)\nval_emb = l2norm(val_emb.astype(np.float32))\n\nbase = identity_balanced_map(val_emb.copy(), val_labels.copy())\nprint(\"RAW CV id-mAP:\", base)\n\n# sweep k (keep it small; big k often hurts)\nKs = [0, 1, 2, 3, 5, 7, 10]\nalpha = 0.5\n\nbest_k, best_map = 0, base\nfor k in Ks:\n    if k == 0:\n        m = base\n    else:\n        emb_qe = qe_mix_topk(val_emb, topk=k, alpha=alpha, exclude_self=True)\n        m = identity_balanced_map(emb_qe, val_labels)\n    print(f\"k={k:2d} -> CV id-mAP {m:.4f} | delta {m-base:+.4f}\")\n    if m > best_map:\n        best_map = m\n        best_k = k\n\nprint(\"\\n✅ Best k =\", best_k, \"best CV =\", best_map, \"delta =\", best_map-base)\nConfig.qe_topk = int(best_k)\nprint(\"Set Config.qe_topk =\", Config.qe_topk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-27T03:06:59.811263Z","iopub.execute_input":"2026-02-27T03:06:59.811623Z","iopub.status.idle":"2026-02-27T03:09:04.831377Z","shell.execute_reply.started":"2026-02-27T03:06:59.811588Z","shell.execute_reply":"2026-02-27T03:09:04.830362Z"}},"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(\"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,"execution":{"iopub.status.busy":"2026-02-27T03:09:25.968723Z","iopub.execute_input":"2026-02-27T03:09:25.969536Z","iopub.status.idle":"2026-02-27T03:11:17.418985Z","shell.execute_reply.started":"2026-02-27T03:09:25.969501Z","shell.execute_reply":"2026-02-27T03:11:17.417990Z"}},"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}]}