{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":126777,"databundleVersionId":15314950}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nfrom pathlib import Path\n\nfor root, dirs, files in os.walk(\"/kaggle/input\"):\n    print(root)\n    if len(files) > 0:\n        print(\"  sample files:\", files[:5])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:13:34.527149Z","iopub.execute_input":"2026-03-10T23:13:34.527891Z","iopub.status.idle":"2026-03-10T23:13:37.553821Z","shell.execute_reply.started":"2026-03-10T23:13:34.527863Z","shell.execute_reply":"2026-03-10T23:13:37.553034Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport math\nimport time\nimport random\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom pathlib import Path\nfrom collections import Counter, defaultdict\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics.pairwise import cosine_similarity\n\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nwarnings.filterwarnings(\"ignore\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:13:37.555119Z","iopub.execute_input":"2026-03-10T23:13:37.555413Z","iopub.status.idle":"2026-03-10T23:13:51.008341Z","shell.execute_reply.started":"2026-03-10T23:13:37.555383Z","shell.execute_reply":"2026-03-10T23:13:51.007591Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    seed = 42\n    \n    # paths\n    DATA_DIR = Path(\"/kaggle/input/competitions/jaguar-re-id\")   # change this\n    TRAIN_DIR = DATA_DIR / \"train\" / \"train\"\n    TEST_DIR = DATA_DIR / \"test\" / \"test\"\n    \n    TRAIN_CSV = DATA_DIR / \"train.csv\"\n    TEST_CSV = DATA_DIR / \"test.csv\"\n    SAMPLE_SUB = DATA_DIR / \"sample_submission.csv\"\n    \n    # columns (change if needed)\n    image_col = \"filename\"\n    target_col = \"ground_truth\"\n    \n    # image + model\n    image_size = 384\n    backbone = \"convnextv2_base.fcmae_ft_in22k_in1k\"\n    embedding_dim = 512\n    dropout = 0.1\n    \n    # training\n    folds = 5\n    epochs = 8\n    batch_size = 16\n    num_workers = 4\n    lr = 1e-4\n    weight_decay = 1e-4\n    label_smoothing = 0.0\n    \n    # arcface\n    margin_s = 30.0\n    margin_m = 0.35\n    \n    # TTA\n    use_tta = True\n    \n    # mask usage\n    use_alpha_as_mask = True   # if PNG has alpha channel\n    \n    # hardware\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    amp = True\n    \n    # save\n    model_dir = Path(\"./models\")\n    \nCFG.model_dir.mkdir(exist_ok=True, parents=True)\nprint(CFG.device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:13:51.009222Z","iopub.execute_input":"2026-03-10T23:13:51.009645Z","iopub.status.idle":"2026-03-10T23:13:51.272719Z","shell.execute_reply.started":"2026-03-10T23:13:51.009624Z","shell.execute_reply":"2026-03-10T23:13:51.271965Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Seed everything","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(CFG.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:13:51.274597Z","iopub.execute_input":"2026-03-10T23:13:51.274833Z","iopub.status.idle":"2026-03-10T23:13:51.354560Z","shell.execute_reply.started":"2026-03-10T23:13:51.274813Z","shell.execute_reply":"2026-03-10T23:13:51.353706Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load data","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(CFG.TRAIN_CSV)\ntest_df = pd.read_csv(CFG.TEST_CSV)\nsample_sub = pd.read_csv(CFG.SAMPLE_SUB)\n\nprint(train_df.head())\nprint(test_df.head())\nprint(sample_sub.head())\nprint(train_df.columns.tolist())\nprint(test_df.columns.tolist())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:13:51.355451Z","iopub.execute_input":"2026-03-10T23:13:51.355958Z","iopub.status.idle":"2026-03-10T23:13:51.546893Z","shell.execute_reply.started":"2026-03-10T23:13:51.355929Z","shell.execute_reply":"2026-03-10T23:13:51.546212Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_df.columns.tolist())\nprint(train_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:13:51.547704Z","iopub.execute_input":"2026-03-10T23:13:51.547941Z","iopub.status.idle":"2026-03-10T23:13:51.554141Z","shell.execute_reply.started":"2026-03-10T23:13:51.547913Z","shell.execute_reply":"2026-03-10T23:13:51.553327Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Label encoding","metadata":{}},{"cell_type":"code","source":"label2id = {k: v for v, k in enumerate(sorted(train_df[CFG.target_col].unique()))}\nid2label = {v: k for k, v in label2id.items()}\n\ntrain_df[\"label\"] = train_df[CFG.target_col].map(label2id)\nnum_classes = train_df[\"label\"].nunique()\n\nprint(\"num_classes =\", num_classes)\nprint(train_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:13:51.555133Z","iopub.execute_input":"2026-03-10T23:13:51.555409Z","iopub.status.idle":"2026-03-10T23:13:51.575275Z","shell.execute_reply.started":"2026-03-10T23:13:51.555378Z","shell.execute_reply":"2026-03-10T23:13:51.574497Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# EDA and visualization","metadata":{}},{"cell_type":"code","source":"vc = train_df[CFG.target_col].value_counts().sort_values(ascending=False)\n\nplt.figure(figsize=(14,5))\nvc.plot(kind=\"bar\")\nplt.title(\"Train identity distribution\")\nplt.xlabel(\"Jaguar ID\")\nplt.ylabel(\"Image count\")\nplt.show()\n\nprint(vc.describe())\nprint(\"Min images per identity:\", vc.min())\nprint(\"Max images per identity:\", vc.max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:13:51.576344Z","iopub.execute_input":"2026-03-10T23:13:51.576724Z","iopub.status.idle":"2026-03-10T23:13:51.938820Z","shell.execute_reply.started":"2026-03-10T23:13:51.576695Z","shell.execute_reply":"2026-03-10T23:13:51.938135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_rgba_or_rgb(path):\n    img = cv2.imread(str(path), cv2.IMREAD_UNCHANGED)\n    if img is None:\n        raise FileNotFoundError(path)\n    if len(img.shape) == 2:\n        img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n        return img, None\n    if img.shape[2] == 4:\n        rgb = cv2.cvtColor(img[:, :, :3], cv2.COLOR_BGR2RGB)\n        alpha = img[:, :, 3]\n        return rgb, alpha\n    rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    return rgb, None\n\ndef show_random_images(df, n=12):\n    idxs = np.random.choice(len(df), n, replace=False)\n    rows = math.ceil(n / 4)\n    plt.figure(figsize=(16, rows * 4))\n    for i, idx in enumerate(idxs, 1):\n        row = df.iloc[idx]\n        path = CFG.TRAIN_DIR / row[CFG.image_col]\n        rgb, alpha = load_rgba_or_rgb(path)\n        plt.subplot(rows, 4, i)\n        plt.imshow(rgb)\n        plt.title(str(row[CFG.target_col]))\n        plt.axis(\"off\")\n    plt.tight_layout()\n    plt.show()\n\nshow_random_images(train_df, n=12)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:13:51.939727Z","iopub.execute_input":"2026-03-10T23:13:51.939993Z","iopub.status.idle":"2026-03-10T23:14:04.984666Z","shell.execute_reply.started":"2026-03-10T23:13:51.939972Z","shell.execute_reply":"2026-03-10T23:14:04.983690Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Show alpha-mask effect","metadata":{}},{"cell_type":"code","source":"def apply_alpha_mask(rgb, alpha):\n    if alpha is None:\n        return rgb\n    alpha_f = alpha.astype(np.float32) / 255.0\n    out = rgb.astype(np.float32) * alpha_f[..., None]\n    return out.astype(np.uint8)\n\ndef show_mask_examples(df, n=6):\n    idxs = np.random.choice(len(df), n, replace=False)\n    plt.figure(figsize=(12, n * 4))\n    for i, idx in enumerate(idxs):\n        row = df.iloc[idx]\n        path = CFG.TRAIN_DIR / row[CFG.image_col]\n        rgb, alpha = load_rgba_or_rgb(path)\n        masked = apply_alpha_mask(rgb, alpha)\n\n        plt.subplot(n, 2, 2*i+1)\n        plt.imshow(rgb)\n        plt.title(f\"Original - {row[CFG.target_col]}\")\n        plt.axis(\"off\")\n\n        plt.subplot(n, 2, 2*i+2)\n        plt.imshow(masked)\n        plt.title(\"Masked foreground\")\n        plt.axis(\"off\")\n    plt.tight_layout()\n    plt.show()\n\nshow_mask_examples(train_df, n=4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:04.987149Z","iopub.execute_input":"2026-03-10T23:14:04.987413Z","iopub.status.idle":"2026-03-10T23:14:12.566273Z","shell.execute_reply.started":"2026-03-10T23:14:04.987392Z","shell.execute_reply":"2026-03-10T23:14:12.565339Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Folds","metadata":{}},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=CFG.folds, shuffle=True, random_state=CFG.seed)\n\ntrain_df[\"fold\"] = -1\nfor fold, (_, val_idx) in enumerate(skf.split(train_df, train_df[\"label\"])):\n    train_df.loc[val_idx, \"fold\"] = fold\n\ntrain_df[\"fold\"].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:12.567749Z","iopub.execute_input":"2026-03-10T23:14:12.568175Z","iopub.status.idle":"2026-03-10T23:14:12.586755Z","shell.execute_reply.started":"2026-03-10T23:14:12.568133Z","shell.execute_reply":"2026-03-10T23:14:12.586028Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Albumentations transforms","metadata":{}},{"cell_type":"code","source":"def get_train_transforms():\n    return A.Compose([\n        A.LongestMaxSize(max_size=CFG.image_size),\n        A.PadIfNeeded(CFG.image_size, CFG.image_size, border_mode=cv2.BORDER_CONSTANT),\n        A.HorizontalFlip(p=0.5),\n        A.ShiftScaleRotate(\n            shift_limit=0.05, scale_limit=0.10, rotate_limit=12,\n            border_mode=cv2.BORDER_CONSTANT, p=0.5\n        ),\n        A.RandomBrightnessContrast(p=0.5),\n        A.HueSaturationValue(p=0.3),\n        A.GaussNoise(p=0.2),\n        A.CoarseDropout(\n            max_holes=2, max_height=CFG.image_size//10, max_width=CFG.image_size//10, p=0.2\n        ),\n        A.Normalize(),\n        ToTensorV2(),\n    ])\n\ndef get_valid_transforms():\n    return A.Compose([\n        A.LongestMaxSize(max_size=CFG.image_size),\n        A.PadIfNeeded(CFG.image_size, CFG.image_size, border_mode=cv2.BORDER_CONSTANT),\n        A.Normalize(),\n        ToTensorV2(),\n    ])\n\ndef get_tta_transforms():\n    return [\n        A.Compose([\n            A.LongestMaxSize(max_size=CFG.image_size),\n            A.PadIfNeeded(CFG.image_size, CFG.image_size, border_mode=cv2.BORDER_CONSTANT),\n            A.Normalize(),\n            ToTensorV2(),\n        ]),\n        A.Compose([\n            A.HorizontalFlip(p=1.0),\n            A.LongestMaxSize(max_size=CFG.image_size),\n            A.PadIfNeeded(CFG.image_size, CFG.image_size, border_mode=cv2.BORDER_CONSTANT),\n            A.Normalize(),\n            ToTensorV2(),\n        ]),\n    ]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:12.587669Z","iopub.execute_input":"2026-03-10T23:14:12.587905Z","iopub.status.idle":"2026-03-10T23:14:12.595375Z","shell.execute_reply.started":"2026-03-10T23:14:12.587884Z","shell.execute_reply":"2026-03-10T23:14:12.594748Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class JaguarDataset(Dataset):\n    def __init__(self, df, img_dir, transforms=None, is_test=False):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.transforms = transforms\n        self.is_test = is_test\n\n    def __len__(self):\n        return len(self.df)\n\n    def _read_image(self, filename):\n        path = self.img_dir / filename\n        rgb, alpha = load_rgba_or_rgb(path)\n\n        if CFG.use_alpha_as_mask and alpha is not None:\n            rgb = apply_alpha_mask(rgb, alpha)\n\n        return rgb\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image = self._read_image(row[CFG.image_col])\n\n        if self.transforms:\n            image = self.transforms(image=image)[\"image\"]\n\n        if self.is_test:\n            return image, row[CFG.image_col]\n\n        label = int(row[\"label\"])\n        return image, label, row[CFG.image_col]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:12.596212Z","iopub.execute_input":"2026-03-10T23:14:12.596530Z","iopub.status.idle":"2026-03-10T23:14:12.611057Z","shell.execute_reply.started":"2026-03-10T23:14:12.596511Z","shell.execute_reply":"2026-03-10T23:14:12.610453Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Weighted sampler for imbalance","metadata":{}},{"cell_type":"code","source":"def make_weighted_sampler(df):\n    counts = df[\"label\"].value_counts().to_dict()\n    weights = df[\"label\"].map(lambda x: 1.0 / counts[x]).values\n    sampler = WeightedRandomSampler(\n        weights=torch.DoubleTensor(weights),\n        num_samples=len(weights),\n        replacement=True\n    )\n    return sampler","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:12.611732Z","iopub.execute_input":"2026-03-10T23:14:12.611960Z","iopub.status.idle":"2026-03-10T23:14:12.624852Z","shell.execute_reply.started":"2026-03-10T23:14:12.611942Z","shell.execute_reply":"2026-03-10T23:14:12.624249Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# GeM pooling","metadata":{}},{"cell_type":"code","source":"class 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        x = x.clamp(min=self.eps).pow(self.p)\n        x = F.avg_pool2d(x, kernel_size=(x.size(-2), x.size(-1))).pow(1. / self.p)\n        return x.flatten(1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:12.625559Z","iopub.execute_input":"2026-03-10T23:14:12.625783Z","iopub.status.idle":"2026-03-10T23:14:12.636396Z","shell.execute_reply.started":"2026-03-10T23:14:12.625756Z","shell.execute_reply":"2026-03-10T23:14:12.635711Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ArcMarginProduct","metadata":{}},{"cell_type":"code","source":"class ArcMarginProduct(nn.Module):\n    def __init__(self, in_features, out_features, s=30.0, m=0.35):\n        super().__init__()\n        self.in_features = in_features\n        self.out_features = out_features\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, embeddings, labels):\n        cosine = F.linear(F.normalize(embeddings), F.normalize(self.weight))\n        sine = torch.sqrt(torch.clamp(1.0 - cosine ** 2, min=1e-9))\n        phi = cosine * math.cos(self.m) - sine * math.sin(self.m)\n\n        one_hot = torch.zeros_like(cosine)\n        one_hot.scatter_(1, labels.view(-1,1), 1.0)\n\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s\n        return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:12.637231Z","iopub.execute_input":"2026-03-10T23:14:12.637541Z","iopub.status.idle":"2026-03-10T23:14:12.649123Z","shell.execute_reply.started":"2026-03-10T23:14:12.637514Z","shell.execute_reply":"2026-03-10T23:14:12.648241Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class JaguarModel(nn.Module):\n    def __init__(self, backbone_name, num_classes, embedding_dim=512, dropout=0.1):\n        super().__init__()\n        self.backbone = timm.create_model(\n            backbone_name,\n            pretrained=True,\n            num_classes=0,\n            global_pool=\"\"\n        )\n        self.gem = GeM()\n        self.dropout = nn.Dropout(dropout)\n\n        backbone_channels = self.backbone.num_features\n        self.embedding = nn.Linear(backbone_channels, embedding_dim)\n        self.bn = nn.BatchNorm1d(embedding_dim)\n        self.arc = ArcMarginProduct(\n            embedding_dim, num_classes, s=CFG.margin_s, m=CFG.margin_m\n        )\n\n    def forward_features(self, x):\n        feat = self.backbone.forward_features(x)\n        if feat.ndim == 4:\n            feat = self.gem(feat)\n        elif feat.ndim == 3:\n            feat = feat.mean(dim=1)\n        feat = self.dropout(feat)\n        emb = self.embedding(feat)\n        emb = self.bn(emb)\n        emb = F.normalize(emb, p=2, dim=1)\n        return emb\n\n    def forward(self, x, labels=None):\n        emb = self.forward_features(x)\n        if labels is None:\n            return emb\n        logits = self.arc(emb, labels)\n        return logits, emb","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:12.650022Z","iopub.execute_input":"2026-03-10T23:14:12.650255Z","iopub.status.idle":"2026-03-10T23:14:12.665869Z","shell.execute_reply.started":"2026-03-10T23:14:12.650237Z","shell.execute_reply":"2026-03-10T23:14:12.664934Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Retrieval metric for validation","metadata":{}},{"cell_type":"code","source":"def average_precision_at_k(y_true_sorted):\n    y_true_sorted = np.asarray(y_true_sorted)\n    if y_true_sorted.sum() == 0:\n        return 0.0\n    cumsum = np.cumsum(y_true_sorted)\n    precision = cumsum / (np.arange(len(y_true_sorted)) + 1)\n    ap = (precision * y_true_sorted).sum() / y_true_sorted.sum()\n    return ap\n\ndef macro_map_by_identity(embeddings, labels):\n    embeddings = np.asarray(embeddings)\n    labels = np.asarray(labels)\n\n    sim = cosine_similarity(embeddings)\n    n = len(labels)\n\n    ap_per_query = []\n    identity_to_aps = defaultdict(list)\n\n    for i in range(n):\n        order = np.argsort(-sim[i])\n        order = order[order != i]\n\n        y_true = (labels[order] == labels[i]).astype(np.int32)\n        ap = average_precision_at_k(y_true)\n\n        ap_per_query.append(ap)\n        identity_to_aps[labels[i]].append(ap)\n\n    per_identity_map = [np.mean(v) for v in identity_to_aps.values()]\n    return float(np.mean(per_identity_map)), float(np.mean(ap_per_query))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:12.666695Z","iopub.execute_input":"2026-03-10T23:14:12.666919Z","iopub.status.idle":"2026-03-10T23:14:12.679497Z","shell.execute_reply.started":"2026-03-10T23:14:12.666898Z","shell.execute_reply":"2026-03-10T23:14:12.678810Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train / valid loops","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, scaler, criterion, device):\n    model.train()\n    losses = []\n\n    for images, labels, _ in tqdm(loader, desc=\"train\", leave=False):\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n\n        with torch.cuda.amp.autocast(enabled=CFG.amp):\n            logits, _ = model(images, labels)\n            loss = criterion(logits, labels)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        losses.append(loss.item())\n\n    return np.mean(losses)\n\n\n@torch.no_grad()\ndef valid_one_epoch(model, loader, criterion, device):\n    model.eval()\n    losses = []\n\n    all_embs = []\n    all_labels = []\n\n    for images, labels, _ in tqdm(loader, desc=\"valid\", leave=False):\n        images = images.to(device)\n        labels = labels.to(device)\n\n        with torch.cuda.amp.autocast(enabled=CFG.amp):\n            logits, embs = model(images, labels)\n            loss = criterion(logits, labels)\n\n        losses.append(loss.item())\n        all_embs.append(embs.detach().cpu().numpy())\n        all_labels.append(labels.detach().cpu().numpy())\n\n    all_embs = np.concatenate(all_embs)\n    all_labels = np.concatenate(all_labels)\n\n    macro_map, query_map = macro_map_by_identity(all_embs, all_labels)\n    return np.mean(losses), macro_map, query_map","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:12.680588Z","iopub.execute_input":"2026-03-10T23:14:12.680896Z","iopub.status.idle":"2026-03-10T23:14:12.695951Z","shell.execute_reply.started":"2026-03-10T23:14:12.680874Z","shell.execute_reply":"2026-03-10T23:14:12.695168Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train one fold","metadata":{}},{"cell_type":"code","source":"def train_fold(fold):\n    print(f\"\\n==================== FOLD {fold} ====================\")\n\n    tr_df = train_df[train_df.fold != fold].reset_index(drop=True)\n    va_df = train_df[train_df.fold == fold].reset_index(drop=True)\n\n    train_dataset = JaguarDataset(tr_df, CFG.TRAIN_DIR, transforms=get_train_transforms(), is_test=False)\n    valid_dataset = JaguarDataset(va_df, CFG.TRAIN_DIR, transforms=get_valid_transforms(), is_test=False)\n\n    sampler = make_weighted_sampler(tr_df)\n\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=CFG.batch_size,\n        sampler=sampler,\n        num_workers=CFG.num_workers,\n        pin_memory=True,\n        drop_last=True\n    )\n\n    valid_loader = DataLoader(\n        valid_dataset,\n        batch_size=CFG.batch_size,\n        shuffle=False,\n        num_workers=CFG.num_workers,\n        pin_memory=True\n    )\n\n    model = JaguarModel(\n        backbone_name=CFG.backbone,\n        num_classes=num_classes,\n        embedding_dim=CFG.embedding_dim,\n        dropout=CFG.dropout\n    ).to(CFG.device)\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n    scaler = torch.cuda.amp.GradScaler(enabled=CFG.amp)\n    criterion = nn.CrossEntropyLoss(label_smoothing=CFG.label_smoothing)\n\n    best_score = -1\n\n    for epoch in range(CFG.epochs):\n        start = time.time()\n\n        train_loss = train_one_epoch(model, train_loader, optimizer, scaler, criterion, CFG.device)\n        valid_loss, macro_map, query_map = valid_one_epoch(model, valid_loader, criterion, CFG.device)\n\n        elapsed = time.time() - start\n        print(\n            f\"Fold {fold} | Epoch {epoch+1:02d}/{CFG.epochs} | \"\n            f\"train_loss={train_loss:.4f} | valid_loss={valid_loss:.4f} | \"\n            f\"macro_mAP={macro_map:.5f} | query_mAP={query_map:.5f} | time={elapsed:.1f}s\"\n        )\n\n        if macro_map > best_score:\n            best_score = macro_map\n            save_path = CFG.model_dir / f\"fold_{fold}.pth\"\n            torch.save(\n                {\n                    \"model\": model.state_dict(),\n                    \"score\": best_score,\n                    \"fold\": fold\n                },\n                save_path\n            )\n            print(f\"Saved best model to {save_path}\")\n\n    del model, optimizer, scaler, train_loader, valid_loader\n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:14:12.696880Z","iopub.execute_input":"2026-03-10T23:14:12.697691Z","iopub.status.idle":"2026-03-10T23:14:12.712767Z","shell.execute_reply.started":"2026-03-10T23:14:12.697663Z","shell.execute_reply":"2026-03-10T23:14:12.712002Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train all folds","metadata":{}},{"cell_type":"code","source":"for fold in range(CFG.folds):\n    train_fold(fold)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:13:24.494473Z","iopub.execute_input":"2026-03-10T23:13:24.494810Z","iopub.status.idle":"2026-03-10T23:13:24.501206Z","shell.execute_reply.started":"2026-03-10T23:13:24.494785Z","shell.execute_reply":"2026-03-10T23:13:24.500316Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load model helper","metadata":{}},{"cell_type":"code","source":"def load_fold_model(fold):\n    model = JaguarModel(\n        backbone_name=CFG.backbone,\n        num_classes=num_classes,\n        embedding_dim=CFG.embedding_dim,\n        dropout=CFG.dropout\n    ).to(CFG.device)\n\n    ckpt = torch.load(CFG.model_dir / f\"fold_{fold}.pth\", map_location=CFG.device)\n    model.load_state_dict(ckpt[\"model\"])\n    model.eval()\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T21:18:58.920644Z","iopub.execute_input":"2026-03-10T21:18:58.920950Z","iopub.status.idle":"2026-03-10T21:18:58.931344Z","shell.execute_reply.started":"2026-03-10T21:18:58.920923Z","shell.execute_reply":"2026-03-10T21:18:58.930644Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Build unique test-image dataframe","metadata":{}},{"cell_type":"code","source":"unique_test_images = sorted(set(test_df[\"query_image\"].unique()) | set(test_df[\"gallery_image\"].unique()))\nunique_test_df = pd.DataFrame({CFG.image_col: unique_test_images})\n\nprint(\"Unique test images:\", len(unique_test_df))\nunique_test_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T21:18:58.932433Z","iopub.execute_input":"2026-03-10T21:18:58.932748Z","iopub.status.idle":"2026-03-10T21:18:59.018991Z","shell.execute_reply.started":"2026-03-10T21:18:58.932718Z","shell.execute_reply":"2026-03-10T21:18:59.018503Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference with TTA","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef extract_embeddings_for_df(model, df, img_dir, transforms):\n    dataset = JaguarDataset(df, img_dir, transforms=transforms, is_test=True)\n    loader = DataLoader(\n        dataset,\n        batch_size=CFG.batch_size,\n        shuffle=False,\n        num_workers=CFG.num_workers,\n        pin_memory=True\n    )\n\n    names = []\n    embs = []\n\n    for images, img_names in tqdm(loader, leave=False):\n        images = images.to(CFG.device)\n\n        with torch.cuda.amp.autocast(enabled=CFG.amp):\n            feat = model(images)\n\n        feat = F.normalize(feat, p=2, dim=1)\n        embs.append(feat.cpu().numpy())\n        names.extend(list(img_names))\n\n    embs = np.concatenate(embs)\n    return names, embs\n\n\n@torch.no_grad()\ndef predict_test_embeddings():\n    all_fold_embeddings = []\n\n    for fold in range(CFG.folds):\n        print(f\"Loading fold {fold}\")\n        model = load_fold_model(fold)\n\n        if CFG.use_tta:\n            tta_embs = []\n            for tta_tfms in get_tta_transforms():\n                names, embs = extract_embeddings_for_df(model, unique_test_df, CFG.TEST_DIR, tta_tfms)\n                tta_embs.append(embs)\n            embs = np.mean(tta_embs, axis=0)\n            embs = embs / np.linalg.norm(embs, axis=1, keepdims=True)\n        else:\n            names, embs = extract_embeddings_for_df(model, unique_test_df, CFG.TEST_DIR, get_valid_transforms())\n\n        all_fold_embeddings.append(embs)\n\n        del model\n        gc.collect()\n        torch.cuda.empty_cache()\n\n    final_embs = np.mean(all_fold_embeddings, axis=0)\n    final_embs = final_embs / np.linalg.norm(final_embs, axis=1, keepdims=True)\n\n    return names, final_embs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T21:18:59.020161Z","iopub.execute_input":"2026-03-10T21:18:59.020428Z","iopub.status.idle":"2026-03-10T21:18:59.028689Z","shell.execute_reply.started":"2026-03-10T21:18:59.020406Z","shell.execute_reply":"2026-03-10T21:18:59.027935Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":" # Run test embedding extraction","metadata":{}},{"cell_type":"code","source":"test_names, test_embs = predict_test_embeddings()\nprint(test_embs.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T21:18:59.029643Z","iopub.execute_input":"2026-03-10T21:18:59.029966Z","iopub.status.idle":"2026-03-10T21:38:33.597094Z","shell.execute_reply.started":"2026-03-10T21:18:59.029943Z","shell.execute_reply":"2026-03-10T21:38:33.594877Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualize embedding similarity matrix","metadata":{}},{"cell_type":"code","source":"sim_matrix = cosine_similarity(test_embs)\n\nplt.figure(figsize=(8,8))\nplt.imshow(sim_matrix, cmap=\"viridis\")\nplt.title(\"Test cosine similarity matrix\")\nplt.colorbar()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T21:38:33.599791Z","iopub.execute_input":"2026-03-10T21:38:33.600154Z","iopub.status.idle":"2026-03-10T21:38:34.128167Z","shell.execute_reply.started":"2026-03-10T21:38:33.600114Z","shell.execute_reply":"2026-03-10T21:38:34.127271Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Map image to embedding index","metadata":{}},{"cell_type":"code","source":"img2idx = {img: i for i, img in enumerate(test_names)}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T21:38:34.129374Z","iopub.execute_input":"2026-03-10T21:38:34.130016Z","iopub.status.idle":"2026-03-10T21:38:34.134008Z","shell.execute_reply.started":"2026-03-10T21:38:34.129991Z","shell.execute_reply":"2026-03-10T21:38:34.133077Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Build submission from pairwise similarity","metadata":{}},{"cell_type":"code","source":"similarities = []\n\nfor row in tqdm(test_df.itertuples(index=False), total=len(test_df)):\n    q_idx = img2idx[row.query_image]\n    g_idx = img2idx[row.gallery_image]\n\n    sim = sim_matrix[q_idx, g_idx]\n    sim = (sim + 1.0) / 2.0\n    sim = float(np.clip(sim, 0.0, 1.0))\n    similarities.append(sim)\n\nsubmission = pd.DataFrame({\n    \"row_id\": test_df[\"row_id\"].values,\n    \"similarity\": similarities\n})\n\nsubmission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:07:44.088895Z","iopub.execute_input":"2026-03-10T23:07:44.089627Z","iopub.status.idle":"2026-03-10T23:07:44.097214Z","shell.execute_reply.started":"2026-03-10T23:07:44.089596Z","shell.execute_reply":"2026-03-10T23:07:44.096043Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Validate submission","metadata":{}},{"cell_type":"code","source":"def validate_submission(submission, test_df):\n    assert submission.shape[0] == test_df.shape[0], \"Wrong row count\"\n    assert list(submission.columns) == [\"row_id\", \"similarity\"], \"Wrong columns\"\n    assert (submission[\"row_id\"].values == test_df[\"row_id\"].values).all(), \"row_id mismatch\"\n\n    sims = submission[\"similarity\"].values\n    assert np.isfinite(sims).all(), \"NaN/Inf found\"\n    assert (sims >= 0).all(), \"similarity < 0 found\"\n    assert (sims <= 1).all(), \"similarity > 1 found\"\n\n    print(\"Submission valid\")\n    print(submission[\"similarity\"].describe())\n\nvalidate_submission(submission, test_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T21:38:35.050306Z","iopub.execute_input":"2026-03-10T21:38:35.050743Z","iopub.status.idle":"2026-03-10T21:38:35.073160Z","shell.execute_reply.started":"2026-03-10T21:38:35.050715Z","shell.execute_reply":"2026-03-10T21:38:35.072415Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Save submission","metadata":{}},{"cell_type":"code","source":"submission.to_csv(\"submission_convnextv2_arcface.csv\", index=False)\nprint(\"Saved submission_convnextv2_arcface.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-10T23:07:18.333932Z","iopub.execute_input":"2026-03-10T23:07:18.334831Z","iopub.status.idle":"2026-03-10T23:07:18.348241Z","shell.execute_reply.started":"2026-03-10T23:07:18.334794Z","shell.execute_reply":"2026-03-10T23:07:18.347356Z"}},"outputs":[],"execution_count":null}]}