{"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":132732,"databundleVersionId":16583342,"isSourceIdPinned":false},{"sourceType":"modelInstanceVersion","sourceId":828885,"databundleVersionId":16634059,"modelInstanceId":630348,"modelId":642258,"isSourceIdPinned":false},{"sourceType":"modelInstanceVersion","sourceId":835748,"databundleVersionId":16735511,"modelInstanceId":635782,"modelId":647795,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# DLMMDD Baseline with ConvNeXt\n\nThis notebook uses a ConvNeXt model to classify faces.\n\nVersions 3 and 5 of this notebook used a 224x224 version of ConvNeXt-tiny; Version 6 uses a 384x384 version.\n\n- Training data: 7000 quadratic images of shape (1024, 1024) or (512, 512)\n- Test data: 3000 near-quadratic images of various resolutions -- 79 % are quadratic, 35 % have shape (1024, 1024). The smaller dimension is at least 80 % of the larger dimension.\n\nReference\n- https://www.kaggle.com/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge\n","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nfrom sklearn.model_selection import StratifiedKFold\nimport matplotlib.pyplot as plt\nimport pickle\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.transforms import v2\nimport timm\n\nis_interactive = os.environ['KAGGLE_KERNEL_RUN_TYPE'] == 'Interactive'\ntqdm_if_int = tqdm if is_interactive else lambda x, **kwargs: x\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:23:55.863689Z","iopub.execute_input":"2026-04-17T14:23:55.864410Z","iopub.status.idle":"2026-04-17T14:24:13.675034Z","shell.execute_reply.started":"2026-04-17T14:23:55.864369Z","shell.execute_reply":"2026-04-17T14:24:13.674078Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    data_dir = '/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/'\n    train_csv = data_dir + \"Data/training.csv\"\n    test_csv = data_dir + \"Data/test.csv\"\n    # model_name = \"convnext_tiny.fb_in22k_ft_in1k\" # 114MByte download, 28.6M parameters\n    # pretrained_cfg_overlay = {'file': '/kaggle/input/models/ambrosm/convnext-tiny-fb-in22k-ft-in1k/pytorch/default/1/model.safetensors'}\n    model_name = \"convnext_tiny.fb_in22k_ft_in1k_384\" # 114MByte download, 28.6M parameters\n    pretrained_cfg_overlay = {'file': '/kaggle/input/models/ambrosm/convnext-tiny-fb-in22k-ft-in1k-384/pytorch/default/1/model.safetensors'}\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    batch_size = 64\n    T_0 = 4\n    epochs = T_0\n    lr = 3e-4\n    label_smoothing = 0.1\n    num_classes = 10\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:24:13.676651Z","iopub.execute_input":"2026-04-17T14:24:13.677397Z","iopub.status.idle":"2026-04-17T14:24:13.685585Z","shell.execute_reply.started":"2026-04-17T14:24:13.677357Z","shell.execute_reply":"2026-04-17T14:24:13.684373Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"complete_train_df = pd.read_csv(CFG.train_csv)\ntest_df = pd.read_csv(CFG.test_csv)\n\nif is_interactive:\n    complete_train_df, test_df = complete_train_df.iloc[::10], test_df.iloc[::10] # Debug with a subset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:24:13.687288Z","iopub.execute_input":"2026-04-17T14:24:13.687695Z","iopub.status.idle":"2026-04-17T14:24:13.754282Z","shell.execute_reply.started":"2026-04-17T14:24:13.687664Z","shell.execute_reply":"2026-04-17T14:24:13.753190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot samples from train and test\nfor label, df in [('Train', complete_train_df), ('Test', test_df)]:\n    _, axs = plt.subplots(3, 7, figsize=(12, 6))\n    for i in range(len(axs.ravel())):\n        r = df.iloc[i + 20]\n        img = Image.open(CFG.data_dir + r[\"path\"]).convert(\"RGB\")\n        axs.ravel()[i].imshow(img)\n        axs.ravel()[i].axis('off')\n    plt.suptitle(label, fontsize=24)\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:24:13.756883Z","iopub.execute_input":"2026-04-17T14:24:13.757290Z","iopub.status.idle":"2026-04-17T14:24:20.604526Z","shell.execute_reply.started":"2026-04-17T14:24:13.757255Z","shell.execute_reply":"2026-04-17T14:24:20.603321Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Plot samples from every class\n# for cls in range(CFG.num_classes):\n#     _, axs = plt.subplots(6, 7, figsize=(12, 12))\n#     for i in range(len(axs.ravel())):\n#         r = complete_train_df.query(\"y == @cls\").iloc[i]\n#         img = Image.open(CFG.data_dir + r[\"path\"]).convert(\"RGB\")\n#         axs.ravel()[i].imshow(img)\n#         axs.ravel()[i].axis('off')\n#     plt.suptitle(f\"Class {cls}\", fontsize=24)\n#     plt.tight_layout()\n#     plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:24:20.606359Z","iopub.execute_input":"2026-04-17T14:24:20.606693Z","iopub.status.idle":"2026-04-17T14:24:20.612124Z","shell.execute_reply.started":"2026-04-17T14:24:20.606661Z","shell.execute_reply":"2026-04-17T14:24:20.611047Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MyDataset(Dataset):\n    \n    def __init__(self, ds_type, df, return_label, model):\n        \"\"\"Image dataset.\n\n        Returns square images of the correct resolution for model.\n\n        Parameters\n        ds_type: 'train', 'val', or 'test'\n        df: either a subset of training.csv or test.csv\n        return_label: True for train and val, False for test\n        model: a timm model, which defines image size and normalization parameters\n        \"\"\"\n        assert ds_type in ['train', 'val', 'test']\n        self.ds_type = ds_type\n        self.df = df\n        self.base_path = CFG.data_dir\n        self.return_label = return_label\n        self.shapes = []\n        self.input_size = model.pretrained_cfg['input_size']\n        self.tfm0 = v2.Compose([\n            v2.ToImage(), # converts PIL to tensor\n            v2.ToDtype(torch.float32, scale=True),\n            v2.Normalize(mean=model.pretrained_cfg['mean'], std=model.pretrained_cfg['std'])\n        ])\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        r = self.df.iloc[i]\n        img = Image.open(self.base_path + r[\"path\"]).convert(\"RGB\")\n        self.shapes.append((img.width, img.height))\n        img = img.resize(self.input_size[-2:])\n        if self.ds_type == 'train':\n            if i % 3 == 0:\n                img = v2.GaussianBlur(kernel_size=5)(img)\n            img = v2.RandomAutocontrast(p=0.25)(img)\n        elif self.ds_type == 'val':\n            if i % 3 == 0:\n                img = v2.functional.gaussian_blur(img, kernel_size=5, sigma=1)\n            if i % 4 == 0:\n                img = v2.functional.autocontrast(img)\n        return self.tfm0(img), torch.tensor(r[\"y\"]).long() if self.return_label else r[\"ID\"]\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:24:20.613564Z","iopub.execute_input":"2026-04-17T14:24:20.613986Z","iopub.status.idle":"2026-04-17T14:24:20.639263Z","shell.execute_reply.started":"2026-04-17T14:24:20.613948Z","shell.execute_reply":"2026-04-17T14:24:20.638049Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training and inference","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, scaler):\n    \"\"\"Train the model for one epoch.\n    \n    Return cross-entropy loss and accuracy.\n    \"\"\"\n    model.train()\n    loss_fn = nn.CrossEntropyLoss(label_smoothing=CFG.label_smoothing)\n    \n    total_loss = 0\n    correct = 0\n    total = 0\n\n    pbar = tqdm_if_int(loader, desc=\"Train\")\n\n    for i, (img_batch, label_batch) in enumerate(pbar):\n        img_batch, label_batch = img_batch.to(CFG.device, memory_format=torch.channels_last), label_batch.to(CFG.device)\n\n        optimizer.zero_grad()\n        with torch.amp.autocast(CFG.device, enabled=CFG.device=='cuda'):\n            logits = model(img_batch)\n            loss = loss_fn(logits, label_batch)\n        \n        # loss.backward()\n        # optimizer.step()\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item() * img_batch.size(0)\n\n        preds = logits.argmax(1)\n        correct += (preds == label_batch).sum().item()\n        total += label_batch.size(0)\n\n        if is_interactive:\n            pbar.set_postfix({\n                \"loss\": total_loss / total,\n                \"acc\": correct / total\n            })\n\n    return total_loss / total, correct / total\n\ndef validate(model, loader):\n    \"\"\"Validate the model.\n\n    Return cross-entropy loss, accuracy and logits.\n    \"\"\"\n    model.eval()\n    loss_fn = nn.CrossEntropyLoss(label_smoothing=CFG.label_smoothing)\n\n    total_loss = 0\n    correct = 0\n    total = 0\n    logits = []\n\n    pbar = tqdm_if_int(loader, desc=\"Valid\")\n\n    with torch.no_grad():\n        for img_batch, label_batch in pbar:\n            img_batch, label_batch = img_batch.to(CFG.device, memory_format=torch.channels_last), label_batch.to(CFG.device)\n\n            logits_batch = model(img_batch)\n            loss = loss_fn(logits_batch, label_batch)\n\n            total_loss += loss.item() * img_batch.size(0)\n\n            preds = logits_batch.argmax(1)\n            logits.append(logits_batch.cpu().numpy())\n            correct += (preds == label_batch).sum().item()\n            total += label_batch.size(0)\n\n            if is_interactive:\n                pbar.set_postfix({\n                    \"val_loss\": total_loss / total,\n                    \"val_acc\": correct / total\n                })\n\n    return total_loss / total, correct / total, np.vstack(logits)\n\ndef predict(model, loader):\n    \"\"\"Classify the images using model.\n    \n    Return a list of ids and a logit array of shape (n_samples, n_classes).\n    \"\"\"\n    model.eval()\n    ids, logits = [], []\n    with torch.no_grad():\n        for img_batch, id_batch in tqdm_if_int(loader, desc=\"Test\"):\n            img_batch = img_batch.to(CFG.device, memory_format=torch.channels_last)\n            logits_batch = model(img_batch)\n            ids.extend(id_batch.numpy())\n            logits.append(logits_batch.cpu().numpy())\n    return ids, np.vstack(logits)\n    ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:24:20.640403Z","iopub.execute_input":"2026-04-17T14:24:20.640793Z","iopub.status.idle":"2026-04-17T14:24:20.668538Z","shell.execute_reply.started":"2026-04-17T14:24:20.640760Z","shell.execute_reply":"2026-04-17T14:24:20.667457Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Device: {CFG.device}   Model: {CFG.model_name}\")\n\nall_preds = []\noof_logits = np.full((len(complete_train_df), CFG.num_classes), np.nan)\nkf = StratifiedKFold(n_splits=5, shuffle=True, random_state=1)\nfor fold, (idx_tr, idx_va) in enumerate(kf.split(complete_train_df, complete_train_df['y'])):\n    train_df = complete_train_df.iloc[idx_tr]\n    val_df = complete_train_df.iloc[idx_va]\n\n    # pretrained_cfg_overlay loads the weights from the Kaggle model, which is more reliable than the download from Huggingface.\n    # The weights are the same as on Huggingface.\n    model = timm.create_model(\n        model_name=CFG.model_name,\n        pretrained_cfg_overlay=CFG.pretrained_cfg_overlay,\n        pretrained=True,\n        num_classes=CFG.num_classes\n    ).to(CFG.device)\n    model = model.to(memory_format=torch.channels_last)\n    \n    if fold == 0:\n        print(timm.data.resolve_data_config(model.pretrained_cfg))\n    \n    train_loader = DataLoader(\n        MyDataset('train', train_df, return_label=True, model=model),\n        batch_size=CFG.batch_size,\n        shuffle=True,\n        num_workers=4,\n        pin_memory=CFG.device=='cuda'\n    )\n    \n    val_loader = DataLoader(\n        MyDataset('val', val_df, return_label=True, model=model),\n        batch_size=CFG.batch_size,\n        shuffle=False,\n        num_workers=4,\n        pin_memory=CFG.device=='cuda'\n    )\n    \n    test_loader = DataLoader(\n        MyDataset('test', test_df, return_label=False, model=model),\n        batch_size=CFG.batch_size,\n        shuffle=False,\n        num_workers=4,\n        pin_memory=CFG.device=='cuda'\n    )\n\n    optimizer = optim.AdamW(model.parameters(), lr=CFG.lr)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer,\n        T_0=CFG.T_0,\n        eta_min=1e-6\n    )\n    scaler = torch.amp.GradScaler(CFG.device, enabled=CFG.device=='cuda')\n    \n    # Train\n    best_acc = 0.0\n    for epoch in range(CFG.epochs):\n        # Train one epoch\n        print(f\"\\nFold {fold}, epoch {epoch+1}/{CFG.epochs}: lr={scheduler.get_last_lr()[0]:.3e}\")\n        train_loss, train_acc = train_one_epoch(model, train_loader, optimizer, scaler)\n        print(f\"Train: loss={train_loss:.3f}, acc={train_acc:.3f}\")\n        scheduler.step()\n\n        # Validate after every epoch\n        val_loss, val_acc, logits = validate(model, val_loader)\n        oof_logits[idx_va] = logits\n        print(f\"Valid:                                 val_loss={val_loss:.3f}, val_acc={val_acc:.3f}\")\n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), f\"best_model{fold}.pth\")\n    \n    # Predict\n    model.load_state_dict(torch.load(f\"best_model{fold}.pth\"))\n    ids, preds = predict(model, test_loader)\n    assert (test_df[\"ID\"] == ids).all()\n    all_preds.append(preds)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:24:20.669690Z","iopub.execute_input":"2026-04-17T14:24:20.669974Z","execution_failed":"2026-04-17T14:25:18.668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(6, 3))\nplt.title('Predicted classes (oof)')\nplt.bar(np.arange(CFG.num_classes), np.bincount(np.argmax(oof_logits, axis=1)), color='chocolate')\nplt.xlabel('class')\nplt.ylabel('count')\nplt.xticks(np.arange(CFG.num_classes))\nplt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-17T14:25:18.670Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"# Ensemble the logits\nlogits = np.mean(all_preds, axis=0)\npreds = np.argmax(logits, axis=1)\n\n# Save the submission file\npd.DataFrame({\"ID\": test_df[\"ID\"], \"TARGET\": preds}).to_csv(\"submission.csv\", index=False)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-17T14:25:18.670Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(6, 3))\nplt.title('Predicted classes (test)')\nplt.bar(np.arange(CFG.num_classes), np.bincount(preds))\nplt.xlabel('class')\nplt.ylabel('count')\nplt.xticks(np.arange(CFG.num_classes))\nplt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-17T14:25:18.670Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open('all_preds.pickle', 'wb') as f:\n    pickle.dump(all_preds, f)","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-17T14:25:18.671Z"}},"outputs":[],"execution_count":null}]}