{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"gpu","dataSources":[{"sourceId":129543,"databundleVersionId":15525987,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14715175,"sourceType":"datasetVersion","datasetId":9401709}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Jaguar Re-Identification - exp001_baseline (5-Fold CV)\n\nTraining and inference pipeline using MegaDescriptor-Large + ArcFace Loss with 5-Fold Cross Validation, EMA, and AMP","metadata":{}},{"cell_type":"markdown","source":"## 1. Setup & Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nfrom collections import defaultdict\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\nfrom dataclasses import dataclass, field\nfrom pathlib import Path\nfrom typing import Literal\n\nimport albumentations as A\nimport numpy as np\nimport pandas as pd\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom albumentations.pytorch import ToTensorV2\nfrom PIL import Image\nfrom sklearn.model_selection import StratifiedKFold\nfrom timm.utils import ModelEmaV3\nfrom torch.amp import GradScaler, autocast\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\n\ntorch.set_float32_matmul_precision(\"high\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:03:42.445640Z","iopub.execute_input":"2026-02-03T05:03:42.446828Z","iopub.status.idle":"2026-02-03T05:03:56.456896Z","shell.execute_reply.started":"2026-02-03T05:03:42.446792Z","shell.execute_reply":"2026-02-03T05:03:56.456288Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    # Paths\n    input_dir = \"/kaggle/input/round-2-jaguar-reidentification-challenge\"\n    output_dir = \"/kaggle/tmp\"\n\n    # Model\n    model_name = \"hf-hub:BVRA/MegaDescriptor-L-384\"  # MegaDescriptor-Large\n    embedding_dim = 1024\n    num_classes = 31  # Number of jaguar individuals\n\n    # Training\n    seed = 42\n    n_splits = 5  # 5-fold CV\n    epochs = 20\n    batch_size = 8\n    num_workers = 4\n    lr = 1e-4\n    weight_decay = 1e-4\n\n    # ArcFace\n    arcface_s = 30.0  # Scale\n    arcface_m = 0.5  # Margin\n\n    # EMA\n    use_ema = True\n    ema_decay = 0.995\n\n    # Image\n    img_size = 384\n\n    # Device\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:03:56.458562Z","iopub.execute_input":"2026-02-03T05:03:56.458998Z","iopub.status.idle":"2026-02-03T05:03:56.515329Z","shell.execute_reply.started":"2026-02-03T05:03:56.458972Z","shell.execute_reply":"2026-02-03T05:03:56.514387Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Convert Images to NPY Format (Optional)","metadata":{}},{"cell_type":"code","source":"@dataclass\nclass ConvertConfig:\n    \"\"\"Configuration for image conversion\"\"\"\n\n    # Input/Output paths\n    input_dir: str = \"/kaggle/input/round-2-jaguar-reidentification-challenge\"\n    output_dir: str = \"/kaggle/tmp\"\n\n    # Resize settings (width, height). None keeps original size\n    resize: tuple[int, int] | None = None\n\n    # If True, save all images in a single npy file\n    batch_save: bool = False\n\n    # Number of parallel workers\n    num_workers: int = 8\n\n    # Target split to convert\n    split: Literal[\"train\", \"test\", \"both\"] = \"both\"\n\n    # Target image extensions\n    image_extensions: list[str] = field(default_factory=lambda: [\".png\", \".jpg\", \".jpeg\"])\n\n\ndef load_image(image_path: str, resize: tuple[int, int] | None = None) -> np.ndarray:\n    \"\"\"Load image and return as numpy array\"\"\"\n    img = Image.open(image_path).convert(\"RGB\")\n    if resize is not None:\n        img = img.resize(resize, Image.LANCZOS)\n    return np.array(img, dtype=np.uint8)\n\n\ndef process_single_image(args: tuple) -> tuple[str, np.ndarray]:\n    \"\"\"Process a single image\"\"\"\n    image_path, resize = args\n    return os.path.basename(image_path), load_image(image_path, resize)\n\n\ndef convert_images_to_npy(\n    input_dir: str,\n    output_dir: str,\n    config: ConvertConfig,\n) -> None:\n    \"\"\"\n    Convert images in a directory to npy format\n\n    Args:\n        input_dir: Input image directory\n        output_dir: Output directory\n        config: Conversion configuration\n    \"\"\"\n    input_path = Path(input_dir)\n    output_path = Path(output_dir)\n    output_path.mkdir(parents=True, exist_ok=True)\n\n    # Get list of image files\n    image_files = sorted([f for f in input_path.iterdir() if f.suffix.lower() in config.image_extensions])\n    print(f\"Found {len(image_files)} images in {input_dir}\")\n\n    if config.batch_save and config.resize is None:\n        raise ValueError(\"batch_save requires resize to be specified (images must have same size)\")\n\n    if config.batch_save:\n        # Save all images in a single npy file\n        images = []\n        filenames = []\n\n        with ThreadPoolExecutor(max_workers=config.num_workers) as executor:\n            futures = {executor.submit(process_single_image, (str(f), config.resize)): f for f in image_files}\n\n            for future in tqdm(as_completed(futures), total=len(image_files), desc=\"Loading images\"):\n                filename, img_array = future.result()\n                images.append(img_array)\n                filenames.append(filename)\n\n        # Sort by filename\n        sorted_indices = np.argsort(filenames)\n        images = np.array([images[i] for i in sorted_indices])\n        filenames = [filenames[i] for i in sorted_indices]\n\n        # Save\n        dir_name = input_path.name\n        np.save(output_path / f\"{dir_name}_images.npy\", images)\n        np.save(output_path / f\"{dir_name}_filenames.npy\", np.array(filenames))\n        print(f\"Saved: {output_path / f'{dir_name}_images.npy'} - shape: {images.shape}\")\n\n    else:\n        # Save as individual npy files\n        def save_single(image_file: Path) -> None:\n            img_array = load_image(str(image_file), config.resize)\n            output_file = output_path / f\"{image_file.stem}.npy\"\n            np.save(output_file, img_array)\n\n        with ThreadPoolExecutor(max_workers=config.num_workers) as executor:\n            list(tqdm(executor.map(save_single, image_files), total=len(image_files), desc=\"Converting images\"))\n\n        print(f\"Saved {len(image_files)} npy files to {output_dir}\")\n\n\ndef run_convert(config: ConvertConfig) -> None:\n    \"\"\"Run image conversion based on configuration\"\"\"\n    splits = [\"train\", \"test\"] if config.split == \"both\" else [config.split]\n\n    for split in splits:\n        input_dir = os.path.join(config.input_dir, split)\n        output_dir = os.path.join(config.output_dir, split)\n\n        if not os.path.exists(input_dir):\n            print(f\"Skipping {split}: {input_dir} does not exist\")\n            continue\n\n        print(f\"\\nProcessing {split}...\")\n        convert_images_to_npy(\n            input_dir=input_dir,\n            output_dir=output_dir,\n            config=config,\n        )\n\n    print(\"\\nDone!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:03:56.516498Z","iopub.execute_input":"2026-02-03T05:03:56.516798Z","iopub.status.idle":"2026-02-03T05:03:56.678203Z","shell.execute_reply.started":"2026-02-03T05:03:56.516767Z","shell.execute_reply":"2026-02-03T05:03:56.677542Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"convert_config = ConvertConfig()\nrun_convert(convert_config)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:03:56.679095Z","iopub.execute_input":"2026-02-03T05:03:56.679442Z","iopub.status.idle":"2026-02-03T05:08:18.788953Z","shell.execute_reply.started":"2026-02-03T05:03:56.679411Z","shell.execute_reply":"2026-02-03T05:08:18.788120Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Utility Functions, Datasets & Model Definitions","metadata":{}},{"cell_type":"code","source":"def set_seed(seed):\n    \"\"\"Set seed for reproducibility\"\"\"\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(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:08:18.790569Z","iopub.execute_input":"2026-02-03T05:08:18.790868Z","iopub.status.idle":"2026-02-03T05:08:18.795507Z","shell.execute_reply.started":"2026-02-03T05:08:18.790842Z","shell.execute_reply":"2026-02-03T05:08:18.794832Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_train_transforms(img_size, seed=42):\n    \"\"\"Get training data augmentation\"\"\"\n    return A.Compose(\n        [\n            A.Resize(img_size, img_size),\n            A.HorizontalFlip(p=0.5),\n            A.Affine(translate_percent=0.1, scale=(0.85, 1.15), rotate=(-15, 15), p=0.5),\n            A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1, p=0.5),\n            A.GaussNoise(std_range=(0.02, 0.1), p=0.3),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ],\n        seed=seed,\n    )\n\n\ndef get_valid_transforms(img_size, seed=42):\n    \"\"\"Get validation/test transforms\"\"\"\n    return A.Compose(\n        [\n            A.Resize(img_size, img_size),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ],\n        seed=seed,\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:08:18.796300Z","iopub.execute_input":"2026-02-03T05:08:18.796912Z","iopub.status.idle":"2026-02-03T05:08:27.262517Z","shell.execute_reply.started":"2026-02-03T05:08:18.796889Z","shell.execute_reply":"2026-02-03T05:08:27.261646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class JaguarDataset(Dataset):\n    \"\"\"Jaguar image dataset\"\"\"\n\n    def __init__(self, df, image_dir, transforms=None, is_test=False, use_npy=True):\n        self.df = df\n        self.image_dir = image_dir\n        self.transforms = transforms\n        self.is_test = is_test\n        self.use_npy = use_npy\n\n        if not is_test:\n            # Create label encoder\n            self.labels = df[\"ground_truth\"].values\n            self.label2idx = {label: idx for idx, label in enumerate(sorted(df[\"ground_truth\"].unique()))}\n            self.idx2label = {idx: label for label, idx in self.label2idx.items()}\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        # Load image\n        if self.use_npy:\n            npy_filename = os.path.splitext(row[\"filename\"])[0] + \".npy\"\n            img_path = os.path.join(self.image_dir, npy_filename)\n            image = np.load(img_path)\n        else:\n            img_path = os.path.join(self.image_dir, row[\"filename\"])\n            image = np.array(Image.open(img_path).convert(\"RGB\"))\n\n        # Apply transforms\n        if self.transforms:\n            image = self.transforms(image=image)[\"image\"]\n\n        if self.is_test:\n            return image, row[\"filename\"]\n        else:\n            label = self.label2idx[row[\"ground_truth\"]]\n            return image, label\n\n\nclass TestImageDataset(Dataset):\n    \"\"\"Test image dataset (unique images only)\"\"\"\n\n    def __init__(self, image_files, image_dir, transforms=None):\n        self.image_files = image_files\n        self.image_dir = image_dir\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.image_files)\n\n    def __getitem__(self, idx):\n        filename = self.image_files[idx]\n        img_path = os.path.join(self.image_dir, filename)\n        image = np.array(Image.open(img_path).convert(\"RGB\"))\n\n        if self.transforms:\n            image = self.transforms(image=image)[\"image\"]\n\n        return image, filename","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:08:27.264761Z","iopub.execute_input":"2026-02-03T05:08:27.265006Z","iopub.status.idle":"2026-02-03T05:08:27.273970Z","shell.execute_reply.started":"2026-02-03T05:08:27.264983Z","shell.execute_reply":"2026-02-03T05:08:27.273268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ArcFaceHead(nn.Module):\n    \"\"\"ArcFace head for metric learning\"\"\"\n\n    def __init__(self, in_features, out_features, s=30.0, m=0.5):\n        super().__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.s = s\n        self.m = m\n\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n    def forward(self, features, labels=None):\n        # Normalize features and weights\n        features = F.normalize(features, dim=1)\n        weight = F.normalize(self.weight, dim=1)\n\n        # Cosine similarity\n        cosine = F.linear(features, weight)\n\n        if labels is None:\n            return cosine\n\n        # ArcFace margin\n        theta = torch.acos(torch.clamp(cosine, -1.0 + 1e-7, 1.0 - 1e-7))\n        target_logits = torch.cos(theta + self.m)\n\n        one_hot = torch.zeros_like(cosine)\n        one_hot.scatter_(1, labels.view(-1, 1), 1)\n\n        output = cosine * (1 - one_hot) + target_logits * one_hot\n        output *= self.s\n\n        return output\n\n\nclass JaguarReIDModel(nn.Module):\n    \"\"\"Jaguar re-identification model using MegaDescriptor backbone\"\"\"\n\n    def __init__(self, model_name, embedding_dim, num_classes, s=30.0, m=0.5, pretrained=True):\n        super().__init__()\n\n        # Backbone\n        self.backbone = timm.create_model(model_name, pretrained=pretrained, num_classes=0)\n        backbone_out = self.backbone.num_features\n\n        # Embedding layer\n        self.embedding = nn.Sequential(\n            nn.Linear(backbone_out, embedding_dim),\n            nn.BatchNorm1d(embedding_dim),\n        )\n\n        # ArcFace head\n        self.arcface = ArcFaceHead(embedding_dim, num_classes, s=s, m=m)\n\n    def get_embedding(self, x):\n        \"\"\"Extract embeddings for inference\"\"\"\n        features = self.backbone(x)\n        embeddings = self.embedding(features)\n        return F.normalize(embeddings, dim=1)\n\n    def forward(self, x, labels=None):\n        features = self.backbone(x)\n        embeddings = self.embedding(features)\n\n        if labels is not None:\n            output = self.arcface(embeddings, labels)\n            return output, embeddings\n\n        return F.normalize(embeddings, dim=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:08:27.274774Z","iopub.execute_input":"2026-02-03T05:08:27.275049Z","iopub.status.idle":"2026-02-03T05:08:27.290458Z","shell.execute_reply.started":"2026-02-03T05:08:27.275027Z","shell.execute_reply":"2026-02-03T05:08:27.289880Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Training","metadata":{}},{"cell_type":"code","source":"class BalancedSampler:\n    \"\"\"Balanced sampling to handle class imbalance\"\"\"\n\n    def __init__(self, labels, batch_size, samples_per_class=4):\n        self.labels = labels\n        self.batch_size = batch_size\n        self.samples_per_class = samples_per_class\n\n        # Group indices by class\n        self.class_indices = defaultdict(list)\n        for idx, label in enumerate(labels):\n            self.class_indices[label].append(idx)\n\n        self.classes = list(self.class_indices.keys())\n        self.num_classes = len(self.classes)\n\n    def __iter__(self):\n        # Number of batches\n        num_batches = len(self.labels) // self.batch_size\n\n        for _ in range(num_batches):\n            batch = []\n            # Select classes for this batch\n            selected_classes = random.sample(\n                self.classes, min(self.batch_size // self.samples_per_class, self.num_classes)\n            )\n\n            for cls in selected_classes:\n                indices = self.class_indices[cls]\n                selected = random.choices(indices, k=self.samples_per_class)\n                batch.extend(selected)\n\n            # Fill remaining with random samples\n            while len(batch) < self.batch_size:\n                batch.append(random.choice(range(len(self.labels))))\n\n            yield batch[: self.batch_size]\n\n    def __len__(self):\n        return len(self.labels) // self.batch_size","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:08:27.291216Z","iopub.execute_input":"2026-02-03T05:08:27.291488Z","iopub.status.idle":"2026-02-03T05:08:27.304179Z","shell.execute_reply.started":"2026-02-03T05:08:27.291466Z","shell.execute_reply":"2026-02-03T05:08:27.303582Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_pbar(dataloader):\n    \"\"\"Create progress bar\"\"\"\n    pbar = tqdm(\n        enumerate(dataloader),\n        total=len(dataloader),\n        bar_format=\"{l_bar}{bar:10}{r_bar}\",\n    )\n    return pbar\n\n\ndef pbar_train_desc(pbar_train, lr, epoch, total_epoch, average_loss):\n    \"\"\"Update training progress bar description\"\"\"\n    learning_rate = f\"LR : {lr:.2E}\"\n    gpu_memory = f\"Mem : {torch.cuda.memory_reserved() / 1e9:.3g}GB\"\n    epoch_info = f\"Epoch {epoch}/{total_epoch}\"\n    loss = f\"Loss: {average_loss:.4f}\"\n    description = f\"{epoch_info:12} {gpu_memory:15} {learning_rate:15} {loss:15}\"\n    pbar_train.set_description(description)\n\n\ndef pbar_valid_desc(pbar_val, average_loss):\n    \"\"\"Update validation progress bar description\"\"\"\n    average_val_loss = f\"Val Loss: {average_loss:.4f}\"\n    description = f\"{average_val_loss:18}\"\n    pbar_val.set_description(description)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:08:27.305140Z","iopub.execute_input":"2026-02-03T05:08:27.305430Z","iopub.status.idle":"2026-02-03T05:08:27.315431Z","shell.execute_reply.started":"2026-02-03T05:08:27.305407Z","shell.execute_reply":"2026-02-03T05:08:27.314711Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_mAP(embeddings, labels):\n    \"\"\"Calculate Identity-Balanced mAP (competition metric)\"\"\"\n    similarity = torch.mm(embeddings, embeddings.t())\n    unique_labels = labels.unique()\n    identity_ap_list = []\n\n    for label in unique_labels:\n        mask = labels == label\n        query_indices = torch.where(mask)[0]\n        query_aps = []\n\n        for q_idx in query_indices:\n            sims = similarity[q_idx].clone()\n            sims[q_idx] = -1  # Exclude self\n            sorted_indices = torch.argsort(sims, descending=True)\n            gt = (labels[sorted_indices] == label).float()\n\n            if gt.sum() == 0:\n                continue\n\n            cumsum = torch.cumsum(gt, dim=0)\n            precision = cumsum / (torch.arange(len(gt), device=gt.device) + 1).float()\n            ap = (precision * gt).sum() / gt.sum()\n            query_aps.append(ap.item())\n\n        if query_aps:\n            identity_ap_list.append(np.mean(query_aps))\n\n    mAP = np.mean(identity_ap_list) if identity_ap_list else 0.0\n    return mAP","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:08:27.316301Z","iopub.execute_input":"2026-02-03T05:08:27.316576Z","iopub.status.idle":"2026-02-03T05:08:27.329515Z","shell.execute_reply.started":"2026-02-03T05:08:27.316554Z","shell.execute_reply":"2026-02-03T05:08:27.328838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_main():\n    print(f\"Device: {CFG.device}\")\n    set_seed(CFG.seed)\n\n    # Create output directory\n    os.makedirs(CFG.output_dir, exist_ok=True)\n\n    # Load data\n    train_df = pd.read_csv(os.path.join(CFG.input_dir, \"train.csv\"))\n\n    print(f\"Number of training samples: {len(train_df)}\")\n    print(f\"Number of unique individuals: {train_df['ground_truth'].nunique()}\")\n\n    # 5-Fold Cross Validation\n    skf = StratifiedKFold(n_splits=CFG.n_splits, shuffle=True, random_state=CFG.seed)\n    cv_scores = []\n\n    for fold, (train_idx, valid_idx) in enumerate(skf.split(train_df, train_df[\"ground_truth\"])):\n        print(f\"\\n{'=' * 60}\")\n        print(f\"Fold {fold + 1}/{CFG.n_splits}\")\n        print(f\"{'=' * 60}\")\n\n        # ============================================================\n        # Dataset\n        # ============================================================\n        fold_train_df = train_df.iloc[train_idx].reset_index(drop=True)\n        fold_valid_df = train_df.iloc[valid_idx].reset_index(drop=True)\n\n        print(f\"Training samples: {len(fold_train_df)}, Validation samples: {len(fold_valid_df)}\")\n\n        train_transforms = get_train_transforms(CFG.img_size, seed=CFG.seed)\n        valid_transforms = get_valid_transforms(CFG.img_size, seed=CFG.seed)\n\n        train_dataset = JaguarDataset(\n            fold_train_df,\n            os.path.join(CFG.output_dir, \"train\"),\n            transforms=train_transforms,\n            is_test=False,\n            use_npy=True,\n        )\n        valid_dataset = JaguarDataset(\n            fold_valid_df,\n            os.path.join(CFG.output_dir, \"train\"),\n            transforms=valid_transforms,\n            is_test=False,\n            use_npy=True,\n        )\n\n        # Balanced sampling\n        train_labels = [train_dataset.label2idx[row] for row in fold_train_df[\"ground_truth\"]]\n        sampler = BalancedSampler(train_labels, CFG.batch_size)\n\n        train_loader = DataLoader(train_dataset, batch_sampler=sampler, num_workers=CFG.num_workers, pin_memory=True)\n        valid_loader = DataLoader(\n            valid_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers, pin_memory=True\n        )\n\n        # ============================================================\n        # Model\n        # ============================================================\n        model = JaguarReIDModel(\n            CFG.model_name, CFG.embedding_dim, CFG.num_classes, s=CFG.arcface_s, m=CFG.arcface_m, pretrained=True\n        )\n        model = model.to(CFG.device)\n        # model = torch.compile(model)\n\n        # EMA\n        ema_model = None\n        if CFG.use_ema:\n            ema_model = ModelEmaV3(model, decay=CFG.ema_decay)\n\n        criterion = nn.CrossEntropyLoss()\n        optimizer = AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n        scheduler = CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-6)\n        scaler = GradScaler()\n\n        best_mAP = 0\n        for epoch in range(CFG.epochs):\n            # ============================================================\n            # Train\n            # ============================================================\n            model.train()\n            total_loss = 0\n\n            pbar = get_pbar(train_loader)\n            for batch_idx, (images, labels) in pbar:\n                images = images.to(CFG.device)\n                labels = labels.to(CFG.device)\n\n                optimizer.zero_grad()\n\n                with autocast(device_type=\"cuda\", dtype=torch.float16):\n                    output, _ = model(images, labels)\n                    loss = criterion(output, labels)\n\n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n\n                if ema_model is not None:\n                    ema_model.update(model)\n\n                total_loss += loss.item()\n                avg_loss = total_loss / (batch_idx + 1)\n                pbar_train_desc(pbar, scheduler.get_last_lr()[0], epoch + 1, CFG.epochs, avg_loss)\n\n            scheduler.step()\n\n            # ============================================================\n            # Validate (every 5 epochs or last epoch)\n            # ============================================================\n            if (epoch + 1) % 5 != 0 and epoch != CFG.epochs - 1:\n                continue\n\n            eval_model = ema_model.module if ema_model is not None else model\n            eval_model.eval()\n\n            embeddings_list = []\n            labels_list = []\n            pbar = get_pbar(valid_loader)\n            valid_loss = 0\n\n            with torch.no_grad():\n                for batch_idx, (images, labels) in pbar:\n                    images = images.to(CFG.device)\n                    labels = labels.to(CFG.device)\n\n                    with autocast(device_type=\"cuda\", dtype=torch.float16):\n                        output, _ = eval_model(images, labels)\n                        loss = F.cross_entropy(output, labels)\n                        embeddings = eval_model.get_embedding(images)\n\n                    valid_loss += loss.item()\n                    avg_loss = valid_loss / (batch_idx + 1)\n\n                    embeddings_list.append(embeddings.float().cpu())\n                    labels_list.append(labels)\n                    pbar_valid_desc(pbar, avg_loss)\n\n            embeddings = torch.cat(embeddings_list, dim=0)\n            labels = torch.cat(labels_list, dim=0)\n\n            mAP = calculate_mAP(embeddings, labels)\n            print(f\"val_mAP: {mAP:.4f}\")\n\n            if mAP > best_mAP:\n                best_mAP = mAP\n                save_dict = {\"model_state_dict\": model.state_dict(), \"best_mAP\": best_mAP}\n                if ema_model is not None:\n                    save_dict[\"ema_state_dict\"] = ema_model.module.state_dict()\n                torch.save(save_dict, os.path.join(CFG.output_dir, f\"best_model_fold{fold}.pth\"))\n                print(f\"Best model saved (mAP: {best_mAP:.4f})\")\n\n        cv_scores.append(best_mAP)\n        print(f\"\\nFold {fold + 1} completed! Best mAP: {best_mAP:.4f}\")\n\n    # ============================================================\n    # CV Summary\n    # ============================================================\n    print(f\"\\n{'=' * 60}\")\n    print(\"CV result summary\")\n    print(f\"{'=' * 60}\")\n    for fold, score in enumerate(cv_scores):\n        print(f\"Fold {fold + 1}: mAP = {score:.4f}\")\n    print(f\"CV average mAP: {np.mean(cv_scores):.4f} (+/- {np.std(cv_scores):.4f})\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:08:27.330346Z","iopub.execute_input":"2026-02-03T05:08:27.330602Z","iopub.status.idle":"2026-02-03T05:08:27.348222Z","shell.execute_reply.started":"2026-02-03T05:08:27.330580Z","shell.execute_reply":"2026-02-03T05:08:27.347635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#uncomment to train\n# train_main()\n\n!cp /kaggle/input/d/welshonionman/jaguarmodels/models/* /kaggle/tmp -r","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:19:05.468623Z","iopub.execute_input":"2026-02-03T05:19:05.469490Z","iopub.status.idle":"2026-02-03T05:19:23.220481Z","shell.execute_reply.started":"2026-02-03T05:19:05.469454Z","shell.execute_reply":"2026-02-03T05:19:23.219427Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Inference","metadata":{}},{"cell_type":"code","source":"class TTATestDataset(Dataset):\n    \"\"\"Test dataset with TTA support (horizontal flip)\"\"\"\n\n    def __init__(self, image_files: list[str], image_dir: str, img_size: int):\n        self.image_files = image_files\n        self.image_dir = image_dir\n        self.img_size = img_size\n\n        self.base_transform = A.Compose(\n            [\n                A.Resize(img_size, img_size),\n                A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n                ToTensorV2(),\n            ]\n        )\n\n        self.hflip_transform = A.Compose(\n            [\n                A.Resize(img_size, img_size),\n                A.HorizontalFlip(p=1.0),\n                A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n                ToTensorV2(),\n            ]\n        )\n\n    def __len__(self):\n        return len(self.image_files)\n\n    def __getitem__(self, idx):\n        filename = self.image_files[idx]\n        img_path = os.path.join(self.image_dir, filename)\n        image = np.array(Image.open(img_path).convert(\"RGB\"))\n\n        img_base = self.base_transform(image=image)[\"image\"]\n        img_hflip = self.hflip_transform(image=image)[\"image\"]\n\n        return img_base, img_hflip, filename\n\n\ndef extract_embeddings(model, loader, device):\n    \"\"\"Extract embeddings for all images (no TTA)\"\"\"\n    model.eval()\n    embeddings_dict = {}\n\n    with torch.no_grad():\n        for images, filenames in tqdm(loader, desc=\"Extracting embeddings\"):\n            images = images.to(device)\n            embeddings = model.get_embedding(images)\n\n            for emb, fname in zip(embeddings.cpu().numpy(), filenames):\n                embeddings_dict[fname] = emb\n\n    return embeddings_dict\n\n\ndef extract_embeddings_with_tta(model, loader, device):\n    \"\"\"Extract embeddings with TTA (horizontal flip)\"\"\"\n    model.eval()\n    embeddings_dict = {}\n\n    with torch.no_grad():\n        for img_base, img_hflip, filenames in tqdm(loader, desc=\"Extracting embeddings (TTA)\"):\n            img_base = img_base.to(device)\n            img_hflip = img_hflip.to(device)\n\n            # Original image embedding\n            emb_base = model.get_embedding(img_base)\n            # Horizontal flip embedding\n            emb_hflip = model.get_embedding(img_hflip)\n            # Average TTA embeddings and normalize\n            emb_tta = (emb_base + emb_hflip) / 2\n            emb_tta = F.normalize(emb_tta, dim=1)\n\n            for emb, fname in zip(emb_tta.cpu().numpy(), filenames):\n                embeddings_dict[fname] = emb\n\n    return embeddings_dict\n\n\ndef compute_similarity(emb1, emb2):\n    \"\"\"Compute cosine similarity between two embeddings\"\"\"\n    return np.dot(emb1, emb2)\n\n\ndef generate_submission(\n    embeddings_dict: dict[str, np.ndarray],\n    test_df: pd.DataFrame,\n    output_path: str,\n) -> pd.DataFrame:\n    \"\"\"Generate submission file\"\"\"\n    similarities = []\n    for _, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"Computing similarities\"):\n        query_emb = embeddings_dict[row[\"query_image\"]]\n        gallery_emb = embeddings_dict[row[\"gallery_image\"]]\n        sim = compute_similarity(query_emb, gallery_emb)\n        sim = (sim + 1) / 2\n        similarities.append(sim)\n    similarities = np.array(similarities)\n\n    submission = pd.DataFrame({\"row_id\": test_df[\"row_id\"], \"similarity\": similarities})\n\n    submission.to_csv(output_path, index=False)\n    print(f\"Submission saved: {output_path}\")\n    return submission","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:17:14.795243Z","iopub.execute_input":"2026-02-03T05:17:14.795623Z","iopub.status.idle":"2026-02-03T05:17:14.809964Z","shell.execute_reply.started":"2026-02-03T05:17:14.795588Z","shell.execute_reply":"2026-02-03T05:17:14.809389Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def strip_compiled_prefix(state_dict):\n    \"\"\"Remove _orig_mod. prefix from state_dict keys (added by torch.compile)\"\"\"\n    new_state_dict = {}\n    for key, value in state_dict.items():\n        if key.startswith(\"_orig_mod.\"):\n            new_key = key[len(\"_orig_mod.\") :]\n        else:\n            new_key = key\n        new_state_dict[new_key] = value\n    return new_state_dict\n\n\ndef inference_main(use_tta: bool = True):\n    \"\"\"\n    Run inference with 5-fold ensemble and optional TTA\n\n    Args:\n        use_tta: If True, use Test Time Augmentation (horizontal flip)\n    \"\"\"\n    print(f\"Device: {CFG.device}\")\n    print(f\"TTA: {use_tta}\")\n    set_seed(CFG.seed)\n\n    # Load data\n    test_df = pd.read_csv(os.path.join(CFG.input_dir, \"test.csv\"))\n    print(f\"Number of test pairs: {len(test_df)}\")\n\n    # Extract test embeddings with 5-fold ensemble\n    test_images = sorted(list(set(test_df[\"query_image\"].tolist() + test_df[\"gallery_image\"].tolist())))\n    print(f\"Number of unique test images: {len(test_images)}\")\n\n    # Collect embeddings from all folds\n    all_fold_embeddings = []\n\n    for fold in range(CFG.n_splits):\n        model_path = os.path.join(CFG.output_dir, f\"best_model_fold{fold}.pth\")\n        print(f\"\\nLoading model: {model_path}\")\n\n        model = JaguarReIDModel(\n            CFG.model_name, CFG.embedding_dim, CFG.num_classes, s=CFG.arcface_s, m=CFG.arcface_m, pretrained=False\n        )\n\n        # Load checkpoint (use EMA weights if available)\n        checkpoint = torch.load(model_path, map_location=CFG.device, weights_only=False)\n        if \"ema_state_dict\" in checkpoint:\n            state_dict = strip_compiled_prefix(checkpoint[\"ema_state_dict\"])\n            model.load_state_dict(state_dict)\n            print(f\"Loaded EMA weights (mAP: {checkpoint['best_mAP']:.4f})\")\n        else:\n            state_dict = strip_compiled_prefix(checkpoint[\"model_state_dict\"])\n            model.load_state_dict(state_dict)\n            print(f\"Loaded model weights (mAP: {checkpoint['best_mAP']:.4f})\")\n\n        model = model.to(CFG.device)\n        model.eval()\n\n        if use_tta:\n            # Use TTA dataset\n            test_dataset = TTATestDataset(test_images, os.path.join(CFG.input_dir, \"test\"), CFG.img_size)\n            test_loader = DataLoader(\n                test_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers, pin_memory=True\n            )\n            embeddings_dict = extract_embeddings_with_tta(model, test_loader, CFG.device)\n        else:\n            # Use standard dataset\n            test_dataset = TestImageDataset(\n                test_images, os.path.join(CFG.input_dir, \"test\"), transforms=get_valid_transforms(CFG.img_size)\n            )\n            test_loader = DataLoader(\n                test_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers, pin_memory=True\n            )\n            embeddings_dict = extract_embeddings(model, test_loader, CFG.device)\n\n        all_fold_embeddings.append(embeddings_dict)\n\n    # Average embeddings from all folds\n    print(\"\\nAveraging embeddings from all folds...\")\n    avg_embeddings_dict = {}\n    for img_name in test_images:\n        fold_embs = [fold_emb[img_name] for fold_emb in all_fold_embeddings]\n        avg_emb = np.mean(fold_embs, axis=0)\n        # Normalize the averaged embedding\n        avg_emb = avg_emb / np.linalg.norm(avg_emb)\n        avg_embeddings_dict[img_name] = avg_emb\n\n    # Generate submission\n    submission = generate_submission(\n        avg_embeddings_dict,\n        test_df,\n        \"/kaggle/working/submission.csv\",\n    )\n\n    print(\"\\nSubmission preview:\")\n    print(submission.head(10))\n    print(f\"\\nSubmission shape: {submission.shape}\")\n    print(f\"Similarity range: [{submission['similarity'].min():.4f}, {submission['similarity'].max():.4f}]\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:17:15.192752Z","iopub.execute_input":"2026-02-03T05:17:15.193083Z","iopub.status.idle":"2026-02-03T05:17:15.570315Z","shell.execute_reply.started":"2026-02-03T05:17:15.193055Z","shell.execute_reply":"2026-02-03T05:17:15.569553Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run inference with TTA\ninference_main(use_tta=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:19:30.563592Z","iopub.execute_input":"2026-02-03T05:19:30.564048Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T05:17:37.090443Z","iopub.execute_input":"2026-02-03T05:17:37.090764Z","iopub.status.idle":"2026-02-03T05:17:37.251705Z","shell.execute_reply.started":"2026-02-03T05:17:37.090738Z","shell.execute_reply":"2026-02-03T05:17:37.250694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\n\n# Delete all files in /kaggle/working except submission.csv and best_model_fold*.pth\nworking_dir = \"/kaggle/working\"\nif os.path.exists(working_dir):\n    for file_path in glob.glob(os.path.join(working_dir, \"*\")):\n        basename = os.path.basename(file_path)\n        # Keep submission.csv and all fold model files\n        if basename == \"submission.csv\" or basename.startswith(\"best_model_fold\"):\n            continue\n        if os.path.isfile(file_path):\n            os.remove(file_path)\n            print(f\"Deleted file: {file_path}\")\n        elif os.path.isdir(file_path):\n            import shutil\n\n            shutil.rmtree(file_path)\n            print(f\"Deleted directory: {file_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T04:47:03.088639Z","iopub.execute_input":"2026-02-03T04:47:03.088877Z","iopub.status.idle":"2026-02-03T04:47:03.094975Z","shell.execute_reply.started":"2026-02-03T04:47:03.088848Z","shell.execute_reply":"2026-02-03T04:47:03.094193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}