{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":"none","dataSources":[{"sourceId":126777,"databundleVersionId":15314950,"sourceType":"competition"},{"sourceId":733318,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":558848,"modelId":571424}],"dockerImageVersionId":31259,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Jaguar Re-ID: Training Notebook\n\n## Interface Notebook is here! **https://www.kaggle.com/code/kawaharataishi/jaguar-re-id-starter-set-inference-notebook**\nIt scores 0.866\n\n## Overview\nThis notebook trains a model for the Jaguar Re-Identification Challenge.\n\n## Approach\n1. **Model**: `ConvNeXt Base` (Pretrained on ImageNet)\n2. **Loss Function**: `ArcFace` (Metric Learning) - Minimizes intra-class variance and maximizes inter-class variance.\n   - Margin: 0.50, Scale: 30.0\n3. **Data**: Uses ALL available Jaguar data (No validation split).\n   - Reason: Given the small dataset size, we use all data for training to improve generalization performance.\n4. **Augmentation**: Resize, RandomHorizontalFlip, RandomRotation, ColorJitter\n5. **Resolution**: 384x384 \n\n## Output\nWhen training is complete, a `.pth` file (model weights) is saved.\nDownload this file and upload it as a Kaggle Dataset to be used in the Inference Notebook. \nHere is the model. **/kaggle/input/m/kawaharataishi/jaguar-re-id/pytorch/default/1**","metadata":{}},{"cell_type":"code","source":"import os\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\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\nimport torchvision.transforms as transforms\nimport torchvision.models as models\nimport timm\n\n# ==========================================\n# Configuration\n# ==========================================\n\nclass Config:\n    DATA_DIR = '/kaggle/input/jaguar-re-id' \n        \n    TRAIN_CSV = os.path.join(DATA_DIR, 'train.csv')\n    TRAIN_IMG_DIR = os.path.join(DATA_DIR, 'train', 'train')\n\n    # Model Config\n    IMG_SIZE = (384, 384) \n    BATCH_SIZE = 16 \n    EPOCHS = 15 \n    LEARNING_RATE = 1e-4\n    WEIGHT_DECAY = 1e-4\n    EMBEDDING_DIM = 512\n    \n    # ArcFace Config\n    S = 30.0\n    M = 0.50\n    \n    # Device\n    if torch.cuda.is_available():\n        DEVICE = torch.device('cuda')\n    elif torch.backends.mps.is_available():\n        DEVICE = torch.device('mps') # for Mac\n    else:\n        DEVICE = torch.device('cpu')\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(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(42)\nprint(f\"Device: {Config.DEVICE}\")","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset Definition","metadata":{}},{"cell_type":"code","source":"def crop_alphachannel(img: Image.Image) -> Image.Image:\n    # Remove extra whitespace using the Alpha channel\n    if img.mode in ('RGBA', 'LA') or (img.mode == 'P' and 'transparency' in img.info):\n        bbox = img.getbbox()\n        if bbox:\n            return img.crop(bbox)\n    return img\n\nclass JaguarDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None, is_train=True, label_map=None):\n        self.df = df\n        self.img_dir = img_dir\n        self.transform = transform\n        self.is_train = is_train\n        self.label_map = 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_path = os.path.join(self.img_dir, row['filename'])\n\n        try:\n            image = Image.open(img_path).convert(\"RGBA\")\n            image = crop_alphachannel(image)\n            image = image.convert(\"RGB\") \n        except:\n             image = Image.new(\"RGB\", Config.IMG_SIZE)\n\n        if self.transform:\n            image = self.transform(image)\n\n        if self.is_train:\n            label_str = row['ground_truth']\n            label = self.label_map[label_str]\n            return image, torch.tensor(label, dtype=torch.long)\n        else:\n            return image, row['filename']\n\n# Data Loading\ntrain_df_raw = pd.read_csv(Config.TRAIN_CSV)\nunique_ids = sorted(train_df_raw['ground_truth'].unique())\nlabel_map = {name: i for i, name in enumerate(unique_ids)}\nNUM_CLASSES = len(unique_ids)\n\nprint(f\"Number of Classes: {NUM_CLASSES}\")\n\n# Transforms\ntrain_transform = transforms.Compose([\n    transforms.Resize(Config.IMG_SIZE),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1, hue=0.1),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\n# Configured to use ALL data (No validation split)\ntrain_dataset = JaguarDataset(train_df_raw, Config.TRAIN_IMG_DIR, transform=train_transform, is_train=True, label_map=label_map)\ntrain_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=2, pin_memory=True)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Definition (ArcFace + ConvNeXt)","metadata":{}},{"cell_type":"code","source":"class ArcMarginProduct(nn.Module):\n    def __init__(self, in_features, out_features, s=30.0, m=0.50, easy_margin=False):\n        super(ArcMarginProduct, self).__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        self.easy_margin = easy_margin\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, input, label):\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        sine = torch.sqrt((1.0 - torch.pow(cosine, 2)).clamp(0, 1))\n        phi = cosine * self.cos_m - sine * self.sin_m\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(cosine.size(), device=input.device)\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s\n        return output\n\nclass JaguarReIDModel(nn.Module):\n    def __init__(self, num_classes, embedding_dim=512, pretrained=True):\n        super(JaguarReIDModel, self).__init__()\n        # Backbone: ConvNeXt Base\n        self.backbone = timm.create_model('convnext_base', pretrained=pretrained, num_classes=0)\n        in_features = self.backbone.num_features\n        \n        # Head\n        self.embedding = nn.Sequential(\n            nn.Linear(in_features, embedding_dim),\n            nn.BatchNorm1d(embedding_dim),\n            nn.PReLU()\n        )\n\n        self.arcface = ArcMarginProduct(embedding_dim, num_classes, s=Config.S, m=Config.M)\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            return self.arcface(embeddings, labels)\n        else:\n            return F.normalize(embeddings)\n\nmodel = JaguarReIDModel(NUM_CLASSES, Config.EMBEDDING_DIM).to(Config.DEVICE)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Loop","metadata":{}},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=Config.LEARNING_RATE, weight_decay=Config.WEIGHT_DECAY)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n\ndef train_one_epoch(model, loader, optimizer, criterion, device):\n    model.train()\n    total_loss = 0\n    pbar = tqdm(loader, desc=\"Training\")\n\n    for images, labels in pbar:\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(images, labels)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n        pbar.set_postfix({'loss': loss.item()})\n\n    return total_loss / len(loader)\n\nprint(\"Start Training...\")\nfor epoch in range(Config.EPOCHS):\n    avg_loss = train_one_epoch(model, train_loader, optimizer, criterion, Config.DEVICE)\n    scheduler.step()\n    print(f\"Epoch {epoch+1}/{Config.EPOCHS} - Loss: {avg_loss:.4f}\")\n    \n    if (epoch + 1) % 5 == 0:\n        save_name = f'jaguar_resnet50_arcface_ep{epoch+1}.pth'\n        torch.save(model.state_dict(), save_name)\n        print(f\"Saved checkpoint: {save_name}\")\n\nprint(\"Training Finished.\")","metadata":{},"outputs":[],"execution_count":null}]}