{"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}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nBaseline Solution: Synthetic Image Attribution Challenge\n=========================================================\nЦель: быстрый базовый запуск за 1-2 часа.\nМодель: EfficientNet-B0 (pretrained) → fine-tune на 10 классов.\n\"\"\"\n\n# ── 0. Установка зависимостей ──────────────────────────────────────────────────\n# pip install timm torch torchvision pandas pillow scikit-learn tqdm albumentations\n\nimport os\nimport random\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\n\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\n\n# ── 1. Конфиг ─────────────────────────────────────────────────────────────────\nclass CFG:\n    DATA_DIR = Path(\"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data\")          # папка с данными (измените при необходимости)\n    TRAIN_CSV     = DATA_DIR / \"training.csv\"\n    TEST_CSV      = DATA_DIR / \"test.csv\"\n    SAMPLE_SUB    = DATA_DIR / \"sample_submission.csv\"\n    SOURCES_TXT   = DATA_DIR / \"sources.txt\"\n\n    MODEL_NAME    = \"efficientnet_b0\"       # быстрая модель\n    NUM_CLASSES   = 10\n    IMG_SIZE      = 224\n    BATCH_SIZE    = 32\n    NUM_EPOCHS    = 9                       # 5 эпох ≈ 30-60 мин на CPU/слабом GPU\n    LR            = 1e-4\n    WEIGHT_DECAY  = 1e-4\n    VAL_SPLIT     = 0.14\n    NUM_WORKERS   = 4\n    SEED          = 42\n    DEVICE        = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    OUTPUT_DIR    = Path(\"/kaggle/working/\")\n    CHECKPOINT    = OUTPUT_DIR / \"best_model.pth\"\n    SUBMISSION    = OUTPUT_DIR / \"submission2.csv\"\n\n\nCFG.OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n\n\ndef seed_everything(seed: int):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_everything(CFG.SEED)\nprint(f\"Device: {CFG.DEVICE}\")\n\n\n# ── 2. Загрузка данных ────────────────────────────────────────────────────────\ntrain_df = pd.read_csv(CFG.TRAIN_CSV)\ntest_df  = pd.read_csv(CFG.TEST_CSV)\n\n# Посмотрим на структуру\nprint(\"Train columns:\", train_df.columns.tolist())\nprint(\"Test  columns:\", test_df.columns.tolist())\nprint(\"Train shape:\", train_df.shape, \"| Test shape:\", test_df.shape)\nprint(\"\\nClass distribution:\\n\", train_df[\"y\"].value_counts().sort_index())\n\n# Источники (опционально)\nif CFG.SOURCES_TXT.exists():\n    with open(CFG.SOURCES_TXT) as f:\n        print(\"\\nSources:\\n\", f.read())\n\n# Разбивка train / val\ntrain_data, val_data = train_test_split(\n    train_df,\n    test_size=CFG.VAL_SPLIT,\n    stratify=train_df[\"y\"],\n    random_state=CFG.SEED,\n)\ntrain_data = train_data.reset_index(drop=True)\nval_data   = val_data.reset_index(drop=True)\nprint(f\"\\nTrain: {len(train_data)} | Val: {len(val_data)} | Test: {len(test_df)}\")\n\n\n# ── 3. Dataset и аугментации ──────────────────────────────────────────────────\ntrain_transforms = T.Compose([\n    T.Resize((CFG.IMG_SIZE, CFG.IMG_SIZE)),\n    T.RandomHorizontalFlip(),\n    T.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1),\n    T.ToTensor(),\n    T.Normalize(mean=[0.485, 0.456, 0.406],\n                std =[0.229, 0.224, 0.225]),\n])\n\nval_transforms = T.Compose([\n    T.Resize((CFG.IMG_SIZE, CFG.IMG_SIZE)),\n    T.ToTensor(),\n    T.Normalize(mean=[0.485, 0.456, 0.406],\n                std =[0.229, 0.224, 0.225]),\n])\n\n\nclass SIADataset(Dataset):\n    \"\"\"Synthetic Image Attribution Dataset.\"\"\"\n\n    def __init__(self, df: pd.DataFrame, data_dir: Path,\n                 transforms=None, is_test: bool = False):\n        self.df         = df\n        self.data_dir   = data_dir\n        self.transforms = transforms\n        self.is_test    = is_test\n        # Колонка с меткой: 'y' (train) или отсутствует (test)\n        self.label_col  = \"y\" if not is_test else None\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n    \n        # путь из csv\n        img_path = Path(row[\"path\"])\n    \n        # если путь относительный\n        if not img_path.is_absolute():\n    \n            # убираем лишний \"Data/\"\n            img_path_str = str(img_path)\n    \n            if img_path_str.startswith(\"Data/\"):\n                img_path_str = img_path_str.replace(\"Data/\", \"\", 1)\n    \n            img_path = self.data_dir / img_path_str\n    \n        # загрузка изображения\n        try:\n            img = Image.open(img_path).convert(\"RGB\")\n    \n        except Exception as e:\n            print(f\"Warning: cannot open {img_path}: {e}\")\n    \n            img = Image.new(\n                \"RGB\",\n                (CFG.IMG_SIZE, CFG.IMG_SIZE)\n            )\n    \n        if self.transforms:\n            img = self.transforms(img)\n    \n        if self.is_test:\n            return img, row[\"ID\"]\n    \n        label = int(row[self.label_col])\n    \n        return img, label\n\n\ntrain_dataset = SIADataset(train_data, CFG.DATA_DIR, train_transforms)\nval_dataset   = SIADataset(val_data,   CFG.DATA_DIR, val_transforms)\ntest_dataset  = SIADataset(test_df,    CFG.DATA_DIR, val_transforms, is_test=True)\n\ntrain_loader = DataLoader(train_dataset, batch_size=CFG.BATCH_SIZE,\n                          shuffle=True,  num_workers=CFG.NUM_WORKERS, pin_memory=True)\nval_loader   = DataLoader(val_dataset,   batch_size=CFG.BATCH_SIZE,\n                          shuffle=False, num_workers=CFG.NUM_WORKERS, pin_memory=True)\ntest_loader  = DataLoader(test_dataset,  batch_size=CFG.BATCH_SIZE,\n                          shuffle=False, num_workers=CFG.NUM_WORKERS, pin_memory=True)\n\nprint(f\"Batches — train: {len(train_loader)} | val: {len(val_loader)} | test: {len(test_loader)}\")\n\n\n# ── 4. Модель ─────────────────────────────────────────────────────────────────\ndef build_model(model_name: str, num_classes: int) -> nn.Module:\n    \"\"\"EfficientNet-B0 pretrained + custom head.\"\"\"\n    model = timm.create_model(model_name, pretrained=True, num_classes=num_classes)\n    return model\n\n\nmodel = build_model(CFG.MODEL_NAME, CFG.NUM_CLASSES).to(CFG.DEVICE)\ntotal_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"Trainable params: {total_params:,}\")\n\n\n# ── 5. Loss, Optimizer, Scheduler ─────────────────────────────────────────────\ncriterion = nn.CrossEntropyLoss(label_smoothing=0.05)\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG.LR, weight_decay=CFG.WEIGHT_DECAY)\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=CFG.LR * 10,\n    steps_per_epoch=len(train_loader),\n    epochs=CFG.NUM_EPOCHS,\n    pct_start=0.1,\n)\n\n\n# ── 6. Train / Eval loop ───────────────────────────────────────────────────────\ndef train_one_epoch(model, loader, optimizer, scheduler, criterion, device):\n    model.train()\n    total_loss, correct, total = 0.0, 0, 0\n\n    pbar = tqdm(loader, desc=\"  Train\", leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        logits = model(imgs)\n        loss   = criterion(logits, labels)\n        loss.backward()\n        nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        scheduler.step()\n\n        total_loss += loss.item() * len(labels)\n        preds       = logits.argmax(dim=1)\n        correct    += (preds == labels).sum().item()\n        total      += len(labels)\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\", acc=f\"{correct/total:.4f}\")\n\n    return total_loss / total, correct / total\n\n\n@torch.no_grad()\ndef evaluate(model, loader, criterion, device):\n    model.eval()\n    total_loss, correct, total = 0.0, 0, 0\n\n    for imgs, labels in tqdm(loader, desc=\"  Val  \", leave=False):\n        imgs, labels = imgs.to(device), labels.to(device)\n        logits = model(imgs)\n        loss   = criterion(logits, labels)\n\n        total_loss += loss.item() * len(labels)\n        preds       = logits.argmax(dim=1)\n        correct    += (preds == labels).sum().item()\n        total      += len(labels)\n\n    return total_loss / total, correct / total\n\n\n# ── 7. Обучение ────────────────────────────────────────────────────────────────\nbest_val_acc = 0.0\nhistory = []\n\nfor epoch in range(1, CFG.NUM_EPOCHS + 1):\n    print(f\"\\nEpoch {epoch}/{CFG.NUM_EPOCHS}  lr={scheduler.get_last_lr()[0]:.2e}\")\n\n    train_loss, train_acc = train_one_epoch(\n        model, train_loader, optimizer, scheduler, criterion, CFG.DEVICE\n    )\n    val_loss, val_acc = evaluate(model, val_loader, criterion, CFG.DEVICE)\n\n    history.append({\"epoch\": epoch, \"train_loss\": train_loss, \"train_acc\": train_acc,\n                    \"val_loss\": val_loss, \"val_acc\": val_acc})\n    print(f\"  train_loss={train_loss:.4f} | train_acc={train_acc:.4f} \"\n          f\"| val_loss={val_loss:.4f} | val_acc={val_acc:.4f}\")\n\n    if val_acc > best_val_acc:\n        best_val_acc = val_acc\n        torch.save(model.state_dict(), CFG.CHECKPOINT)\n        print(f\"  ✓ Saved best model (val_acc={best_val_acc:.4f})\")\n\nprint(f\"\\nBest val accuracy: {best_val_acc:.4f}\")\npd.DataFrame(history).to_csv(CFG.OUTPUT_DIR / \"history.csv\", index=False)\n\n\n# ── 8. Инференс на тестовом наборе ────────────────────────────────────────────\n# Загружаем лучшую модель\nmodel.load_state_dict(torch.load(CFG.CHECKPOINT, map_location=CFG.DEVICE))\nmodel.eval()\n\nall_ids, all_preds = [], []\n\nwith torch.no_grad():\n    for imgs, ids in tqdm(test_loader, desc=\"Inference\"):\n        imgs   = imgs.to(CFG.DEVICE)\n        logits = model(imgs)\n        preds  = logits.argmax(dim=1).cpu().numpy()\n        all_preds.extend(preds.tolist())\n        # ids может быть тензором или списком строк\n        if isinstance(ids, torch.Tensor):\n            all_ids.extend(ids.numpy().tolist())\n        else:\n            all_ids.extend(list(ids))\n\nprint(f\"Collected {len(all_ids)} predictions for {len(test_df)} test samples.\")\n\n\n# ── 9. Создание сабмита ────────────────────────────────────────────────────────\nsubmission = pd.DataFrame({\"ID\": all_ids, \"TARGET\": all_preds})\n\n# Проверяем порядок по sample_submission\nif CFG.SAMPLE_SUB.exists():\n    sample = pd.read_csv(CFG.SAMPLE_SUB)\n    submission = sample[[\"ID\"]].merge(submission, on=\"ID\", how=\"left\")\n    missing = submission[\"TARGET\"].isna().sum()\n    if missing > 0:\n        print(f\"Warning: {missing} missing predictions — filling with 0\")\n        submission[\"TARGET\"] = submission[\"TARGET\"].fillna(0).astype(int)\n    else:\n        submission[\"TARGET\"] = submission[\"TARGET\"].astype(int)\n\nsubmission.to_csv(CFG.SUBMISSION, index=False)\nprint(f\"\\nSubmission saved → {CFG.SUBMISSION}\")\nprint(submission.head(10))\nprint(\"Class distribution in predictions:\\n\", submission[\"TARGET\"].value_counts().sort_index())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T13:03:27.731059Z","iopub.execute_input":"2026-05-15T13:03:27.731948Z","iopub.status.idle":"2026-05-15T13:20:21.177272Z","shell.execute_reply.started":"2026-05-15T13:03:27.731891Z","shell.execute_reply":"2026-05-15T13:20:21.176500Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nAdvanced Baseline: Synthetic Image Attribution Challenge\n=========================================================\nМодель    : Swin-Transformer-Tiny  (timm)\nLoss      : ArcFace  +  Focal Loss (автовыбор по дисбалансу классов)\nАугменты  : albumentations pipeline  +  CutMix  +  MixUp  (batch-level)\nExtras    : визуализация 5 изображений двух разных классов\n\"\"\"\n\n# ── 0. Зависимости ─────────────────────────────────────────────────────────────\n# pip install timm albumentations torch torchvision pandas pillow scikit-learn tqdm matplotlib\n\nimport math, os, random\nimport numpy as np\nimport pandas as pd\nimport matplotlib\nmatplotlib.use(\"Agg\")          # headless на Kaggle\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 1. КОНФИГ\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef _find_data_dir() -> Path:\n    candidates = [\n        Path(\"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data\"),\n        Path(\"/kaggle/input/dlmmdd-workshop-synthetic-source-attribution-challenge\"),\n        Path(\"/kaggle/input\"),\n        Path(\"./data\"),\n    ]\n    for c in candidates:\n        if c.exists() and any(c.rglob(\"train.csv\")):\n            return c\n    return Path(\"./data\")\n\ndef _find_file(root: Path, name: str) -> Path:\n    hits = list(root.rglob(name))\n    return hits[0] if hits else root / name\n\n#_ROOT = _find_data_dir()\n\n        # папка с данными (измените при необходимости)\n\nclass CFG:\n\n    # ── Пути ─────────────────────────────────────────────────────\n    DATA_DIR    = Path(\"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data\")\n\n    TRAIN_CSV   = DATA_DIR / \"training.csv\"\n    TEST_CSV    = DATA_DIR / \"test.csv\"\n    SAMPLE_SUB  = DATA_DIR / \"sample_submission.csv\"\n    SOURCES_TXT = DATA_DIR / \"sources.txt\"\n\n    OUTPUT_DIR  = Path(\"/kaggle/working/\")\n\n    # ── Модель ──────────────────────────────────────────────────\n    MODEL_NAME  = \"convnext_small\"\n\n    NUM_CLASSES = 10\n\n    EMBED_DIM   = 512\n\n    # ── Обучение ────────────────────────────────────────────────\n    IMG_SIZE    = 224\n\n    BATCH_SIZE  = 16\n\n    NUM_EPOCHS  = 5\n\n    LR          = 1e-4\n\n    WEIGHT_DECAY = 5e-2\n\n    VAL_SPLIT   = 0.15\n\n    NUM_WORKERS = 4\n\n    SEED        = 42\n\n    DEVICE      = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    # ── ArcFace ─────────────────────────────────────────────────\n    ARC_S       = 32.0\n\n    ARC_M       = 0.4\n\n    # ── Focal Loss ──────────────────────────────────────────────\n    FOCAL_GAMMA = 2.0\n\n    # ── Augmentations ───────────────────────────────────────────\n    CUTMIX_PROB = 0.3\n\n    MIXUP_PROB  = 0.2\n\n    MIXUP_ALPHA = 0.4\n\n    # ── EMA ─────────────────────────────────────────────────────\n    USE_EMA     = True\n\n    # ── Derived ─────────────────────────────────────────────────\n    CHECKPOINT  = OUTPUT_DIR / \"best_convnext.pth\"\n\n    SUBMISSION  = OUTPUT_DIR / \"submission.csv\"\n\n    VIZ_PATH    = OUTPUT_DIR / \"class_samples.png\"\n\n\nCFG.OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(f\"Data root : {CFG.DATA_DIR}\")\nprint(f\"Device    : {CFG.DEVICE}\")\n\n\ndef seed_everything(seed: int):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_everything(CFG.SEED)\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 2. ДАННЫЕ\n# ══════════════════════════════════════════════════════════════════════════════\n\ntrain_df = pd.read_csv(CFG.TRAIN_CSV)\ntest_df  = pd.read_csv(CFG.TEST_CSV)\n\nprint(f\"Train: {train_df.shape}  Test: {test_df.shape}\")\nprint(\"Class distribution:\\n\", train_df[\"y\"].value_counts().sort_index())\n\n# Читаем источники если есть\nsource_names: dict[int, str] = {}\nif CFG.SOURCES_TXT.exists():\n    with open(CFG.SOURCES_TXT) as f:\n        for line in f:\n            line = line.strip()\n            if not line:\n                continue\n            # формат: \"0 AuraFlow\" или \"0: AuraFlow\"\n            parts = line.replace(\":\", \"\").split(None, 1)\n            if len(parts) == 2 and parts[0].isdigit():\n                source_names[int(parts[0])] = parts[1]\n    print(\"Sources:\", source_names)\n\n# Проверяем баланс → выбираем loss\ncounts = train_data[\"y\"].value_counts()\nimbalance_ratio = counts.max() / counts.min()\nUSE_FOCAL = imbalance_ratio > 2.0\nprint(f\"Imbalance ratio: {imbalance_ratio:.2f}  →  {'Focal Loss' if USE_FOCAL else 'CrossEntropy'}\")\n\n# Train / Val split\ntrain_data, val_data = train_test_split(\n    train_df, test_size=CFG.VAL_SPLIT,\n    stratify=train_df[\"y\"], random_state=CFG.SEED\n)\ntrain_data = train_data.reset_index(drop=True)\nval_data   = val_data.reset_index(drop=True)\nprint(f\"Train: {len(train_data)}  Val: {len(val_data)}  Test: {len(test_df)}\")\n\n\n# ── Файловый индекс (решает проблему Data/Data/Data/…) ────────────────────────\nprint(\"Building file index …\")\n_FILE_INDEX: dict[str, Path] = {}\nfor ext in (\"*.png\", \"*.jpg\", \"*.jpeg\", \"*.webp\"):\n    for p in CFG.DATA_DIR.rglob(ext):\n        _FILE_INDEX[p.name] = p\nprint(f\"File index: {len(_FILE_INDEX)} images\")\n\ndef resolve_path(raw: str) -> Path:\n    p = Path(raw)\n    if p.is_absolute() and p.exists():\n        return p\n    c = CFG.DATA_DIR / p\n    if c.exists():\n        return c\n    if p.name in _FILE_INDEX:\n        return _FILE_INDEX[p.name]\n    parts = p.parts\n    for s in range(1, len(parts)):\n        c2 = CFG.DATA_DIR / Path(*parts[s:])\n        if c2.exists():\n            return c2\n    return c   # вернём как есть; упадёт в try/except\n\n# Быстрая проверка\n_tp = resolve_path(train_df[\"path\"].iloc[0])\nprint(f\"Path check: {_tp}  exists={_tp.exists()}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 3. ВИЗУАЛИЗАЦИЯ ДВУХ КЛАССОВ\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef visualize_two_classes(df: pd.DataFrame, class_a: int, class_b: int,\n                          n: int = 5, save_path: Path = CFG.VIZ_PATH):\n    \"\"\"\n    Показывает n картинок из class_a (строка 1) и n из class_b (строка 2).\n    Сохраняет PNG в save_path.\n    \"\"\"\n    fig, axes = plt.subplots(2, n, figsize=(3 * n, 7))\n    name_a = source_names.get(class_a, f\"Class {class_a}\")\n    name_b = source_names.get(class_b, f\"Class {class_b}\")\n\n    for row_idx, (cls, name) in enumerate([(class_a, name_a), (class_b, name_b)]):\n        samples = df[df[\"y\"] == cls].sample(n=n, random_state=CFG.SEED)\n        for col_idx, (_, rec) in enumerate(samples.iterrows()):\n            ax = axes[row_idx][col_idx]\n            img_path = resolve_path(rec[\"path\"])\n            try:\n                img = Image.open(img_path).convert(\"RGB\")\n            except Exception:\n                img = Image.new(\"RGB\", (224, 224), color=(200, 200, 200))\n            ax.imshow(img)\n            ax.axis(\"off\")\n            if col_idx == 0:\n                ax.set_ylabel(name, fontsize=11, fontweight=\"bold\", rotation=90,\n                              labelpad=6, va=\"center\")\n\n    fig.suptitle(f\"Sample images: '{name_a}'  vs  '{name_b}'\", fontsize=13, y=1.01)\n    plt.tight_layout()\n    plt.savefig(save_path, bbox_inches=\"tight\", dpi=120)\n    plt.close()\n    print(f\"Visualization saved → {save_path}\")\n\n# Визуализируем класс 0 и класс 1 (можно поменять)\nvisualize_two_classes(train_df, class_a=0, class_b=1)\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 4. АУГМЕНТАЦИИ (albumentations)\n# ══════════════════════════════════════════════════════════════════════════════\n\n_MEAN = (0.485, 0.456, 0.406)\n_STD  = (0.229, 0.224, 0.225)\n\ntrain_aug = A.Compose([\n    A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.1),\n    A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=15, p=0.5),\n    A.OneOf([\n        A.GaussianBlur(blur_limit=(3, 5)),\n        A.MotionBlur(blur_limit=5),\n        A.MedianBlur(blur_limit=3),\n    ], p=0.3),\n    A.OneOf([\n        A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2),\n        A.HueSaturationValue(hue_shift_limit=10, sat_shift_limit=20, val_shift_limit=10),\n        A.CLAHE(clip_limit=2.0),\n    ], p=0.5),\n    A.CoarseDropout(max_holes=8, max_height=32, max_width=32, p=0.3),\n    A.Normalize(mean=_MEAN, std=_STD),\n    ToTensorV2(),\n])\n\nval_aug = A.Compose([\n    A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n    A.Normalize(mean=_MEAN, std=_STD),\n    ToTensorV2(),\n])\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 5. DATASET\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass SIADataset(Dataset):\n\n    def __init__(\n        self,\n        df: pd.DataFrame,\n        data_dir: Path,\n        aug=None,\n        is_test: bool = False\n    ):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.aug = aug\n        self.is_test = is_test\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n\n        # путь из CSV\n        img_path = Path(row[\"path\"])\n\n        # если путь относительный\n        if not img_path.is_absolute():\n\n            img_path_str = str(img_path)\n\n            # убираем лишний префикс Data/\n            if img_path_str.startswith(\"Data/\"):\n                img_path_str = img_path_str.replace(\"Data/\", \"\", 1)\n\n            img_path = self.data_dir / img_path_str\n\n        # fallback через индекс\n        if not img_path.exists():\n            img_path = resolve_path(str(row[\"path\"]))\n\n        # загрузка\n        try:\n            img = np.array(Image.open(img_path).convert(\"RGB\"))\n\n        except Exception as e:\n\n            print(f\"Warning: cannot open {img_path}: {e}\")\n\n            img = np.zeros(\n                (CFG.IMG_SIZE, CFG.IMG_SIZE, 3),\n                dtype=np.uint8\n            )\n\n        # augmentations\n        if self.aug:\n            img = self.aug(image=img)[\"image\"]\n\n        if self.is_test:\n            return img, row[\"ID\"]\n\n        label = int(row[\"y\"])\n\n        return img, label\n\n\ntrain_dataset = SIADataset(\n    train_data,\n    CFG.DATA_DIR,\n    aug=train_aug\n)\n\nval_dataset = SIADataset(\n    val_data,\n    CFG.DATA_DIR,\n    aug=val_aug\n)\n\ntest_dataset = SIADataset(\n    test_df,\n    CFG.DATA_DIR,\n    aug=val_aug,\n    is_test=True\n)\n\ntrain_loader = DataLoader(train_dataset, CFG.BATCH_SIZE, shuffle=True,\n                          num_workers=CFG.NUM_WORKERS, pin_memory=True)\nval_loader   = DataLoader(val_dataset,   CFG.BATCH_SIZE, shuffle=False,\n                          num_workers=CFG.NUM_WORKERS, pin_memory=True)\ntest_loader  = DataLoader(test_dataset,  CFG.BATCH_SIZE, shuffle=False,\n                          num_workers=CFG.NUM_WORKERS, pin_memory=True)\n\nprint(f\"Batches → train:{len(train_loader)} val:{len(val_loader)} test:{len(test_loader)}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 6. CUTMIX / MIXUP  (batch-level)\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef mixup_data(x: torch.Tensor, y: torch.Tensor, alpha: float = 0.4):\n    \"\"\"Возвращает смешанные данные и лямбду.\"\"\"\n    lam = np.random.beta(alpha, alpha) if alpha > 0 else 1.0\n    idx = torch.randperm(x.size(0), device=x.device)\n    mixed_x = lam * x + (1 - lam) * x[idx]\n    return mixed_x, y, y[idx], lam\n\n\ndef rand_bbox(size, lam):\n    \"\"\"Вычисляет случайный bbox для CutMix.\"\"\"\n    W, H = size[-1], size[-2]\n    cut_ratio = math.sqrt(1.0 - lam)\n    cut_w, cut_h = int(W * cut_ratio), int(H * cut_ratio)\n    cx, cy = random.randint(0, W), random.randint(0, H)\n    x1, y1 = max(cx - cut_w // 2, 0), max(cy - cut_h // 2, 0)\n    x2, y2 = min(cx + cut_w // 2, W), min(cy + cut_h // 2, H)\n    return x1, y1, x2, y2\n\n\ndef cutmix_data(x: torch.Tensor, y: torch.Tensor):\n    lam = np.random.beta(1.0, 1.0)\n    idx = torch.randperm(x.size(0), device=x.device)\n    x1, y1, x2, y2 = rand_bbox(x.size(), lam)\n    x_mix = x.clone()\n    x_mix[:, :, y1:y2, x1:x2] = x[idx, :, y1:y2, x1:x2]\n    lam_adj = 1 - (x2 - x1) * (y2 - y1) / (x.size(-1) * x.size(-2))\n    return x_mix, y, y[idx], lam_adj\n\n\ndef mixed_criterion(criterion, pred, y_a, y_b, lam):\n    \"\"\"Взвешенный loss для MixUp / CutMix.\"\"\"\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 7. FOCAL LOSS\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass FocalLoss(nn.Module):\n    \"\"\"Focal Loss для многоклассовой классификации.\"\"\"\n    def __init__(self, gamma: float = 2.0, weight=None, reduction: str = \"mean\"):\n        super().__init__()\n        self.gamma     = gamma\n        self.weight    = weight\n        self.reduction = reduction\n\n    def forward(self, logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:\n        log_p  = F.log_softmax(logits, dim=1)\n        p      = torch.exp(log_p)\n        log_p_t = log_p.gather(1, targets.unsqueeze(1)).squeeze(1)\n        p_t     = p.gather(1, targets.unsqueeze(1)).squeeze(1)\n        fl      = -((1 - p_t) ** self.gamma) * log_p_t\n        if self.weight is not None:\n            w  = self.weight.to(logits.device)\n            fl = fl * w[targets]\n        return fl.mean() if self.reduction == \"mean\" else fl.sum()\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 8. ARCFACE HEAD\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass ArcFaceHead(nn.Module):\n    \"\"\"\n    ArcFace (Additive Angular Margin) classification head.\n    Во время инференса работает как обычный линейный слой (margin=0).\n    \"\"\"\n    def __init__(self, in_features: int, num_classes: int,\n                 s: float = 30.0, m: float = 0.50):\n        super().__init__()\n        self.s = s\n        self.m = m\n        self.weight = nn.Parameter(torch.empty(num_classes, in_features))\n        nn.init.xavier_uniform_(self.weight)\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.th    = math.cos(math.pi - m)   # cos(π - m)\n        self.mm    = math.sin(math.pi - m) * m\n\n    def forward(self, embeddings: torch.Tensor,\n                labels: torch.Tensor | None = None) -> torch.Tensor:\n        # Нормализуем эмбеддинги и веса → косинусные сходства\n        cos_theta = F.linear(F.normalize(embeddings), F.normalize(self.weight))\n\n        if labels is None or not self.training:\n            return cos_theta * self.s   # plain cosine logits во время eval/inference\n\n        # ArcFace margin\n        sin_theta  = torch.sqrt((1.0 - cos_theta.pow(2)).clamp(1e-6, 1.0))\n        cos_theta_m = cos_theta * self.cos_m - sin_theta * self.sin_m\n        # Безопасность: если cos(θ) < cos(π-m) → используем cos(θ)-mm\n        cos_theta_m = torch.where(cos_theta > self.th, cos_theta_m,\n                                  cos_theta - self.mm)\n\n        one_hot = torch.zeros_like(cos_theta)\n        one_hot.scatter_(1, labels.unsqueeze(1), 1.0)\n        output = one_hot * cos_theta_m + (1.0 - one_hot) * cos_theta\n        return output * self.s\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 9. МОДЕЛЬ  (Swin-T backbone + embedding bottleneck + ArcFace head)\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass SwinArcFace(nn.Module):\n    \"\"\"\n    Swin-Tiny → GlobalAvgPool → BN+Linear bottleneck → ArcFace head.\n    \"\"\"\n    def __init__(self, model_name: str, num_classes: int,\n                 embed_dim: int, s: float, m: float):\n        super().__init__()\n        # Backbone без classification head\n        self.backbone = timm.create_model(model_name, pretrained=True, num_classes=0)\n        feat_dim = self.backbone.num_features\n\n        # Embedding bottleneck\n        self.neck = nn.Sequential(\n            nn.Linear(feat_dim, embed_dim, bias=False),\n            nn.BatchNorm1d(embed_dim),\n        )\n        # ArcFace head\n        self.head = ArcFaceHead(embed_dim, num_classes, s=s, m=m)\n\n    def forward(self, x: torch.Tensor,\n                labels: torch.Tensor | None = None) -> tuple[torch.Tensor, torch.Tensor]:\n        feat   = self.backbone(x)           # (B, feat_dim)\n        embed  = self.neck(feat)             # (B, embed_dim)\n        logits = self.head(embed, labels)    # (B, num_classes)\n        return logits, embed\n\n\nmodel = SwinArcFace(\n    CFG.MODEL_NAME, CFG.NUM_CLASSES,\n    CFG.EMBED_DIM, CFG.ARC_S, CFG.ARC_M\n).to(CFG.DEVICE)\n\ntotal_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"\\nModel : {CFG.MODEL_NAME}\")\nprint(f\"Params: {total_params:,}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 10. LOSS, OPTIMIZER, SCHEDULER\n# ══════════════════════════════════════════════════════════════════════════════\n\n# Веса классов (для Focal Loss при сильном дисбалансе)\nclass_counts = train_df[\"y\"].value_counts().sort_index().values.astype(float)\nclass_weights = torch.tensor(class_counts.sum() / (len(class_counts) * class_counts),\n                              dtype=torch.float32)\n\nif USE_FOCAL:\n    criterion = FocalLoss(gamma=CFG.FOCAL_GAMMA, weight=class_weights)\n    print(\"Using: FocalLoss (γ={})\".format(CFG.FOCAL_GAMMA))\nelse:\n    criterion = nn.CrossEntropyLoss(label_smoothing=0.05, weight=class_weights.to(CFG.DEVICE))\n    print(\"Using: CrossEntropyLoss (label_smoothing=0.05)\")\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG.LR, weight_decay=CFG.WEIGHT_DECAY)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, T_max=CFG.NUM_EPOCHS, eta_min=CFG.LR / 20\n)\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 11. TRAIN / EVAL\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef train_one_epoch(model, loader, optimizer, criterion, device):\n    model.train()\n    total_loss, correct, total = 0.0, 0, 0\n    pbar = tqdm(loader, desc=\"  Train\", leave=False)\n\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n\n        # Выбираем аугментацию для батча\n        r = random.random()\n        if r < CFG.CUTMIX_PROB:\n            imgs, y_a, y_b, lam = cutmix_data(imgs, labels)\n            logits, _ = model(imgs, labels)   # ArcFace margin применяется к y_a\n            loss = mixed_criterion(criterion, logits, y_a, y_b, lam)\n        elif r < CFG.CUTMIX_PROB + CFG.MIXUP_PROB:\n            imgs, y_a, y_b, lam = mixup_data(imgs, labels, CFG.MIXUP_ALPHA)\n            logits, _ = model(imgs, labels)\n            loss = mixed_criterion(criterion, logits, y_a, y_b, lam)\n        else:\n            logits, _ = model(imgs, labels)\n            loss = criterion(logits, labels)\n\n        optimizer.zero_grad()\n        loss.backward()\n        nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n\n        total_loss += loss.item() * len(labels)\n        correct    += (logits.argmax(1) == labels).sum().item()\n        total      += len(labels)\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\", acc=f\"{correct/total:.4f}\")\n\n    return total_loss / total, correct / total\n\n\n@torch.no_grad()\ndef evaluate(model, loader, criterion, device):\n    model.eval()\n    total_loss, correct, total = 0.0, 0, 0\n\n    for imgs, labels in tqdm(loader, desc=\"  Val  \", leave=False):\n        imgs, labels = imgs.to(device), labels.to(device)\n        logits, _  = model(imgs)          # без labels → нет margin\n        loss = criterion(logits, labels)\n        total_loss += loss.item() * len(labels)\n        correct    += (logits.argmax(1) == labels).sum().item()\n        total      += len(labels)\n\n    return total_loss / total, correct / total\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 12. ЦИКЛ ОБУЧЕНИЯ\n# ══════════════════════════════════════════════════════════════════════════════\n\nbest_val_acc = 0.0\nhistory = []\n\nfor epoch in range(1, CFG.NUM_EPOCHS + 1):\n    lr_now = optimizer.param_groups[0][\"lr\"]\n    print(f\"\\nEpoch {epoch}/{CFG.NUM_EPOCHS}  lr={lr_now:.2e}\")\n\n    tr_loss, tr_acc = train_one_epoch(model, train_loader, optimizer, criterion, CFG.DEVICE)\n    vl_loss, vl_acc = evaluate(model, val_loader, criterion, CFG.DEVICE)\n    scheduler.step()\n\n    history.append(dict(epoch=epoch, train_loss=tr_loss, train_acc=tr_acc,\n                        val_loss=vl_loss, val_acc=vl_acc))\n    print(f\"  train  loss={tr_loss:.4f}  acc={tr_acc:.4f}\")\n    print(f\"  val    loss={vl_loss:.4f}  acc={vl_acc:.4f}\")\n\n    if vl_acc > best_val_acc:\n        best_val_acc = vl_acc\n        torch.save(model.state_dict(), CFG.CHECKPOINT)\n        print(f\"  ✓ Checkpoint saved  (val_acc={best_val_acc:.4f})\")\n\nprint(f\"\\nBest val accuracy: {best_val_acc:.4f}\")\npd.DataFrame(history).to_csv(CFG.OUTPUT_DIR / \"history.csv\", index=False)\n\n\n# ── График обучения ────────────────────────────────────────────────────────────\nhist_df = pd.DataFrame(history)\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))\nax1.plot(hist_df[\"epoch\"], hist_df[\"train_loss\"], label=\"train\")\nax1.plot(hist_df[\"epoch\"], hist_df[\"val_loss\"],   label=\"val\")\nax1.set_title(\"Loss\"); ax1.legend(); ax1.set_xlabel(\"Epoch\")\nax2.plot(hist_df[\"epoch\"], hist_df[\"train_acc\"], label=\"train\")\nax2.plot(hist_df[\"epoch\"], hist_df[\"val_acc\"],   label=\"val\")\nax2.set_title(\"Accuracy\"); ax2.legend(); ax2.set_xlabel(\"Epoch\")\nplt.tight_layout()\nplt.savefig(CFG.OUTPUT_DIR / \"training_curves.png\", dpi=120)\nplt.close()\nprint(f\"Training curves saved → {CFG.OUTPUT_DIR / 'training_curves.png'}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 13. ИНФЕРЕНС\n# ══════════════════════════════════════════════════════════════════════════════\n\nmodel.load_state_dict(torch.load(CFG.CHECKPOINT, map_location=CFG.DEVICE))\nmodel.eval()\n\nall_ids, all_preds = [], []\n\nwith torch.no_grad():\n    for imgs, ids in tqdm(test_loader, desc=\"Inference\"):\n        imgs   = imgs.to(CFG.DEVICE)\n        logits, _ = model(imgs)    # inference: нет margin\n        preds  = logits.argmax(1).cpu().numpy()\n        all_preds.extend(preds.tolist())\n        all_ids.extend(ids.numpy().tolist() if isinstance(ids, torch.Tensor) else list(ids))\n\nprint(f\"Predictions: {len(all_preds)} / {len(test_df)}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 14. SUBMISSION\n# ══════════════════════════════════════════════════════════════════════════════\n\nsubmission = pd.DataFrame({\"ID\": all_ids, \"TARGET\": all_preds})\n\nif CFG.SAMPLE_SUB.exists():\n    sample     = pd.read_csv(CFG.SAMPLE_SUB)\n    submission = sample[[\"ID\"]].merge(submission, on=\"ID\", how=\"left\")\n    missing    = submission[\"TARGET\"].isna().sum()\n    if missing:\n        print(f\"Warning: {missing} missing predictions → filled with 0\")\n        submission[\"TARGET\"] = submission[\"TARGET\"].fillna(0)\n    submission[\"TARGET\"] = submission[\"TARGET\"].astype(int)\n\nsubmission.to_csv(CFG.SUBMISSION, index=False)\nprint(f\"\\n✓ Submission saved → {CFG.SUBMISSION}\")\nprint(submission.head(10))\nprint(\"\\nPrediction distribution:\\n\", submission[\"TARGET\"].value_counts().sort_index())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T07:09:50.232522Z","iopub.execute_input":"2026-05-19T07:09:50.233381Z","iopub.status.idle":"2026-05-19T07:10:14.803418Z","shell.execute_reply.started":"2026-05-19T07:09:50.233350Z","shell.execute_reply":"2026-05-19T07:10:14.802251Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_two_classes(df: pd.DataFrame, class_a: int, class_b: int,\n                          n: int = 5, save_path: Path = CFG.VIZ_PATH):\n    \"\"\"\n    Показывает n картинок из class_a (строка 1) и n из class_b (строка 2).\n    Сохраняет PNG в save_path.\n    \"\"\"\n    fig, axes = plt.subplots(2, n, figsize=(3 * n, 7))\n    name_a = source_names.get(class_a, f\"Class {class_a}\")\n    name_b = source_names.get(class_b, f\"Class {class_b}\")\n\n    for row_idx, (cls, name) in enumerate([(class_a, name_a), (class_b, name_b)]):\n        samples = df[df[\"y\"] == cls].sample(n=n, random_state=CFG.SEED)\n        for col_idx, (_, rec) in enumerate(samples.iterrows()):\n            ax = axes[row_idx][col_idx]\n            img_path = resolve_path(rec[\"path\"])\n            try:\n                img = Image.open(img_path).convert(\"RGB\")\n            except Exception:\n                img = Image.new(\"RGB\", (224, 224), color=(200, 200, 200))\n            ax.imshow(img)\n            ax.axis(\"off\")\n            if col_idx == 0:\n                ax.set_ylabel(name, fontsize=11, fontweight=\"bold\", rotation=90,\n                              labelpad=6, va=\"center\")\n\n    fig.suptitle(f\"Sample images: '{name_a}'  vs  '{name_b}'\", fontsize=13, y=1.01)\n    plt.tight_layout()\n    plt.savefig(save_path, bbox_inches=\"tight\", dpi=120)\n    plt.close()\n    print(f\"Visualization saved → {save_path}\")\n\n# Визуализируем класс 0 и класс 1 (можно поменять)\nvisualize_two_classes(train_df, class_a=0, class_b=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T13:27:13.774824Z","iopub.execute_input":"2026-05-15T13:27:13.775531Z","iopub.status.idle":"2026-05-15T13:27:13.786773Z","shell.execute_reply.started":"2026-05-15T13:27:13.775487Z","shell.execute_reply":"2026-05-15T13:27:13.785698Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BinaryDataset(Dataset):\n\n    def __init__(self, df, target_class, transforms=None):\n        self.df = df.reset_index(drop=True)\n        self.target_class = target_class\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n\n        img_path = resolve_path(row[\"path\"])\n\n        img = np.array(\n            Image.open(img_path).convert(\"RGB\")\n        )\n\n        if self.transforms:\n            img = self.transforms(image=img)[\"image\"]\n\n        label = 1 if row[\"y\"] == self.target_class else 0\n\n        return img, torch.tensor(label, dtype=torch.float32)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BinaryEffNet(nn.Module):\n\n    def __init__(self):\n        super().__init__()\n\n        self.backbone = timm.create_model(\n            \"efficientnet_b0\",\n            pretrained=True,\n            num_classes=0\n        )\n\n        self.head = nn.Linear(\n            self.backbone.num_features,\n            1\n        )\n\n    def forward(self, x):\n\n        feat = self.backbone(x)\n\n        logit = self.head(feat)\n\n        return logit.squeeze(1)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models = []\n\nfor cls in range(10):\n\n    print(f\"\\nTraining class {cls}\")\n\n    train_ds = BinaryDataset(\n        train_data,\n        target_class=cls,\n        transforms=train_aug\n    )\n\n    train_loader = DataLoader(\n        train_ds,\n        batch_size=32,\n        shuffle=True,\n        num_workers=4\n    )\n\n    model = BinaryEffNet().to(device)\n\n    criterion = nn.BCEWithLogitsLoss()\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=1e-4\n    )\n\n    for epoch in range(5):\n\n        model.train()\n\n        for imgs, labels in train_loader:\n\n            imgs = imgs.to(device)\n            labels = labels.to(device)\n\n            logits = model(imgs)\n\n            loss = criterion(logits, labels)\n\n            optimizer.zero_grad()\n\n            loss.backward()\n\n            optimizer.step()\n\n    torch.save(model.state_dict(), f\"model_{cls}.pth\")\n\n    models.append(model)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models = []\n\nfor cls in range(10):\n\n    model = BinaryEffNet().to(device)\n\n    model.load_state_dict(\n        torch.load(f\"model_{cls}.pth\")\n    )\n\n    model.eval()\n\n    models.append(model)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef predict_image(img):\n\n    probs = []\n\n    x = val_aug(image=img)[\"image\"]\n    x = x.unsqueeze(0).to(device)\n\n    for model in models:\n\n        logit = model(x)\n\n        prob = torch.sigmoid(logit)\n\n        probs.append(prob.item())\n\n    pred = np.argmax(probs)\n\n    return pred, probs","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════════════════════\n# Synthetic Source Attribution\n# Strong Ensemble Pipeline\n#\n# Features:\n# - ConvNeXt + EfficientNet + Swin ensemble\n# - ArcFace + CrossEntropy\n# - FFT spectrum features\n# - JPEG / Blur / Crop / Rotation / Brightness augmentations\n# - Random augmentation combinations (1-3 ops)\n# - AdamW + Cosine + Warmup\n# - Label smoothing\n# - TTA inference\n# - Regularization against overfitting\n# - ImageNet pretrained\n#\n# Kaggle-ready\n# ══════════════════════════════════════════════════════════════════════════════\n\nimport os\nimport cv2\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\n\nfrom pathlib import Path\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\n\nimport timm\n\nfrom tqdm import tqdm\n\nfrom sklearn.model_selection import train_test_split\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom timm.scheduler.cosine_lr import CosineLRScheduler\n\n# ══════════════════════════════════════════════════════════════════════════════\n# CONFIG\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass CFG:\n    MODELS = [\n    \"convnext_small\",\n    \"efficientnet_b0\",\n    \"swin_tiny_patch4_window7_224\"\n    ]\n\n    DATA_DIR = Path(\n        \"/kaggle/input/competitions/\"\n        \"dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data\"\n    )\n\n    TRAIN_CSV = DATA_DIR / \"training.csv\"\n    TEST_CSV  = DATA_DIR / \"test.csv\"\n\n    OUTPUT_DIR = Path(\"/kaggle/working/\")\n\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    NUM_CLASSES = 10\n\n    IMG_SIZE = 224\n\n    BATCH_SIZE = 16\n\n    NUM_WORKERS = 4\n\n    NUM_EPOCHS = 12\n\n    LR = 1e-4\n\n    MIN_LR = 1e-6\n\n    WEIGHT_DECAY = 5e-2\n\n    WARMUP_EPOCHS = 1\n\n    SEED = 42\n\n    VAL_SPLIT = 0.15\n\n    EMBED_DIM = 512\n\n    ARC_S = 32.0\n    ARC_M = 0.4\n\n    LABEL_SMOOTHING = 0.1\n\n    DROPOUT = 0.3\n\n    USE_FFT = True\n\n    TTA = True\n\n    MODELS = [\n        \"convnext_small\",\n        \"tf_efficientnet_b3\",\n        \"swin_tiny_patch4_window7_224\"\n    ]\n\n\nCFG.OUTPUT_DIR.mkdir(exist_ok=True, parents=True)\n\n# ══════════════════════════════════════════════════════════════════════════════\n# SEED\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef seed_everything(seed):\n\n    random.seed(seed)\n    np.random.seed(seed)\n\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(CFG.SEED)\n\n# ══════════════════════════════════════════════════════════════════════════════\n# DATA\n# ══════════════════════════════════════════════════════════════════════════════\n\ntrain_df = pd.read_csv(CFG.TRAIN_CSV)\ntest_df  = pd.read_csv(CFG.TEST_CSV)\n\ntrain_df, val_df = train_test_split(\n    train_df,\n    test_size=CFG.VAL_SPLIT,\n    stratify=train_df[\"y\"],\n    random_state=CFG.SEED\n)\n\ntrain_df = train_df.reset_index(drop=True)\nval_df   = val_df.reset_index(drop=True)\n\n# ══════════════════════════════════════════════════════════════════════════════\n# FFT\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef fft_feature(img):\n\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n\n    fft = np.fft.fft2(gray)\n\n    fft_shift = np.fft.fftshift(fft)\n\n    magnitude = np.log(np.abs(fft_shift) + 1)\n\n    magnitude = cv2.normalize(\n        magnitude,\n        None,\n        0,\n        255,\n        cv2.NORM_MINMAX\n    ).astype(np.uint8)\n\n    magnitude = cv2.cvtColor(\n        magnitude,\n        cv2.COLOR_GRAY2RGB\n    )\n\n    return magnitude\n\n\n\n\nclass AttributionModel(nn.Module):\n\n    def __init__(self, model_name):\n        super().__init__()\n\n        self.backbone = timm.create_model(\n            model_name,\n            pretrained=False,\n            num_classes=CFG.NUM_CLASSES\n        )\n\n    def forward(self, x):\n\n        return self.backbone(x)\n# ══════════════════════════════════════════════════════════════════════════════\n# AUGMENTATIONS\n# ══════════════════════════════════════════════════════════════════════════════\n\n_MEAN = (0.485, 0.456, 0.406)\n_STD  = (0.229, 0.224, 0.225)\n\ntrain_aug = A.Compose([\n\n    A.Resize(\n        CFG.IMG_SIZE,\n        CFG.IMG_SIZE\n    ),\n\n    A.Lambda(\n        image=lambda img, **kw: random_post_ops()(image=img)[\"image\"]\n    ),\n\n    A.HorizontalFlip(p=0.5),\n\n    A.Normalize(\n        mean=_MEAN,\n        std=_STD\n    ),\n\n    ToTensorV2()\n\n])\n\nA.OneOf([\n\n\n\n    A.SomeOf([\n        A.ImageCompression(\n            quality_range=(40, 100),\n            p=1.0\n        ),\n\n        A.Rotate(\n            limit=20,\n            border_mode=cv2.BORDER_REFLECT_101,\n            p=1.0\n        ),\n\n        A.GaussianBlur(\n            blur_limit=(3, 7),\n            p=1.0\n        ),\n\n        A.MotionBlur(\n            blur_limit=7,\n            p=1.0\n        ),\n\n        A.ToGray(\n            p=1.0\n        ),\n\n        A.RandomBrightnessContrast(\n            brightness_limit=0.2,\n            contrast_limit=0.2,\n            p=1.0\n        ),\n\n        A.Sharpen(\n            alpha=(0.1, 0.3),\n            p=1.0\n        ),\n\n    ],\n    n=2,\n    replace=False,\n    p=1.0),\n\n    A.SomeOf([\n            A.ImageCompression(\n                quality_range=(40, 100),\n                p=1.0\n            ),\n    \n            A.Rotate(\n                limit=20,\n                border_mode=cv2.BORDER_REFLECT_101,\n                p=1.0\n            ),\n    \n            A.GaussianBlur(\n                blur_limit=(3, 7),\n                p=1.0\n            ),\n    \n            A.MotionBlur(\n                blur_limit=7,\n                p=1.0\n            ),\n    \n            A.ToGray(\n                p=1.0\n            ),\n    \n            A.RandomBrightnessContrast(\n                brightness_limit=0.2,\n                contrast_limit=0.2,\n                p=1.0\n            ),\n    \n            A.Sharpen(\n                alpha=(0.1, 0.3),\n                p=1.0\n            ),\n    \n        ],\n        n=3,\n        replace=False,\n        p=1.0),\n    \n    ], p=0.9)\n\n    n=(1, 3),\n    replace=False,\n    p=0.9),\n\n    A.HorizontalFlip(p=0.5),\n\n    A.Normalize(mean=_MEAN, std=_STD),\n\n    ToTensorV2()\n\n])\n\nval_aug = A.Compose([\n\n    A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n\n    A.Normalize(mean=_MEAN, std=_STD),\n\n    ToTensorV2()\n\n])\n\n\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# DATASET\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass SIADataset(Dataset):\n\n    def __init__(\n        self,\n        df,\n        transforms=None,\n        is_test=False\n    ):\n\n        self.df = df.reset_index(drop=True)\n\n        self.transforms = transforms\n\n        self.is_test = is_test\n\n    def __len__(self):\n\n        return len(self.df)\n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n\n        img_path = CFG.DATA_DIR / row[\"path\"].replace(\"Data/\", \"\")\n\n        img = np.array(\n            Image.open(img_path).convert(\"RGB\")\n        )\n\n        # FFT branch\n        if CFG.USE_FFT:\n\n            fft_img = fft_feature(img)\n\n            img = cv2.addWeighted(\n                img,\n                0.7,\n                fft_img,\n                0.3,\n                0\n            )\n\n        if self.transforms:\n\n            img = self.transforms(image=img)[\"image\"]\n\n        if self.is_test:\n\n            return img, row[\"ID\"]\n\n        label = int(row[\"y\"])\n\n        return img, label\n\n# ══════════════════════════════════════════════════════════════════════════════\n# DATALOADERS\n# ══════════════════════════════════════════════════════════════════════════════\ntrain_ds = SIADataset(\n    train_df,\n    transforms=build_train_aug(),   # ✅\n)\n\nval_ds = SIADataset(\n    val_df,\n    transforms=build_val_aug(),     # ✅\n)\n\ntest_ds = SIADataset(\n    test_df,\n    transforms=build_val_aug(),     # ✅\n    is_test=True\n)\n\nval_loader = DataLoader(\n    val_ds,\n    batch_size=CFG.BATCH_SIZE,\n    shuffle=False,\n    num_workers=CFG.NUM_WORKERS,\n    pin_memory=True\n)\n\ntest_loader = DataLoader(\n    test_ds,\n    batch_size=CFG.BATCH_SIZE,\n    shuffle=False,\n    num_workers=CFG.NUM_WORKERS,\n    pin_memory=True\n)\n# ══════════════════════════════════════════════════════════════════════════════\n# ARCFACE\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass ArcMarginProduct(nn.Module):\n\n    def __init__(\n        self,\n        in_features,\n        out_features,\n        s=32.0,\n        m=0.4\n    ):\n\n        super().__init__()\n\n        self.weight = nn.Parameter(\n            torch.FloatTensor(out_features, in_features)\n        )\n\n        nn.init.xavier_uniform_(self.weight)\n\n        self.s = s\n        self.m = m\n\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n\n    def forward(self, x, labels=None):\n\n        cosine = F.linear(\n            F.normalize(x),\n            F.normalize(self.weight)\n        )\n\n        if labels is None:\n\n            return cosine * self.s\n\n        sine = torch.sqrt(\n            1.0 - torch.pow(cosine, 2)\n        )\n\n        phi = cosine * self.cos_m - sine * self.sin_m\n\n        one_hot = torch.zeros_like(cosine)\n\n        one_hot.scatter_(\n            1,\n            labels.view(-1, 1),\n            1\n        )\n\n        output = (\n            one_hot * phi +\n            (1.0 - one_hot) * cosine\n        )\n\n        output *= self.s\n\n        return output\n\n# ══════════════════════════════════════════════════════════════════════════════\n# MODEL\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass AttributionModel(nn.Module):\n\n    def __init__(self, backbone_name):\n\n        super().__init__()\n\n        self.backbone = timm.create_model(\n\n            backbone_name,\n\n            pretrained=True,\n\n            num_classes=0,\n\n            drop_rate=CFG.DROPOUT,\n\n            drop_path_rate=0.2\n        )\n\n        feat_dim = self.backbone.num_features\n\n        self.embedding = nn.Sequential(\n\n            nn.Linear(feat_dim, CFG.EMBED_DIM),\n\n            nn.LayerNorm(CFG.EMBED_DIM),\n\n            nn.GELU(),\n\n            nn.Dropout(CFG.DROPOUT)\n        )\n\n        self.arc = ArcMarginProduct(\n            CFG.EMBED_DIM,\n            CFG.NUM_CLASSES,\n            CFG.ARC_S,\n            CFG.ARC_M\n        )\n\n    def forward(self, x, labels=None):\n\n        feat = self.backbone(x)\n\n        emb = self.embedding(feat)\n\n        logits = self.arc(emb, labels)\n\n        return logits\n\n# ══════════════════════════════════════════════════════════════════════════════\n# LOSS\n# ══════════════════════════════════════════════════════════════════════════════\n\ncriterion = nn.CrossEntropyLoss(\n    label_smoothing=CFG.LABEL_SMOOTHING\n)\n\n# ══════════════════════════════════════════════════════════════════════════════\n# TRAIN\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef train_fn(model, loader, optimizer, scheduler):\n\n    model.train()\n\n    total_loss = 0\n\n    correct = 0\n    total = 0\n\n    pbar = tqdm(loader)\n\n    for imgs, labels in pbar:\n\n        imgs = imgs.to(CFG.DEVICE)\n        labels = labels.to(CFG.DEVICE)\n\n        optimizer.zero_grad()\n\n        logits = model(imgs, labels)\n\n        loss = criterion(logits, labels)\n\n        loss.backward()\n\n        nn.utils.clip_grad_norm_(\n            model.parameters(),\n            1.0\n        )\n\n        optimizer.step()\n\n        scheduler.step_update(\n            num_updates=scheduler._num_updates + 1\n        )\n\n        total_loss += loss.item() * len(labels)\n\n        preds = logits.argmax(1)\n\n        correct += (preds == labels).sum().item()\n\n        total += len(labels)\n\n        pbar.set_postfix(\n            loss=f\"{loss.item():.4f}\",\n            acc=f\"{correct/total:.4f}\"\n        )\n\n    return total_loss / total, correct / total\n\n# ══════════════════════════════════════════════════════════════════════════════\n# VALID\n# ══════════════════════════════════════════════════════════════════════════════\n\n@torch.no_grad()\ndef valid_fn(model, loader):\n\n    model.eval()\n\n    total_loss = 0\n\n    correct = 0\n    total = 0\n\n    for imgs, labels in tqdm(loader):\n\n        imgs = imgs.to(CFG.DEVICE)\n        labels = labels.to(CFG.DEVICE)\n\n        logits = model(imgs)\n\n        loss = criterion(logits, labels)\n\n        total_loss += loss.item() * len(labels)\n\n        preds = logits.argmax(1)\n\n        correct += (preds == labels).sum().item()\n\n        total += len(labels)\n\n    return total_loss / total, correct / total\n\n# ══════════════════════════════════════════════════════════════════════════════\n# TRAIN ENSEMBLE\n# ══════════════════════════════════════════════════════════════════════════════\n\nall_models = []\n\nfor model_name in CFG.MODELS:\n\n    print(f\"\\n{'='*60}\")\n    print(model_name)\n    print(f\"{'='*60}\")\n\n    model = AttributionModel(model_name)\n\n    model = model.to(CFG.DEVICE)\n\n    optimizer = torch.optim.AdamW(\n\n        model.parameters(),\n\n        lr=CFG.LR,\n\n        weight_decay=CFG.WEIGHT_DECAY\n    )\n\n    scheduler = CosineLRScheduler(\n\n        optimizer,\n\n        t_initial=CFG.NUM_EPOCHS * len(train_loader),\n\n        lr_min=CFG.MIN_LR,\n\n        warmup_lr_init=1e-6,\n\n        warmup_t=CFG.WARMUP_EPOCHS * len(train_loader),\n\n        cycle_limit=1,\n\n        t_in_epochs=False\n    )\n\n    best_acc = 0\n\n    for epoch in range(CFG.NUM_EPOCHS):\n\n        print(f\"\\nEpoch {epoch+1}\")\n\n        train_loss, train_acc = train_fn(\n            model,\n            train_loader,\n            optimizer,\n            scheduler\n        )\n\n        val_loss, val_acc = valid_fn(\n            model,\n            val_loader\n        )\n\n        print(f\"train_acc={train_acc:.4f}\")\n        print(f\"val_acc={val_acc:.4f}\")\n\n        if val_acc > best_acc:\n\n            best_acc = val_acc\n\n            torch.save(\n\n                model.state_dict(),\n\n                CFG.OUTPUT_DIR / f\"{model_name}.pth\"\n            )\n\n            print(\"saved\")\n\n    all_models.append(model_name)\n# ══════════════════════════════════════════════════════════════════════════════\n# TTA\n# ══════════════════════════════════════════════════════════════════════════════\n\n@torch.no_grad()\ndef tta_predict(model, imgs):\n\n    logits1 = model(imgs)\n\n    # horizontal flip\n    logits2 = model(\n        torch.flip(imgs, dims=[3])\n    )\n\n    # rotate 90\n    logits3 = model(\n        torch.rot90(imgs, 1, [2, 3])\n    )\n\n    logits = (\n        logits1 +\n        logits2 +\n        logits3\n    ) / 3.0\n\n    return logits\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# LOAD ENSEMBLE\n# ══════════════════════════════════════════════════════════════════════════════\n\nmodels = []\n\nfor model_name in CFG.MODELS:\n\n    print(f\"Loading {model_name}\")\n\n    model = AttributionModel(model_name)\n\n    checkpoint_path = CFG.OUTPUT_DIR / f\"{model_name}.pth\"\n\n    model.load_state_dict(\n        torch.load(\n            checkpoint_path,\n            map_location=CFG.DEVICE\n        )\n    )\n\n    model = model.to(CFG.DEVICE)\n\n    model.eval()\n\n    models.append(model)\n\nprint(f\"Loaded {len(models)} models\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# ENSEMBLE INFERENCE\n# ══════════════════════════════════════════════════════════════════════════════\n\nall_ids = []\nall_preds = []\n\nwith torch.no_grad():\n\n    for imgs, ids in tqdm(test_loader, desc=\"Inference\"):\n\n        imgs = imgs.to(CFG.DEVICE)\n\n        ensemble_logits = None\n\n        for model in models:\n\n            logits = tta_predict(model, imgs)\n\n            if ensemble_logits is None:\n                ensemble_logits = logits\n            else:\n                ensemble_logits += logits\n\n        ensemble_logits /= len(models)\n\n        preds = ensemble_logits.argmax(dim=1).cpu().numpy()\n\n        all_preds.extend(preds.tolist())\n\n        # ids может быть tensor или list\n        if isinstance(ids, torch.Tensor):\n            all_ids.extend(ids.cpu().numpy().tolist())\n        else:\n            all_ids.extend(list(ids))\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# SUBMISSION\n# ══════════════════════════════════════════════════════════════════════════════\n\nsubmission = pd.DataFrame({\n    \"ID\": all_ids,\n    \"TARGET\": all_preds\n})\n\nsubmission.to_csv(\n    CFG.OUTPUT_DIR / \"submission.csv\",\n    index=False\n)\n\nprint(submission.head())\n\nprint(\"\\nPrediction distribution:\")\nprint(\n    submission[\"TARGET\"]\n    .value_counts()\n    .sort_index()\n)\n\nprint(\"\\nDone!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-17T10:11:02.329840Z","iopub.execute_input":"2026-05-17T10:11:02.330776Z","iopub.status.idle":"2026-05-17T10:11:02.379248Z","shell.execute_reply.started":"2026-05-17T10:11:02.330732Z","shell.execute_reply":"2026-05-17T10:11:02.377863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nEnsemble Solution: Synthetic Image Attribution Challenge\n=========================================================\nBackbone   : ConvNeXt-Small  +  EfficientNet-B2  (+  Swin-Tiny optional)\nLoss       : ArcFace  +  CrossEntropy(label_smoothing=0.1)\nAugments   : JPEG / resize / crop / rotation / grayscale / blur /\n             brightness / super-res  →  1–3 ops per image (albumentations)\n             + FFT/DCT spectral channel appended to input\nOptimizer  : AdamW + CosineAnnealingWarmRestarts + LR warmup\nEnsemble   : soft-vote over per-model logits\nTTA        : 5 augmented views per test image\n\"\"\"\n\n# ─────────────────────────────────────────────────────────────────────────────\n# 0. Dependencies\n#    pip install timm albumentations torch torchvision pandas pillow\n#                scikit-learn tqdm matplotlib scipy\n# ─────────────────────────────────────────────────────────────────────────────\n\nimport math, os, random, warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib; matplotlib.use(\"Agg\")\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom PIL import Image\nfrom scipy.fft import dctn\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms.functional as TF\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 1.  CONFIG\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef _find_data_dir() -> Path:\n    for c in [\n        Path(\"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge\"),\n        Path(\"/kaggle/input/dlmmdd-workshop-synthetic-source-attribution-challenge\"),\n        Path(\"/kaggle/input\"),\n        Path(\"./data\"),\n    ]:\n        if c.exists() and any(c.rglob(\"train.csv\")):\n            return c\n    return Path(\"./data\")\n\ndef _find_file(root: Path, name: str) -> Path:\n    hits = list(root.rglob(name))\n    return hits[0] if hits else root / name\n\n_ROOT = _find_data_dir()\n\n\nfrom pathlib import Path\n\nclass CFG:\n\n    # ── paths ───────────────────────────────────────────────────────────────\n    DATA_DIR = Path(\n        \"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data\"\n    )\n\n    TRAIN_CSV   = DATA_DIR / \"training.csv\"\n    TEST_CSV    = DATA_DIR / \"test.csv\"\n    SAMPLE_SUB  = DATA_DIR / \"sample_submission.csv\"\n    SOURCES_TXT = DATA_DIR / \"sources.txt\"\n\n    OUTPUT_DIR = (\n        Path(\"/kaggle/working\")\n        if Path(\"/kaggle/working\").exists()\n        else Path(\"./outputs\")\n    )\n\n    # ── models ──────────────────────────────────────────────────────────────\n    MODELS = [\n    \n        dict(\n            name=\"convnext_small\",\n            embed_dim=512,\n            tag=\"convnext\",\n            type=\"rgb\"\n        ),\n    \n        dict(\n            name=\"efficientnet_b2\",\n            embed_dim=512,\n            tag=\"efficientnet\",\n            type=\"rgb\"\n        ),\n    \n        dict(\n            name=\"davit_small.msft_in1k\",\n            embed_dim=512,\n            tag=\"davit\",\n            type=\"rgb\"\n        ),\n    \n        # NEW FFT MODEL\n        dict(\n            name=\"eva02_small_patch14_224.mim_in22k\",\n            embed_dim=384,\n            tag=\"eva02\",\n            type=\"rgb\"\n        ),\n    ]\n\n    # ── training ────────────────────────────────────────────────────────────\n    NUM_CLASSES   = 10\n    IMG_SIZE      = 224\n    BATCH_SIZE = 16\n    NUM_EPOCHS    = 7\n    WARMUP_EPOCHS = 2\n\n    LR            = 3e-4\n    LR_MIN        = 1e-6\n    WEIGHT_DECAY  = 1e-2\n\n    VAL_SPLIT     = 0.15\n    NUM_WORKERS   = 4\n\n    SEED          = 42\n\n    DEVICE = (\n        \"cuda\"\n        if torch.cuda.is_available()\n        else \"cpu\"\n    )\n\n    # ── ArcFace ─────────────────────────────────────────────────────────────\n    ARC_S = 30.0\n    ARC_M = 0.45\n\n    # ── losses ──────────────────────────────────────────────────────────────\n    ARCFACE_W     = 0.5\n    CE_W          = 0.5\n    LABEL_SMOOTH = 0.05\n\n    # ── augmentations ──────────────────────────────────────────────────────\n    MAX_POST_OPS = 3\n\n    # ── spectral ────────────────────────────────────────────────────────────\n    USE_SPECTRAL = True\n\n    # RGB + DCT\n    IN_CHANNELS = 4 if USE_SPECTRAL else 3\n\n    # ── TTA ─────────────────────────────────────────────────────────────────\n    TTA_STEPS = 5\n\n    # ── regularization ──────────────────────────────────────────────────────\n    DROP_RATE      = 0.3\n    DROP_PATH_RATE = 0.05\n\n    # ── output ──────────────────────────────────────────────────────────────\n    SUBMISSION = OUTPUT_DIR / \"submission.csv\"\n\n\nCFG.OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(f\"Data root : {CFG.DATA_DIR}\")\nprint(f\"Device    : {CFG.DEVICE}\")\nprint(f\"Models    : {[m['tag'] for m in CFG.MODELS]}\")\nprint(f\"Spectral  : {CFG.USE_SPECTRAL} (in_channels={CFG.IN_CHANNELS})\")\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 2.  DATA LOADING\n# ══════════════════════════════════════════════════════════════════════════════\n\ntrain_df = pd.read_csv(CFG.TRAIN_CSV)\ntest_df  = pd.read_csv(CFG.TEST_CSV)\n\nprint(f\"\\nTrain {train_df.shape}  Test {test_df.shape}\")\nprint(\"Class distribution:\\n\", train_df[\"y\"].value_counts().sort_index())\n\nsource_names: dict[int, str] = {}\nif CFG.SOURCES_TXT.exists():\n    for line in CFG.SOURCES_TXT.read_text().splitlines():\n        parts = line.strip().replace(\":\", \"\").split(None, 1)\n        if len(parts) == 2 and parts[0].isdigit():\n            source_names[int(parts[0])] = parts[1]\n    print(\"Sources:\", source_names)\n\ntrain_data, val_data = train_test_split(\n    train_df, test_size=CFG.VAL_SPLIT,\n    stratify=train_df[\"y\"], random_state=CFG.SEED\n)\ntrain_data = train_data.reset_index(drop=True)\nval_data   = val_data.reset_index(drop=True)\nprint(f\"Train {len(train_data)}  Val {len(val_data)}  Test {len(test_df)}\")\n\n\n# ─── File index ───────────────────────────────────────────────────────────────\nprint(\"\\nBuilding file index …\")\n_FILE_INDEX: dict[str, Path] = {}\nfor ext in (\"*.png\", \"*.jpg\", \"*.jpeg\", \"*.webp\"):\n    for p in CFG.DATA_DIR.rglob(ext):\n        _FILE_INDEX[p.name] = p\nprint(f\"  {len(_FILE_INDEX)} images indexed\")\n\ndef resolve_path(raw: str) -> Path:\n    p = Path(raw)\n    if p.is_absolute() and p.exists(): return p\n    c = CFG.DATA_DIR / p\n    if c.exists(): return c\n    if p.name in _FILE_INDEX: return _FILE_INDEX[p.name]\n    parts = p.parts\n    for s in range(1, len(parts)):\n        c2 = CFG.DATA_DIR / Path(*parts[s:])\n        if c2.exists(): return c2\n    return c\n\n\n# ── Quick sanity check ────────────────────────────────────────────────────────\n_tp = resolve_path(train_df[\"path\"].iloc[0])\nprint(f\"Path check → {_tp}  exists={_tp.exists()}\")\n# ═══════════════════════════════════════════════════════════════\n# FFT SPECTRAL IMAGE\n# ═══════════════════════════════════════════════════════════════\n\ndef make_fft_image(img):\n\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n\n    gray = cv2.resize(\n        gray,\n        (CFG.IMG_SIZE, CFG.IMG_SIZE)\n    )\n\n    gray = np.float32(gray) / 255.0\n\n    # FFT\n    fft = np.fft.fft2(gray)\n\n    fft = np.fft.fftshift(fft)\n\n    magnitude = np.abs(fft)\n\n    magnitude = np.log1p(magnitude)\n\n    magnitude = (\n        magnitude - magnitude.min()\n    ) / (\n        magnitude.max() - magnitude.min() + 1e-6\n    )\n\n    magnitude = (magnitude * 255).astype(np.uint8)\n\n    # делаем pseudo RGB\n    fft_img = np.stack([\n        magnitude,\n        magnitude,\n        magnitude\n    ], axis=-1)\n\n    return fft_img\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 3.  SPECTRAL FEATURE  (DCT / FFT magnitude → single gray channel)\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef compute_dct_channel(img_np: np.ndarray, size: int = CFG.IMG_SIZE) -> np.ndarray:\n    \"\"\"\n    Computes 2-D DCT magnitude of the luminance channel.\n    Returns float32 array shape (size, size) normalised to [0, 1].\n\n    Why DCT?\n    --------\n    AI-generated images leave characteristic spectral fingerprints in the\n    frequency domain (e.g. periodic grid artefacts from VAE upsampling,\n    spectral peaks from diffusion noise schedules). Appending the DCT\n    magnitude as a 4th channel gives the classifier a direct spectral view.\n    \"\"\"\n    gray = 0.299 * img_np[:, :, 0] + 0.587 * img_np[:, :, 1] + 0.114 * img_np[:, :, 2]\n    dct  = dctn(gray.astype(np.float32), norm=\"ortho\")\n    mag  = np.log1p(np.abs(dct))              # log-compress dynamic range\n    mag  = (mag - mag.min()) / (mag.max() - mag.min() + 1e-8)\n    return mag.astype(np.float32)\n\ndef random_post_ops():\n\n    n_ops = random.randint(1, CFG.MAX_POST_OPS)\n\n    selected = random.sample(post_ops, n_ops)\n\n    return A.Compose(selected)\n# ══════════════════════════════════════════════════════════════════════════════\n# 4.  AUGMENTATION PIPELINE\n# ══════════════════════════════════════════════════════════════════════════════\n\n# ══════════════════════════════════════════════════════════════════════════════\n# RANDOM POST-PROCESSING OPS\n# ══════════════════════════════════════════════════════════════════════════════\n\npost_ops = [\n\n    # JPEG artifacts\n    A.ImageCompression(\n        quality_range=(40, 85),\n        p=1.0\n    ),\n\n    # resize artifacts\n    A.Downscale(\n        scale_range=(0.5, 0.9),\n        p=1.0\n    ),\n\n    # random crop\n    A.RandomResizedCrop(\n        size=(CFG.IMG_SIZE, CFG.IMG_SIZE),\n        scale=(0.6, 1.0),\n        ratio=(0.75, 1.33),\n        p=1.0\n    ),\n\n    # rotation\n    A.Rotate(\n        limit=15,\n        border_mode=cv2.BORDER_REFLECT_101,\n        p=1.0\n    ),\n\n    # grayscale\n    A.ToGray(\n        p=1.0\n    ),\n\n    # blur\n    A.OneOf([\n        A.GaussianBlur(\n            blur_limit=(3, 7),\n            p=1.0\n        ),\n\n        A.MotionBlur(\n            blur_limit=7,\n            p=1.0\n        ),\n    ], p=1.0),\n\n    # brightness / contrast\n    A.RandomBrightnessContrast(\n        brightness_limit=0.2,\n        contrast_limit=0.2,\n        p=1.0\n    ),\n\n    # sharpen\n    A.Sharpen(\n        alpha=(0.1, 0.3),\n        p=1.0\n    ),\n]\n\n\ndef random_post_ops(n_ops=None):\n\n    ops = [\n\n        A.ImageCompression(\n            quality_range=(40, 90),\n            p=1.0\n        ),\n\n        A.Downscale(\n            scale_range=(0.5, 0.9),\n            p=1.0\n        ),\n\n        A.GaussianBlur(\n            blur_limit=(3, 7),\n            p=1.0\n        ),\n\n        A.MotionBlur(\n            blur_limit=5,\n            p=1.0\n        ),\n\n        A.ToGray(\n            p=1.0\n        ),\n\n        A.RandomBrightnessContrast(\n            brightness_limit=0.2,\n            contrast_limit=0.2,\n            p=1.0\n        ),\n\n        A.Rotate(\n            limit=15,\n            border_mode=cv2.BORDER_REFLECT_101,\n            p=1.0\n        ),\n    ]\n\n    # случайно выбираем 1-3 операции\n    if n_ops is None:\n        n_ops = random.randint(1, CFG.MAX_POST_OPS)\n\n    return A.SomeOf(\n        ops,\n        n=n_ops,\n        replace=False,\n        p=1.0\n    )\n\n_MEAN = (0.485, 0.456, 0.406)\n_STD  = (0.229, 0.224, 0.225)\ndef build_train_aug():\n\n    return A.Compose([\n\n        A.RandomResizedCrop(\n            size=(CFG.IMG_SIZE, CFG.IMG_SIZE),\n            scale=(0.7, 1.0),\n            ratio=(0.8, 1.2),\n            p=1.0\n        ),\n\n        A.HorizontalFlip(p=0.5),\n\n        A.Lambda(\n            image=lambda img, **kw: random_post_ops()(image=img)[\"image\"],\n            p=0.7\n        ),\n\n        A.Normalize(\n            mean=_MEAN,\n            std=_STD\n        ),\n\n        ToTensorV2()\n\n    ])\n\n\ndef build_val_aug():\n\n    return A.Compose([\n\n        A.Resize(\n            CFG.IMG_SIZE,\n            CFG.IMG_SIZE\n        ),\n\n        A.Normalize(\n            mean=_MEAN,\n            std=_STD\n        ),\n\n        ToTensorV2()\n\n    ])\n\ndef build_tta_aug():\n\n    return A.Compose([\n\n        A.Resize(\n            CFG.IMG_SIZE,\n            CFG.IMG_SIZE\n        ),\n\n        A.HorizontalFlip(p=0.5),\n\n        A.ShiftScaleRotate(\n            shift_limit=0.02,\n            scale_limit=0.05,\n            rotate_limit=5,\n            border_mode=cv2.BORDER_REFLECT_101,\n            p=0.3\n        ),\n\n        A.RandomBrightnessContrast(\n            brightness_limit=0.05,\n            contrast_limit=0.05,\n            p=0.2\n        ),\n\n        A.Normalize(\n            mean=_MEAN,\n            std=_STD\n        ),\n\n        ToTensorV2(),\n    ])\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 5.  DATASET\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass SIADataset(Dataset):\n\n    def __init__(self, df, transforms=None, is_test=False):\n\n        self.df = df.reset_index(drop=True)\n\n        self.transforms = transforms\n\n        self.is_test = is_test\n\n    def __len__(self):\n\n        return len(self.df)\n\n    def _augment(self, img):\n\n        if self.transforms is not None:\n            x = self.transforms(image=img)[\"image\"]\n        else:\n            x = ToTensorV2()(image=img)[\"image\"]\n\n        if CFG.USE_SPECTRAL:\n\n            dct = make_dct_image(img)\n\n            dct = torch.tensor(\n                dct,\n                dtype=torch.float32\n            ).unsqueeze(0)\n\n            x = torch.cat([x, dct], dim=0)\n\n        return x\n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n\n        img_path = resolve_path(row[\"path\"])\n\n        img = np.array(\n            Image.open(img_path).convert(\"RGB\")\n        )\n\n        x = self._augment(img)\n\n        if self.is_test:\n\n            return x, row[\"ID\"]\n\n        y = int(row[\"y\"])\n\n        return x, y\n\n\n\n\n\n# ═══════════════════════════════════════════════════════════════\n# FFT DATASET\n# ═══════════════════════════════════════════════════════════════\n\nclass FFTDataset(Dataset):\n\n    def __init__(self, df, transforms=None, is_test=False):\n\n        self.df = df.reset_index(drop=True)\n\n        self.transforms = transforms\n\n        self.is_test = is_test\n\n    def __len__(self):\n\n        return len(self.df)\n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n\n        img_path = resolve_path(row[\"path\"])\n\n        img = np.array(\n            Image.open(img_path).convert(\"RGB\")\n        )\n\n        fft_img = make_fft_image(img)\n\n        if self.transforms is not None:\n\n            x = self.transforms(\n                image=fft_img\n            )[\"image\"]\n\n        else:\n\n            x = ToTensorV2()(image=fft_img)[\"image\"]\n\n        if self.is_test:\n\n            return x, row[\"ID\"]\n\n        y = int(row[\"y\"])\n\n        return x, y\ndef make_loaders(train_df, val_df, test_df):\n\n    tr = SIADataset(\n        train_df,\n        transforms=build_train_aug(),\n        is_test=False\n    )\n\n    vl = SIADataset(\n        val_df,\n        transforms=build_val_aug(),\n        is_test=False\n    )\n\n    te = SIADataset(\n        test_df,\n        transforms=build_val_aug(),\n        is_test=True\n    )\n\n    kw = dict(\n        num_workers=CFG.NUM_WORKERS,\n        pin_memory=True\n    )\n\n    return (\n\n        DataLoader(\n            tr,\n            batch_size=CFG.BATCH_SIZE,\n            shuffle=True,\n            **kw\n        ),\n\n        DataLoader(\n            vl,\n            batch_size=CFG.BATCH_SIZE,\n            shuffle=False,\n            **kw\n        ),\n\n        DataLoader(\n            te,\n            batch_size=CFG.BATCH_SIZE,\n            shuffle=False,\n            **kw\n        ),\n    )\n\n\n\n\ndef make_fft_loaders(train_df, val_df, test_df):\n\n    tr = FFTDataset(\n        train_df,\n        transforms=build_train_aug(),\n        is_test=False\n    )\n\n    vl = FFTDataset(\n        val_df,\n        transforms=build_val_aug(),\n        is_test=False\n    )\n\n    te = FFTDataset(\n        test_df,\n        transforms=build_val_aug(),\n        is_test=True\n    )\n\n    kw = dict(\n        num_workers=CFG.NUM_WORKERS,\n        pin_memory=True\n    )\n\n    return (\n\n        DataLoader(\n            tr,\n            batch_size=CFG.BATCH_SIZE,\n            shuffle=True,\n            **kw\n        ),\n\n        DataLoader(\n            vl,\n            batch_size=CFG.BATCH_SIZE,\n            shuffle=False,\n            **kw\n        ),\n\n        DataLoader(\n            te,\n            batch_size=CFG.BATCH_SIZE,\n            shuffle=False,\n            **kw\n        ),\n    )\ntrain_loader, val_loader, test_loader = make_loaders(train_data, val_data, test_df)\nprint(f\"\\nBatches → train:{len(train_loader)} val:{len(val_loader)} test:{len(test_loader)}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 6.  ARCFACE HEAD\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass ArcFaceHead(nn.Module):\n    def __init__(self, in_dim: int, n_cls: int, s: float = 30.0, m: float = 0.45):\n        super().__init__()\n        self.s, self.m = s, m\n        self.W     = nn.Parameter(torch.empty(n_cls, in_dim))\n        nn.init.xavier_uniform_(self.W)\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, emb: torch.Tensor,\n                labels: torch.Tensor | None = None) -> torch.Tensor:\n        cos = F.linear(F.normalize(emb), F.normalize(self.W))\n        if labels is None or not self.training:\n            return cos * self.s\n        sin = torch.sqrt((1 - cos.pow(2)).clamp(1e-6))\n        cos_m = cos * self.cos_m - sin * self.sin_m\n        cos_m = torch.where(cos > self.th, cos_m, cos - self.mm)\n        one_hot = torch.zeros_like(cos).scatter_(1, labels.unsqueeze(1), 1.0)\n        return (one_hot * cos_m + (1 - one_hot) * cos) * self.s\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 7.  BACKBONE WRAPPER\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass AttributionModel(nn.Module):\n\n    def __init__(\n        self,\n        backbone_name: str,\n        embed_dim: int,\n        in_channels: int = 4\n    ):\n\n        super().__init__()\n\n        self.in_channels = in_channels\n\n        self.backbone = timm.create_model(\n            backbone_name,\n            pretrained=True,\n            num_classes=0,\n            in_chans=in_channels,\n            drop_path_rate=CFG.DROP_PATH_RATE,\n        )\n\n        # ─────────────────────────────────────────\n        # адаптация только если 4 канала\n        # ─────────────────────────────────────────\n        if in_channels == 4:\n\n            first_conv = None\n\n            for name, module in self.backbone.named_modules():\n\n                if isinstance(module, nn.Conv2d):\n\n                    if module.in_channels == 3:\n\n                        first_conv = module\n                        first_name = name\n                        break\n\n            if first_conv is not None:\n\n                new_conv = nn.Conv2d(\n                    4,\n                    first_conv.out_channels,\n                    kernel_size=first_conv.kernel_size,\n                    stride=first_conv.stride,\n                    padding=first_conv.padding,\n                    bias=(first_conv.bias is not None)\n                )\n\n                with torch.no_grad():\n\n                    new_conv.weight[:, :3] = first_conv.weight\n\n                    new_conv.weight[:, 3:] = (\n                        first_conv.weight.mean(dim=1, keepdim=True)\n                    )\n\n                parent = self.backbone\n                split = first_name.split(\".\")\n\n                for s in split[:-1]:\n                    parent = getattr(parent, s)\n\n                setattr(parent, split[-1], new_conv)\n\n        feat_dim = self.backbone.num_features\n\n        self.neck = nn.Sequential(\n\n            nn.Linear(\n                feat_dim,\n                embed_dim,\n                bias=False\n            ),\n\n            nn.BatchNorm1d(embed_dim),\n\n            nn.Dropout(\n                p=CFG.DROP_RATE\n            ),\n        )\n\n        self.arc_head = ArcFaceHead(\n            embed_dim,\n            CFG.NUM_CLASSES,\n            s=CFG.ARC_S,\n            m=CFG.ARC_M\n        )\n\n        self.ce_head = nn.Linear(\n            embed_dim,\n            CFG.NUM_CLASSES\n        )\n\n    def forward(self, x, labels=None):\n\n        feat = self.backbone(x)\n\n        emb = self.neck(feat)\n\n        arc_logits = self.arc_head(\n            emb,\n            labels\n        )\n\n        ce_logits = self.ce_head(emb)\n\n        return arc_logits, ce_logits, emb\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 8.  LOSS\n# ══════════════════════════════════════════════════════════════════════════════\nimport cv2\nimport numpy as np\n\n# ═══════════════════════════════════════════════════════════════════════\n# DCT SPECTRAL FEATURE\n# ═══════════════════════════════════════════════════════════════════════\n\ndef make_dct_image(img):\n\n    # RGB -> GRAY\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n\n    # resize for stability\n    gray = cv2.resize(\n        gray,\n        (CFG.IMG_SIZE, CFG.IMG_SIZE)\n    )\n\n    # float32\n    gray = np.float32(gray) / 255.0\n\n    # DCT\n    dct = cv2.dct(gray)\n\n    # log spectrum\n    dct = np.log(np.abs(dct) + 1e-6)\n\n    # normalize\n    dct = (dct - dct.min()) / (dct.max() - dct.min() + 1e-6)\n\n    return dct.astype(np.float32)\n# Class weights for mild imbalance handling\ncounts        = train_df[\"y\"].value_counts().sort_index().values.astype(float)\nclass_weights = torch.tensor(counts.sum() / (len(counts) * counts), dtype=torch.float32)\n\nce_criterion  = nn.CrossEntropyLoss(\n    weight=class_weights.to(CFG.DEVICE),\n    label_smoothing=CFG.LABEL_SMOOTH,\n)\n\n\ndef combined_loss(arc_logits, ce_logits, labels):\n    \"\"\"λ·ArcFace + (1-λ)·CE  with label smoothing on CE branch.\"\"\"\n    arc_loss = F.cross_entropy(arc_logits, labels)           # ArcFace already scaled\n    ce_loss  = ce_criterion(ce_logits, labels)\n    return CFG.ARCFACE_W * arc_loss + CFG.CE_W * ce_loss\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 9.  OPTIMIZER & SCHEDULER\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef build_optimizer_scheduler(model: nn.Module, n_batches: int):\n    optimizer = torch.optim.AdamW(\n        model.parameters(), lr=CFG.LR, weight_decay=CFG.WEIGHT_DECAY\n    )\n    total_steps  = CFG.NUM_EPOCHS * n_batches\n    warmup_steps = CFG.WARMUP_EPOCHS * n_batches\n\n    # Linear warmup → Cosine decay\n    def lr_lambda(step: int) -> float:\n        if step < warmup_steps:\n            return float(step) / max(1, warmup_steps)\n        progress = float(step - warmup_steps) / max(1, total_steps - warmup_steps)\n        return max(CFG.LR_MIN / CFG.LR,\n                   0.5 * (1.0 + math.cos(math.pi * progress)))\n\n    scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n    return optimizer, scheduler\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 10.  TRAIN / EVAL LOOP\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef train_one_epoch(model, loader, optimizer, scheduler, device):\n    def random_post_ops(n_ops=None):\n    \n        ops = [\n    \n            A.ImageCompression(\n                quality_range=(40, 90),\n                p=1.0\n            ),\n    \n            A.Downscale(\n                scale_range=(0.5, 0.9),\n                p=1.0\n            ),\n    \n            A.GaussianBlur(\n                blur_limit=(3, 7),\n                p=1.0\n            ),\n    \n            A.MotionBlur(\n                blur_limit=5,\n                p=1.0\n            ),\n    \n            A.ToGray(\n                p=1.0\n            ),\n    \n            A.RandomBrightnessContrast(\n                brightness_limit=0.2,\n                contrast_limit=0.2,\n                p=1.0\n            ),\n    \n            A.Rotate(\n                limit=15,\n                border_mode=cv2.BORDER_REFLECT_101,\n                p=1.0\n            ),\n        ]\n    \n        # случайно выбираем 1-3 операции\n        if n_ops is None:\n            n_ops = random.randint(1, CFG.MAX_POST_OPS)\n    \n        return A.SomeOf(\n            ops,\n            n=n_ops,\n            replace=False,\n            p=1.0\n        )\n    model.train()\n    total_loss = correct = total = 0\n    pbar = tqdm(loader, desc=\"  train\", leave=False)\n\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n\n        arc_l, ce_l, _ = model(imgs, labels)\n        loss = combined_loss(arc_l, ce_l, labels)\n\n        optimizer.zero_grad()\n        loss.backward()\n        nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        scheduler.step()\n\n        total_loss += loss.item() * len(labels)\n        correct    += (ce_l.argmax(1) == labels).sum().item()  # accuracy via CE head\n        total      += len(labels)\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\", acc=f\"{correct/total:.4f}\")\n\n    return total_loss / total, correct / total\n\n\n@torch.no_grad()\ndef evaluate(model, loader, device):\n    model.eval()\n    total_loss = correct = total = 0\n\n    for imgs, labels in tqdm(loader, desc=\"  val  \", leave=False):\n        imgs, labels = imgs.to(device), labels.to(device)\n        arc_l, ce_l, _ = model(imgs)\n        loss = combined_loss(arc_l, ce_l, labels)\n        total_loss += loss.item() * len(labels)\n        correct    += (ce_l.argmax(1) == labels).sum().item()\n        total      += len(labels)\n\n    return total_loss / total, correct / total\n\n\ndef train_model(cfg_model: dict):\n\n    tag = cfg_model[\"tag\"]\n\n    model_type = cfg_model[\"type\"]\n\n    print(f\"\\n{'='*60}\")\n    print(f\"Training: {cfg_model['name']} [{tag}]\")\n    print(f\"Type: {model_type}\")\n    print(f\"{'='*60}\")\n\n    # ─────────────────────────────────────\n    # FFT MODEL\n    # ─────────────────────────────────────\n    if model_type == \"fft\":\n\n        train_loader_local, val_loader_local, _ = make_fft_loaders(\n            train_data,\n            val_data,\n            test_df\n        )\n\n        in_channels = 3\n\n    # ─────────────────────────────────────\n    # RGB + DCT MODELS\n    # ─────────────────────────────────────\n    else:\n\n        train_loader_local, val_loader_local, _ = make_loaders(\n            train_data,\n            val_data,\n            test_df\n        )\n\n        in_channels = 4\n\n    model = AttributionModel(\n        cfg_model[\"name\"],\n        cfg_model[\"embed_dim\"],\n        in_channels=in_channels\n    ).to(CFG.DEVICE)\n\n    n_params = sum(\n        p.numel()\n        for p in model.parameters()\n        if p.requires_grad\n    )\n\n    print(f\"Params: {n_params:,}\")\n\n    optimizer, scheduler = build_optimizer_scheduler(\n        model,\n        len(train_loader_local)\n    )\n\n    best_acc = 0.0\n\n    ckpt_path = CFG.OUTPUT_DIR / f\"best_{tag}.pth\"\n\n    history = []\n\n    for epoch in range(1, CFG.NUM_EPOCHS + 1):\n\n        lr_now = optimizer.param_groups[0][\"lr\"]\n\n        print(f\"\\nEpoch {epoch}/{CFG.NUM_EPOCHS} lr={lr_now:.2e}\")\n\n        tr_loss, tr_acc = train_one_epoch(\n            model,\n            train_loader_local,\n            optimizer,\n            scheduler,\n            CFG.DEVICE\n        )\n\n        vl_loss, vl_acc = evaluate(\n            model,\n            val_loader_local,\n            CFG.DEVICE\n        )\n\n        history.append(dict(\n            epoch=epoch,\n            tr_loss=tr_loss,\n            tr_acc=tr_acc,\n            vl_loss=vl_loss,\n            vl_acc=vl_acc\n        ))\n\n        print(\n            f\"train loss={tr_loss:.4f} acc={tr_acc:.4f} | \"\n            f\"val loss={vl_loss:.4f} acc={vl_acc:.4f}\"\n        )\n\n        if vl_acc > best_acc:\n\n            best_acc = vl_acc\n\n            torch.save(\n                model.state_dict(),\n                ckpt_path\n            )\n\n            print(f\"✓ saved ({best_acc:.4f})\")\n\n    print(f\"Best val acc: {best_acc:.4f}\")\n\n    model.load_state_dict(\n        torch.load(\n            ckpt_path,\n            map_location=CFG.DEVICE\n        )\n    )\n\n    return model, history\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 11.  TTA INFERENCE\n# ══════════════════════════════════════════════════════════════════════════════\n\n\n@torch.no_grad()\ndef predict_with_probs(model, loader):\n\n    model.eval()\n\n    probs_all = []\n    labels_all = []\n\n    for imgs, labels in loader:\n\n        imgs = imgs.to(CFG.DEVICE)\n\n        _, logits, _ = model(imgs)\n\n        probs = torch.softmax(\n            logits,\n            dim=1\n        )\n\n        probs_all.append(\n            probs.detach().cpu()\n        )\n\n        labels_all.append(labels)\n\n    probs_all = torch.cat(probs_all).numpy()\n    labels_all = torch.cat(labels_all).numpy()\n\n    return probs_all, labels_all\n\n\nclass EnsembleMLP(nn.Module):\n\n    def __init__(self, in_dim, n_classes):\n\n        super().__init__()\n\n        self.net = nn.Sequential(\n\n            nn.Linear(in_dim, 128),\n\n            nn.ReLU(),\n\n            nn.Dropout(0.2),\n\n            nn.Linear(128, 64),\n\n            nn.ReLU(),\n\n            nn.Dropout(0.2),\n\n            nn.Linear(64, n_classes)\n        )\n\n    def forward(self, x):\n\n        return self.net(x)\n\n@torch.no_grad()\ndef tta_predict(model, test_df, n_views=5):\n\n    probs_sum = None\n    stored_ids = None\n\n    for view_idx in range(n_views):\n\n        test_ds = SIADataset(\n            test_df,\n            transforms=build_tta_aug(),\n            is_test=True\n        )\n\n        test_loader = DataLoader(\n            test_ds,\n            batch_size=CFG.BATCH_SIZE,\n            shuffle=False,\n            num_workers=CFG.NUM_WORKERS,\n            pin_memory=True\n        )\n\n        current_probs = []\n        current_ids = []\n\n        for imgs, ids in tqdm(\n            test_loader,\n            desc=f\"TTA {view_idx+1}/{n_views}\",\n            leave=False\n        ):\n\n            imgs = imgs.to(CFG.DEVICE)\n\n            _, logits, _ = model(imgs)\n\n            probs = torch.softmax(\n                logits,\n                dim=1\n            )\n\n            current_probs.append(\n                probs.cpu().numpy()\n            )\n\n            current_ids.extend(ids)\n\n        current_probs = np.concatenate(current_probs)\n\n        if probs_sum is None:\n            probs_sum = current_probs\n            stored_ids = current_ids\n        else:\n            probs_sum += current_probs\n\n    probs_mean = probs_sum / n_views\n\n    tta_predict._ids = stored_ids\n\n    return probs_mean\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 12.  TRAIN ALL MODELS\n# ══════════════════════════════════════════════════════════════════════════════\n\ntrained_models = []\nall_histories  = {}\n\nfor cfg_m in CFG.MODELS:\n    model, hist = train_model(cfg_m)\n    trained_models.append((\n    cfg_m[\"tag\"],\n    model,\n    cfg_m[\"type\"]\n    ))\n    all_histories[cfg_m[\"tag\"]] = hist\n\n\n# ── Training curves ───────────────────────────────────────────────────────────\nfig, axes = plt.subplots(len(CFG.MODELS), 2,\n                         figsize=(12, 4 * len(CFG.MODELS)), squeeze=False)\nfor row, (tag, _, _) in enumerate(trained_models):\n    h = pd.DataFrame(all_histories[tag])\n    axes[row][0].plot(h.epoch, h.tr_loss, label=\"train\")\n    axes[row][0].plot(h.epoch, h.vl_loss, label=\"val\")\n    axes[row][0].set_title(f\"{tag} – Loss\"); axes[row][0].legend()\n    axes[row][1].plot(h.epoch, h.tr_acc,  label=\"train\")\n    axes[row][1].plot(h.epoch, h.vl_acc,  label=\"val\")\n    axes[row][1].set_title(f\"{tag} – Accuracy\"); axes[row][1].legend()\nplt.tight_layout()\nplt.savefig(CFG.OUTPUT_DIR / \"training_curves.png\", dpi=120)\nplt.close()\nprint(f\"\\nTraining curves → {CFG.OUTPUT_DIR / 'training_curves.png'}\")\n\n\n# ══════════════════════════════════════════════════════════════════════\n# 13. ENSEMBLE\n# ══════════════════════════════════════════════════════════════════════\n# ═══════════════════════════════════════════════════════\n# BUILD OOF FEATURES\n# ═══════════════════════════════════════════════════════\n\nmeta_features = []\nmeta_labels = None\n\nfor tag, model, model_type in trained_models:\n\n    print(f\"OOF: {tag}\")\n\n    if model_type == \"fft\":\n\n        _, val_loader_local, _ = make_fft_loaders(\n            train_data,\n            val_data,\n            test_df\n        )\n\n    else:\n\n        _, val_loader_local, _ = make_loaders(\n            train_data,\n            val_data,\n            test_df\n        )\n\n    probs, labels = predict_with_probs(\n        model,\n        val_loader_local\n    )\n\n    meta_features.append(probs)\n\n    if meta_labels is None:\n        meta_labels = labels\n\nX_meta = np.concatenate(\n    meta_features,\n    axis=1\n)\n\ny_meta = meta_labels\n\nprint(X_meta.shape)\n\n\nmeta_model = EnsembleMLP(\n    in_dim=X_meta.shape[1],\n    n_classes=CFG.NUM_CLASSES\n).to(CFG.DEVICE)\n\noptimizer = torch.optim.AdamW(\n    meta_model.parameters(),\n    lr=1e-3\n)\n\ncriterion = nn.CrossEntropyLoss()\n\nX_tensor = torch.tensor(\n    X_meta,\n    dtype=torch.float32\n).to(CFG.DEVICE)\n\ny_tensor = torch.tensor(\n    y_meta,\n    dtype=torch.long\n).to(CFG.DEVICE)\n\nmeta_model.train()\n\nfor epoch in range(30):\n\n    optimizer.zero_grad()\n\n    logits = meta_model(X_tensor)\n\n    loss = criterion(\n        logits,\n        y_tensor\n    )\n\n    loss.backward()\n\n    optimizer.step()\n\n    acc = (\n        logits.argmax(1) == y_tensor\n    ).float().mean()\n\n    print(\n        epoch,\n        loss.item(),\n        acc.item()\n    )\n\n\n\nprint(\n    f\"\\nRunning TTA ensemble \"\n    f\"({len(trained_models)} models × {CFG.TTA_STEPS} views)...\"\n)\n\nensemble_probs = None\nstored_ids = None\n\nfor tag, model, model_type in trained_models:\n\n    print(f\"\\nModel: {tag}\")\n\n    probs_sum = None\n\n    for tta_idx in range(CFG.TTA_STEPS):\n\n        # FFT MODEL\n        if model_type == \"fft\":\n\n            test_ds = FFTDataset(\n                test_df,\n                transforms=build_tta_aug(),\n                is_test=True\n            )\n\n        # RGB + DCT MODELS\n        else:\n\n            test_ds = SIADataset(\n                test_df,\n                transforms=build_tta_aug(),\n                is_test=True\n            )\n\n        test_loader = DataLoader(\n            test_ds,\n            batch_size=CFG.BATCH_SIZE,\n            shuffle=False,\n            num_workers=CFG.NUM_WORKERS,\n            pin_memory=True\n        )\n\n        current_probs = []\n        current_ids = []\n\n        for imgs, ids in tqdm(\n            test_loader,\n            desc=f\"{tag} TTA {tta_idx+1}/{CFG.TTA_STEPS}\",\n            leave=False\n        ):\n\n            imgs = imgs.to(CFG.DEVICE)\n\n            _, logits, _ = model(imgs)\n\n            probs = torch.softmax(\n                logits,\n                dim=1\n            )\n\n            current_probs.append(\n                probs.detach().cpu().numpy()\n            )\n\n            current_ids.extend(ids)\n\n        current_probs = np.concatenate(current_probs)\n\n        if probs_sum is None:\n\n            probs_sum = current_probs\n            stored_ids = current_ids\n\n        else:\n\n            probs_sum += current_probs\n\n    probs_mean = probs_sum / CFG.TTA_STEPS\n\n    if ensemble_probs is None:\n        ensemble_probs = probs_mean\n    else:\n        ensemble_probs += probs_mean\n\nmeta_test_features = []\n\nfor tag, model, model_type in trained_models:\n\n    probs_mean = ...\n\n    meta_test_features.append(\n        probs_mean\n    )\n\nX_test_meta = np.concatenate(\n    meta_test_features,\n    axis=1\n)\n\nX_test_meta = torch.tensor(\n    X_test_meta,\n    dtype=torch.float32\n).to(CFG.DEVICE)\n\nmeta_model.eval()\n\nwith torch.no_grad():\n\n    final_logits = meta_model(\n        X_test_meta\n    )\n\nfinal_preds = final_logits.argmax(1).cpu().numpy()\n\nsubmission = pd.DataFrame({\n    \"ID\": stored_ids,\n    \"TARGET\": final_preds\n})\n\nsubmission.to_csv(\n    CFG.SUBMISSION,\n    index=False\n)\n\nprint(f\"\\n✓ Submission saved → {CFG.SUBMISSION}\")\n\nprint(submission.head())\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 14.  VALIDATION ENSEMBLE ACCURACY  (sanity check)\n# ══════════════════════════════════════════════════════════════════════════════\n\nprint(\"\\nEvaluating ensemble on validation set …\")\n\n@torch.no_grad()\ndef val_ensemble_acc(models, val_df):\n    val_ds = SIADataset(val_df, build_val_aug)\n    vl     = DataLoader(val_ds, CFG.BATCH_SIZE, shuffle=False,\n                        num_workers=CFG.NUM_WORKERS, pin_memory=True)\n\n    all_probs, all_labels = None, []\n    for imgs, labels in tqdm(vl, desc=\"  val-ensemble\", leave=False):\n        imgs = imgs.to(CFG.DEVICE)\n        batch_probs = None\n        for _, m in models:\n            m.eval()\n            _, ce_l, _ = m(imgs)\n            p = F.softmax(ce_l, dim=1).cpu().numpy()\n            batch_probs = p if batch_probs is None else batch_probs + p\n        batch_probs /= len(models)\n        all_probs = batch_probs if all_probs is None else np.concatenate([all_probs, batch_probs])\n        all_labels.extend(labels.numpy().tolist())\n\n    preds   = all_probs.argmax(axis=1)\n    acc     = (preds == np.array(all_labels)).mean()\n    return acc\n\nval_acc = val_ensemble_acc(trained_models, val_data)\nprint(f\"  Ensemble val accuracy: {val_acc:.4f}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 15.  CLASS VISUALISATION\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef visualize_two_classes(df, class_a: int = 0, class_b: int = 1, n: int = 5):\n    fig, axes = plt.subplots(2, n, figsize=(3 * n, 7))\n    name_a = source_names.get(class_a, f\"Class {class_a}\")\n    name_b = source_names.get(class_b, f\"Class {class_b}\")\n    for row_i, (cls, name) in enumerate([(class_a, name_a), (class_b, name_b)]):\n        samples = df[df[\"y\"] == cls].sample(n=n, random_state=CFG.SEED)\n        for col_i, (_, rec) in enumerate(samples.iterrows()):\n            ax = axes[row_i][col_i]\n            try:\n                img = Image.open(resolve_path(rec[\"path\"])).convert(\"RGB\")\n            except Exception:\n                img = Image.new(\"RGB\", (224, 224), (200, 200, 200))\n            ax.imshow(img); ax.axis(\"off\")\n            if col_i == 0:\n                ax.set_ylabel(name, fontsize=10, fontweight=\"bold\",\n                              rotation=90, labelpad=6, va=\"center\")\n    fig.suptitle(f\"Samples: '{name_a}'  vs  '{name_b}'\", fontsize=13, y=1.01)\n    plt.tight_layout()\n    out = CFG.OUTPUT_DIR / \"class_samples.png\"\n    plt.savefig(out, bbox_inches=\"tight\", dpi=120); plt.close()\n    print(f\"Visualisation → {out}\")\n\nvisualize_two_classes(train_df, 0, 1)\nprint(\"\\n✓ All done.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-24T12:02:51.293462Z","iopub.execute_input":"2026-05-24T12:02:51.294243Z","iopub.status.idle":"2026-05-24T13:31:16.370653Z","shell.execute_reply.started":"2026-05-24T12:02:51.294200Z","shell.execute_reply":"2026-05-24T13:31:16.369443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"probs_mean.shape == (len(test_df), CFG.NUM_CLASSES)\nmeta_test_features.append(probs_mean)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-24T13:38:33.476862Z","iopub.execute_input":"2026-05-24T13:38:33.478000Z","iopub.status.idle":"2026-05-24T13:38:33.485823Z","shell.execute_reply.started":"2026-05-24T13:38:33.477943Z","shell.execute_reply":"2026-05-24T13:38:33.484757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\n    f\"\\nRunning TTA ensemble \"\n    f\"({len(trained_models)} models × {CFG.TTA_STEPS} views)...\"\n)\n\nensemble_probs = None\nstored_ids = None\n\nmeta_test_features = []\n\nfor tag, model, model_type in trained_models:\n\n    print(f\"\\nModel: {tag}\")\n\n    probs_sum = None\n\n    for tta_idx in range(CFG.TTA_STEPS):\n\n        if model_type == \"fft\":\n\n            test_ds = FFTDataset(\n                test_df,\n                transforms=build_tta_aug(),\n                is_test=True\n            )\n\n        else:\n\n            test_ds = SIADataset(\n                test_df,\n                transforms=build_tta_aug(),\n                is_test=True\n            )\n\n        test_loader = DataLoader(\n            test_ds,\n            batch_size=CFG.BATCH_SIZE,\n            shuffle=False,\n            num_workers=CFG.NUM_WORKERS,\n            pin_memory=True\n        )\n\n        current_probs = []\n        current_ids = []\n\n        for imgs, ids in tqdm(\n            test_loader,\n            desc=f\"{tag} TTA {tta_idx+1}/{CFG.TTA_STEPS}\",\n            leave=False\n        ):\n\n            imgs = imgs.to(CFG.DEVICE)\n\n            with torch.no_grad():\n\n                _, logits, _ = model(imgs)\n\n                probs = torch.softmax(\n                    logits,\n                    dim=1\n                )\n\n            current_probs.append(\n                probs.detach().cpu().numpy()\n            )\n\n            current_ids.extend(ids)\n\n        current_probs = np.concatenate(\n            current_probs,\n            axis=0\n        )\n\n        if probs_sum is None:\n\n            probs_sum = current_probs\n            stored_ids = current_ids\n\n        else:\n\n            probs_sum += current_probs\n\n    # mean TTA probs\n    probs_mean = probs_sum / CFG.TTA_STEPS\n\n    # сохраняем META FEATURES\n    meta_test_features.append(\n        probs_mean\n    )\n\n    # ensemble\n    if ensemble_probs is None:\n        ensemble_probs = probs_mean\n    else:\n        ensemble_probs += probs_mean\n\n# average ensemble\nensemble_probs /= len(trained_models)\n\n# META FEATURES\nX_test_meta = np.concatenate(\n    meta_test_features,\n    axis=1\n)\n\nprint(\"META TEST SHAPE:\", X_test_meta.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-24T13:38:53.555759Z","iopub.execute_input":"2026-05-24T13:38:53.556380Z","iopub.status.idle":"2026-05-24T13:53:51.131103Z","shell.execute_reply.started":"2026-05-24T13:38:53.556354Z","shell.execute_reply":"2026-05-24T13:53:51.130325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nmeta_model.eval()\n\nwith torch.no_grad():\n\n    X_test_tensor = torch.tensor(\n        X_test_meta,\n        dtype=torch.float32\n    ).to(CFG.DEVICE)\n\n    final_logits = meta_model(\n        X_test_tensor\n    )\n\n    final_preds = (\n        final_logits.argmax(1)\n        .cpu()\n        .numpy()\n    )\n\nsubmission = pd.DataFrame({\n    \"ID\": stored_ids,\n    \"TARGET\": final_preds\n})\n\nsubmission.to_csv(\n    CFG.SUBMISSION,\n    index=False\n)\n\nprint(f\"\\n✓ Submission saved → {CFG.SUBMISSION}\")\n\nprint(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-24T13:55:43.573964Z","iopub.execute_input":"2026-05-24T13:55:43.574833Z","iopub.status.idle":"2026-05-24T13:55:43.742390Z","shell.execute_reply.started":"2026-05-24T13:55:43.574800Z","shell.execute_reply":"2026-05-24T13:55:43.741626Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\n    f\"\\nRunning TTA ensemble \"\n    f\"({len(trained_models)} models × {CFG.TTA_STEPS} views)...\"\n)\n\nensemble_probs = None\nstored_ids = None\n\nfor tag, model, model_type in trained_models:\n\n    print(f\"\\nModel: {tag}\")\n\n    probs_sum = None\n\n    for tta_idx in range(CFG.TTA_STEPS):\n\n        # FFT MODEL\n        if model_type == \"fft\":\n\n            test_ds = FFTDataset(\n                test_df,\n                transforms=build_tta_aug(),\n                is_test=True\n            )\n\n        # RGB + DCT MODELS\n        else:\n\n            test_ds = SIADataset(\n                test_df,\n                transforms=build_tta_aug(),\n                is_test=True\n            )\n\n        test_loader = DataLoader(\n            test_ds,\n            batch_size=CFG.BATCH_SIZE,\n            shuffle=False,\n            num_workers=CFG.NUM_WORKERS,\n            pin_memory=True\n        )\n\n        current_probs = []\n        current_ids = []\n\n        for imgs, ids in tqdm(\n            test_loader,\n            desc=f\"{tag} TTA {tta_idx+1}/{CFG.TTA_STEPS}\",\n            leave=False\n        ):\n\n            imgs = imgs.to(CFG.DEVICE)\n\n            _, logits, _ = model(imgs)\n\n            probs = torch.softmax(\n                logits,\n                dim=1\n            )\n\n            current_probs.append(\n                probs.detach().cpu().numpy()\n            )\n\n            current_ids.extend(ids)\n\n        current_probs = np.concatenate(current_probs)\n\n        if probs_sum is None:\n\n            probs_sum = current_probs\n            stored_ids = current_ids\n\n        else:\n\n            probs_sum += current_probs\n\n    probs_mean = probs_sum / CFG.TTA_STEPS\n\n    if ensemble_probs is None:\n        ensemble_probs = probs_mean\n    else:\n        ensemble_probs += probs_mean\n\nensemble_probs /= len(trained_models)\n\nfinal_preds = ensemble_probs.argmax(axis=1)\n\nsubmission = pd.DataFrame({\n    \"ID\": stored_ids,\n    \"TARGET\": final_preds\n})\n\nsubmission.to_csv(\n    CFG.SUBMISSION,\n    index=False\n)\n\nprint(f\"\\n✓ Submission saved → {CFG.SUBMISSION}\")\n\nprint(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T12:03:42.047794Z","iopub.execute_input":"2026-05-22T12:03:42.048088Z","iopub.status.idle":"2026-05-22T12:19:27.104842Z","shell.execute_reply.started":"2026-05-22T12:03:42.048056Z","shell.execute_reply":"2026-05-22T12:19:27.103914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"\\nRunning TTA ensemble ({len(trained_models)} models × {CFG.TTA_STEPS} views) …\")\n\nensemble_probs = None\nstored_ids = None\n\nfor item in trained_models:\n    # Безопасная распаковка\n    if len(item) == 3:\n        tag, model, model_type = item\n    elif len(item) == 2:\n        tag, model = item\n        # Определяем тип по названию или ставим дефолт\n        model_type = \"fft\" if \"fft\" in str(tag).lower() else \"rgb\"\n    else:\n        raise ValueError(f\"Неожиданный формат в trained_models: {item}\")\n\n    print(f\"  Model: {tag}\")\n    model.eval()  # ⬅️ Обязательно для инференса\n    \n    # ─────────────────────────────────────\n    # FFT MODEL\n    # ─────────────────────────────────────\n    if model_type == \"fft\":\n        test_ds = FFTDataset(test_df, transforms=build_tta_aug(), is_test=True)\n        # ─────────────────────────────────────\n    # RGB + DCT MODELS (4 канала)\n    # ─────────────────────────────────────\n    else:\n        # Проверяем, нужна ли модели 4 канала\n        # Обычно EfficientNet/ConvNext в таких задачах ждут 4 канала, если это SIA\n        if model_type == \"dct\" or model_type == \"sia\": \n            # Используйте трансформы, которые добавляют 4-й канал!\n            test_ds = SIADataset(\n                test_df,\n                transforms=build_tta_aug_4ch(),  # <--- Убедитесь, что эта ф-я существует и добавляет 4-й канал\n                is_test=True\n            )\n        else:\n            # Если модель всё-таки 3-канальная (обычный RGB)\n            test_ds = SIADataset(\n                test_df,\n                transforms=build_tta_aug(),      # <--- Обычные трансформы\n                is_test=True\n            )\n\n    test_loader = DataLoader(\n        test_ds,\n        batch_size=CFG.BATCH_SIZE,\n        shuffle=False,\n        num_workers=CFG.NUM_WORKERS,\n        pin_memory=True\n    )\n\n    probs_sum = None\n\n    # ⬇️ Отключаем градиенты для всего TTA цикла\n    with torch.no_grad():\n        for tta_idx in range(CFG.TTA_STEPS):\n            current_probs = []\n            current_ids = []\n\n            for imgs, ids in tqdm(\n                test_loader,\n                desc=f\"{tag} TTA {tta_idx+1}/{CFG.TTA_STEPS}\",\n                leave=False\n            ):\n                imgs = imgs.to(CFG.DEVICE)\n\n                _, logits, _ = model(imgs)\n\n                probs = torch.softmax(logits, dim=1)\n\n                # Теперь .cpu().numpy() сработает без ошибок\n                current_probs.append(probs.cpu().numpy())\n                current_ids.extend(ids)\n\n            current_probs = np.concatenate(current_probs)\n\n            if probs_sum is None:\n                probs_sum = current_probs\n                stored_ids = current_ids\n            else:\n                probs_sum += current_probs\n\n    probs_mean = probs_sum / CFG.TTA_STEPS\n    # ... остальной код без изменений ...\n\n    if ensemble_probs is None:\n        ensemble_probs = probs_mean\n    else:\n        ensemble_probs += probs_mean\n\n# final averaging\nensemble_probs /= len(trained_models)\n\nfinal_preds = ensemble_probs.argmax(axis=1)\n\nsubmission = pd.DataFrame({\n    \"ID\": stored_ids,\n    \"TARGET\": final_preds\n})\n\nsubmission.to_csv(\n    CFG.SUBMISSION,\n    index=False\n)\n\nprint(f\"\\n✓ Submission saved → {CFG.SUBMISSION}\")\n\nprint(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-21T16:04:01.595838Z","iopub.execute_input":"2026-05-21T16:04:01.596621Z","iopub.status.idle":"2026-05-21T16:15:37.436027Z","shell.execute_reply.started":"2026-05-21T16:04:01.596582Z","shell.execute_reply":"2026-05-21T16:15:37.434724Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# FIX tensor(6) -> 6\n\nclean_ids = []\n\nfor x in stored_ids:\n\n    if torch.is_tensor(x):\n        clean_ids.append(x.item())\n    else:\n        clean_ids.append(int(x))\n\nsubmission = pd.DataFrame({\n    \"ID\": clean_ids,\n    \"TARGET\": final_preds.astype(int)\n})\n\nsubmission.to_csv(\n    CFG.SUBMISSION,\n    index=False\n)\n\nprint(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-24T13:55:54.089160Z","iopub.execute_input":"2026-05-24T13:55:54.089723Z","iopub.status.idle":"2026-05-24T13:55:54.104912Z","shell.execute_reply.started":"2026-05-24T13:55:54.089694Z","shell.execute_reply":"2026-05-24T13:55:54.104070Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef tta_predict(model, test_df, n_views=5):\n\n    test_ds = SIADataset(\n        test_df,\n        transforms=build_val_aug(),\n        is_test=True\n    )\n\n    test_loader = DataLoader(\n        test_ds,\n        batch_size=CFG.BATCH_SIZE,\n        shuffle=False,\n        num_workers=CFG.NUM_WORKERS,\n        pin_memory=True\n    )\n\n    probs_sum = []\n    stored_ids = None\n\n    for view_idx in range(n_views):\n\n        current_probs = []\n\n        for imgs, ids in tqdm(\n            test_loader,\n            desc=f\"TTA {view_idx+1}/{n_views}\",\n            leave=False\n        ):\n\n            imgs = imgs.to(CFG.DEVICE)\n\n            _, logits, _ = model(imgs)\n\n            probs = torch.softmax(logits, dim=1)\n\n            current_probs.append(\n                probs.cpu().numpy()\n            )\n\n            if stored_ids is None:\n                stored_ids = list(ids)\n\n        current_probs = np.concatenate(current_probs)\n\n        probs_sum.append(current_probs)\n\n    probs_mean = np.mean(probs_sum, axis=0)\n\n    tta_predict._ids = stored_ids\n\n    return probs_mean","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-17T10:03:34.742054Z","iopub.execute_input":"2026-05-17T10:03:34.742748Z","iopub.status.idle":"2026-05-17T10:03:34.749228Z","shell.execute_reply.started":"2026-05-17T10:03:34.742717Z","shell.execute_reply":"2026-05-17T10:03:34.748431Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"probs.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-17T10:03:43.567511Z","iopub.execute_input":"2026-05-17T10:03:43.568073Z","iopub.status.idle":"2026-05-17T10:03:43.573452Z","shell.execute_reply.started":"2026-05-17T10:03:43.568044Z","shell.execute_reply":"2026-05-17T10:03:43.572666Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"\\nRunning TTA ensemble ({len(trained_models)} models × {CFG.TTA_STEPS} views) …\")\n\nensemble_probs = None\nstored_ids     = None\n\nfor tag, model in trained_models:\n    print(f\"  Model: {tag}\")\n    probs = tta_predict(model, test_df, CFG.TTA_STEPS)   # (N, 10)\n    if stored_ids is None and hasattr(tta_predict, \"_ids\"):\n        stored_ids = tta_predict._ids\n    if ensemble_probs is None:\n        ensemble_probs = probs\n    else:\n        ensemble_probs += probs   # accumulate; divide later\n\nensemble_probs /= len(trained_models)\nfinal_preds = ensemble_probs.argmax(axis=1).tolist()\n\nprint(f\"Ensemble done.  Predictions: {len(final_preds)}\")\n\n\n# ── Validate ID alignment ─────────────────────────────────────────────────────\nif stored_ids is None:\n    stored_ids = test_df[\"ID\"].tolist()\n\nsubmission = pd.DataFrame({\"ID\": stored_ids, \"TARGET\": final_preds})\n\nif CFG.SAMPLE_SUB.exists():\n    sample     = pd.read_csv(CFG.SAMPLE_SUB)\n    submission = sample[[\"ID\"]].merge(submission, on=\"ID\", how=\"left\")\n    missing    = submission[\"TARGET\"].isna().sum()\n    if missing:\n        print(f\"[WARN] {missing} missing predictions → 0\")\n        submission[\"TARGET\"] = submission[\"TARGET\"].fillna(0)\n    submission[\"TARGET\"] = submission[\"TARGET\"].astype(int)\n\nsubmission.to_csv(CFG.SUBMISSION, index=False)\nprint(f\"\\n✓ Submission → {CFG.SUBMISSION}\")\nprint(submission.head(10))\nprint(\"\\nPrediction distribution:\\n\", submission[\"TARGET\"].value_counts().sort_index())\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 14.  VALIDATION ENSEMBLE ACCURACY  (sanity check)\n# ══════════════════════════════════════════════════════════════════════════════\n\nprint(\"\\nEvaluating ensemble on validation set …\")\n\n@torch.no_grad()\ndef val_ensemble_acc(models, val_df):\n    val_ds = SIADataset(val_df, build_val_aug)\n    vl     = DataLoader(val_ds, CFG.BATCH_SIZE, shuffle=False,\n                        num_workers=CFG.NUM_WORKERS, pin_memory=True)\n\n    all_probs, all_labels = None, []\n    for imgs, labels in tqdm(vl, desc=\"  val-ensemble\", leave=False):\n        imgs = imgs.to(CFG.DEVICE)\n        batch_probs = None\n        for _, m in models:\n            m.eval()\n            _, ce_l, _ = m(imgs)\n            p = F.softmax(ce_l, dim=1).cpu().numpy()\n            batch_probs = p if batch_probs is None else batch_probs + p\n        batch_probs /= len(models)\n        all_probs = batch_probs if all_probs is None else np.concatenate([all_probs, batch_probs])\n        all_labels.extend(labels.numpy().tolist())\n\n    preds   = all_probs.argmax(axis=1)\n    acc     = (preds == np.array(all_labels)).mean()\n    return acc\n\nval_acc = val_ensemble_acc(trained_models, val_data)\nprint(f\"  Ensemble val accuracy: {val_acc:.4f}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 15.  CLASS VISUALISATION\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef visualize_two_classes(df, class_a: int = 0, class_b: int = 1, n: int = 5):\n    fig, axes = plt.subplots(2, n, figsize=(3 * n, 7))\n    name_a = source_names.get(class_a, f\"Class {class_a}\")\n    name_b = source_names.get(class_b, f\"Class {class_b}\")\n    for row_i, (cls, name) in enumerate([(class_a, name_a), (class_b, name_b)]):\n        samples = df[df[\"y\"] == cls].sample(n=n, random_state=CFG.SEED)\n        for col_i, (_, rec) in enumerate(samples.iterrows()):\n            ax = axes[row_i][col_i]\n            try:\n                img = Image.open(resolve_path(rec[\"path\"])).convert(\"RGB\")\n            except Exception:\n                img = Image.new(\"RGB\", (224, 224), (200, 200, 200))\n            ax.imshow(img); ax.axis(\"off\")\n            if col_i == 0:\n                ax.set_ylabel(name, fontsize=10, fontweight=\"bold\",\n                              rotation=90, labelpad=6, va=\"center\")\n    fig.suptitle(f\"Samples: '{name_a}'  vs  '{name_b}'\", fontsize=13, y=1.01)\n    plt.tight_layout()\n    out = CFG.OUTPUT_DIR / \"class_samples.png\"\n    plt.savefig(out, bbox_inches=\"tight\", dpi=120); plt.close()\n    print(f\"Visualisation → {out}\")\n\nvisualize_two_classes(train_df, 0, 1)\nprint(\"\\n✓ All done.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-17T10:03:56.268583Z","iopub.execute_input":"2026-05-17T10:03:56.269271Z","iopub.status.idle":"2026-05-17T10:03:56.294453Z","shell.execute_reply.started":"2026-05-17T10:03:56.269198Z","shell.execute_reply":"2026-05-17T10:03:56.293447Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_preds = ensemble_probs.argmax(1)\n\nsubmission = pd.DataFrame({\n    \"ID\": stored_ids,\n    \"TARGET\": final_preds\n})","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nEnsemble Solution v2: Synthetic Image Attribution Challenge\n============================================================\nBackbones:\n  1. ConvNeXt-Small          — spatial CNN,  pretrained ImageNet-22k\n  2. EfficientNet-B2          — efficient CNN, pretrained ImageNet\n  3. DaViT-Tiny               — Dual-Attention ViT (window + channel attn),\n                                 < Swin по склонности к переобучению,\n                                 хорошо ловит локальные артефакты генераторов\n  4. FreqCNN                  — лёгкий (~1.2 M) CNN работающий ТОЛЬКО\n                                 с FFT-спектром (magnitude + phase),\n                                 специализируется на частотных следах\n\nLoss       : ArcFace + CrossEntropy (label_smoothing=0.1)\nAugments   : albumentations pipeline с симуляцией пост-обработки теста\n             (JPEG, resize, crop, rotation, grayscale, blur, brightness,\n              super-res) → 1-3 op combos  +  DCT 4th channel\nOptimizer  : AdamW + linear warmup + cosine decay\nTTA        : 5 views × каждая модель → soft-vote ensemble\n\"\"\"\n\n# ─────────────────────────────────────────────────────────────────────────────\n# 0.  DEPENDENCIES\n#     pip install timm albumentations torch torchvision pandas pillow\n#                 scikit-learn tqdm matplotlib scipy\n# ─────────────────────────────────────────────────────────────────────────────\n\nimport cv2, math, os, random, warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib; matplotlib.use(\"Agg\")\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom PIL import Image\nfrom scipy.fft import dctn\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 1.  CONFIG\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef _find_data_dir() -> Path:\n    for c in [\n        Path(\"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge\"),\n        Path(\"/kaggle/input/dlmmdd-workshop-synthetic-source-attribution-challenge\"),\n        Path(\"/kaggle/input\"),\n        Path(\"./data\"),\n    ]:\n        if c.exists() and any(c.rglob(\"train.csv\")):\n            return c\n    return Path(\"./data\")\n\ndef _find_file(root: Path, name: str) -> Path:\n    hits = list(root.rglob(name))\n    return hits[0] if hits else root / name\n\n_ROOT = _find_data_dir()\n\n\nclass CFG:\n    # ── paths ──────────────────────────────────────────────────────────────────\n    DATA_DIR = Path(\"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data\")\n    TRAIN_CSV   = DATA_DIR / \"training.csv\"\n    TEST_CSV    = DATA_DIR / \"test.csv\"\n    SAMPLE_SUB  = DATA_DIR / \"sample_submission.csv\"\n    SOURCES_TXT = DATA_DIR / \"sources.txt\"\n    OUTPUT_DIR  = Path(\"/kaggle/working\") if Path(\"/kaggle/working\").exists() else Path(\"./outputs\")\n\n    # ── models ─────────────────────────────────────────────────────────────────\n    #\n    #   type=\"rgb\"   → модель получает 3 или 4 канала (RGB + опц. DCT)\n    #   type=\"freq\"  → модель получает только FFT-спектр (2 канала: mag + phase)\n    #\n    #   epochs_override — переопределяет NUM_EPOCHS для конкретной модели\n    #   (DaViT даём меньше эпох, чтобы не переобучился как Swin)\n    #\n    MODELS = [\n        dict(name=\"convnext_small\",   embed_dim=512, tag=\"convnext\",      type=\"rgb\"),\n        dict(name=\"efficientnet_b2\",  embed_dim=512, tag=\"efficientnet\",  type=\"rgb\"),\n        dict(name=\"davit_tiny\",       embed_dim=256, tag=\"davit\",         type=\"rgb\",\n             epochs_override=6),          # DaViT: 6 эп. достаточно, не переобучается\n        dict(name=\"freq_cnn\",         embed_dim=256, tag=\"freq_cnn\",      type=\"freq\",\n             epochs_override=10),         # лёгкая сеть — можно дольше\n    ]\n\n    # ── training ───────────────────────────────────────────────────────────────\n    NUM_CLASSES   = 10\n    IMG_SIZE      = 224\n    BATCH_SIZE    = 32\n    NUM_EPOCHS    = 1          # default (overridable per model)\n    WARMUP_EPOCHS = 2\n    LR            = 3e-4\n    LR_MIN        = 1e-6\n    WEIGHT_DECAY  = 1e-2\n    VAL_SPLIT     = 0.15\n    NUM_WORKERS   = 4\n    SEED          = 42\n    DEVICE        = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    # ── ArcFace ────────────────────────────────────────────────────────────────\n    ARC_S = 30.0\n    ARC_M = 0.45\n\n    # ── loss ───────────────────────────────────────────────────────────────────\n    ARCFACE_W    = 0.5\n    CE_W         = 0.5\n    LABEL_SMOOTH = 0.1\n\n    # ── augmentation ──────────────────────────────────────────────────────────\n    MAX_POST_OPS = 3\n\n    # ── spectral (4th channel for RGB models) ─────────────────────────────────\n    USE_SPECTRAL = True\n    IN_CHANNELS  = 4 if USE_SPECTRAL else 3   # used by RGB models\n\n    # ── TTA ───────────────────────────────────────────────────────────────────\n    TTA_STEPS    = 5\n\n    # ── regularisation ────────────────────────────────────────────────────────\n    DROP_RATE      = 0.3\n    DROP_PATH_RATE = 0.1\n\n    # ── derived ───────────────────────────────────────────────────────────────\n    SUBMISSION = OUTPUT_DIR / \"submission.csv\"\n\n\nCFG.OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\nprint(f\"Data root : {CFG.DATA_DIR}\")\nprint(f\"Device    : {CFG.DEVICE}\")\nprint(f\"Models    : {[m['tag'] for m in CFG.MODELS]}\")\n\n\ndef seed_everything(seed: int):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_everything(CFG.SEED)\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 2.  DATA LOADING\n# ══════════════════════════════════════════════════════════════════════════════\n\ntrain_df = pd.read_csv(CFG.TRAIN_CSV)\ntest_df  = pd.read_csv(CFG.TEST_CSV)\n\nprint(f\"Train {train_df.shape}  Test {test_df.shape}\")\nprint(\"Class distribution:\\n\", train_df[\"y\"].value_counts().sort_index())\n\nsource_names: dict[int, str] = {}\n\n\ntrain_data, val_data = train_test_split(\n    train_df, test_size=CFG.VAL_SPLIT,\n    stratify=train_df[\"y\"], random_state=CFG.SEED,\n)\ntrain_data = train_data.reset_index(drop=True)\nval_data   = val_data.reset_index(drop=True)\nprint(f\"Train {len(train_data)}  Val {len(val_data)}  Test {len(test_df)}\")\n\n# ── file index (решает Data/Data/Data/… пути) ─────────────────────────────────\nfrom pathlib import Path\n\nCFG.DATA_DIR = Path(CFG.DATA_DIR)\n\nprint(\"\\nBuilding file index …\")\n\n_FILE_INDEX: dict[str, Path] = {}\n\nfor ext in (\"*.png\", \"*.jpg\", \"*.jpeg\", \"*.webp\"):\n\n    for p in CFG.DATA_DIR.rglob(ext):\n\n        _FILE_INDEX[p.name] = p\n\nprint(f\"  {len(_FILE_INDEX)} images indexed\")\n\ndef resolve_path(raw: str) -> Path:\n    p = Path(raw)\n    if p.is_absolute() and p.exists(): return p\n    c = CFG.DATA_DIR / p\n    if c.exists(): return c\n    if p.name in _FILE_INDEX: return _FILE_INDEX[p.name]\n    parts = p.parts\n    for s in range(1, len(parts)):\n        c2 = CFG.DATA_DIR / Path(*parts[s:])\n        if c2.exists(): return c2\n    return c\n\n_tp = resolve_path(train_df[\"path\"].iloc[0])\nprint(f\"Path check → {_tp}  exists={_tp.exists()}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 3.  SPECTRAL FEATURES\n# ══════════════════════════════════════════════════════════════════════════════\n\n# ── 3a.  DCT channel  (append as 4th channel to RGB models) ───────────────────\ndef compute_dct_channel(img_np: np.ndarray) -> np.ndarray:\n    \"\"\"\n    2-D DCT of luminance → log-magnitude → normalised float32 (H, W).\n    Captures periodic grid artefacts from VAE/upsampling stages.\n    \"\"\"\n    gray = (0.299 * img_np[:, :, 0] +\n            0.587 * img_np[:, :, 1] +\n            0.114 * img_np[:, :, 2])\n    dct  = dctn(gray.astype(np.float32), norm=\"ortho\")\n    mag  = np.log1p(np.abs(dct))\n    mag  = (mag - mag.min()) / (mag.max() - mag.min() + 1e-8)\n    return mag.astype(np.float32)\n\n\n# ── 3b.  FFT spectrum  (2-channel input for FreqCNN) ──────────────────────────\ndef compute_fft_spectrum(img_np: np.ndarray,\n                          size: int = CFG.IMG_SIZE) -> np.ndarray:\n    \"\"\"\n    Computes 2-D FFT of luminance and returns a 2-channel float32 array\n    (size, size, 2)  →  [log_magnitude, normalised_phase].\n\n    Why FFT for a dedicated branch?\n    --------------------------------\n    • FFT has global support: one coefficient encodes a frequency present\n      everywhere in the image → ideal for detecting diffusion-model noise\n      schedules, GAN frequency bias, and JPEG grid artefacts.\n    • DCT (used for the 4th channel) uses only magnitude;\n      FFT phase adds complementary structural information.\n    • Keeping it in a *separate* small CNN avoids polluting the RGB\n      feature space and lets the model learn frequency patterns freely.\n    \"\"\"\n    gray  = (0.299 * img_np[:, :, 0] +\n             0.587 * img_np[:, :, 1] +\n             0.114 * img_np[:, :, 2]).astype(np.float32) / 255.0\n    gray  = cv2.resize(gray, (size, size))\n\n    fft   = np.fft.fft2(gray)\n    fft_s = np.fft.fftshift(fft)               # center low-freq\n\n    mag   = np.log1p(np.abs(fft_s))\n    mag   = (mag - mag.min()) / (mag.max() - mag.min() + 1e-8)\n\n    phase = np.angle(fft_s) / np.pi            # → [-1, 1]\n    phase = (phase + 1.0) / 2.0               # → [ 0, 1]\n\n    return np.stack([mag, phase], axis=-1).astype(np.float32)  # (H, W, 2)\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 4.  AUGMENTATION PIPELINE\n# ══════════════════════════════════════════════════════════════════════════════\n\n_POST_OPS = [\n    A.ImageCompression(quality_lower=40, quality_upper=85, p=1.0),\n    A.Downscale(scale_min=0.5, scale_max=0.9, p=1.0),\n    A.RandomResizedCrop(CFG.IMG_SIZE, CFG.IMG_SIZE,\n                        scale=(0.6, 1.0), ratio=(0.75, 1.33), p=1.0),\n    A.Rotate(limit=15, p=1.0),\n    A.ToGray(p=1.0),\n    A.GaussianBlur(blur_limit=(3, 7), p=1.0),\n    A.RandomBrightnessContrast(brightness_limit=0.3, contrast_limit=0.3, p=1.0),\n    A.Sharpen(alpha=(0.3, 0.7), lightness=(0.8, 1.2), p=1.0),\n]\n\ndef random_post_ops(max_ops: int = CFG.MAX_POST_OPS) -> A.Compose:\n    k   = random.randint(1, max_ops)\n    ops = random.sample(_POST_OPS, k)\n    return A.Compose(ops)\n\n_MEAN = (0.485, 0.456, 0.406)\n_STD  = (0.229, 0.224, 0.225)\n\n\ndef build_train_aug() -> A.Compose:\n    return A.Compose([\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n        A.HorizontalFlip(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1,\n                           rotate_limit=10, p=0.5),\n        A.OneOf([A.GaussianBlur(blur_limit=(3, 5)),\n                 A.MotionBlur(blur_limit=5)], p=0.3),\n        A.RandomBrightnessContrast(0.2, 0.2, p=0.4),\n        A.HueSaturationValue(10, 20, 10, p=0.3),\n        A.CoarseDropout(max_holes=8, max_height=24, max_width=24, p=0.3),\n        A.Lambda(\n            image=lambda img, **kw: random_post_ops()(image=img)[\"image\"],\n            p=0.7,\n        ),\n        A.Normalize(mean=_MEAN, std=_STD),\n        ToTensorV2(),\n    ])\n\ndef build_val_aug() -> A.Compose:\n    return A.Compose([\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n        A.Normalize(mean=_MEAN, std=_STD),\n        ToTensorV2(),\n    ])\n\ndef build_tta_aug() -> A.Compose:\n    return A.Compose([\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n        A.HorizontalFlip(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.04, scale_limit=0.08,\n                           rotate_limit=8, p=0.5),\n        A.RandomBrightnessContrast(0.1, 0.1, p=0.3),\n        A.Lambda(\n            image=lambda img, **kw: random_post_ops(2)(image=img)[\"image\"],\n            p=0.5,\n        ),\n        A.Normalize(mean=_MEAN, std=_STD),\n        ToTensorV2(),\n    ])\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 5.  DATASET\n#     Два режима через аргумент `input_type`:\n#       \"rgb\"  → tensor (3 или 4 каналов)    — для ConvNeXt / EffNet / DaViT\n#       \"freq\" → tensor (2 канала FFT)        — для FreqCNN\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass SIADataset(Dataset):\n    def __init__(self, df: pd.DataFrame,\n                 aug_fn=None,\n                 is_test: bool = False,\n                 input_type: str = \"rgb\"):      # \"rgb\" | \"freq\"\n        self.df         = df.reset_index(drop=True)\n        self.aug_fn     = aug_fn\n        self.is_test    = is_test\n        self.input_type = input_type\n\n    def __len__(self): return len(self.df)\n\n    def _load(self, raw: str) -> np.ndarray:\n        p = resolve_path(raw)\n        try:\n            return np.array(Image.open(p).convert(\"RGB\"))\n        except Exception as e:\n            print(f\"[WARN] {p}: {e}\")\n            return np.zeros((CFG.IMG_SIZE, CFG.IMG_SIZE, 3), dtype=np.uint8)\n\n    def _make_rgb_tensor(self, img: np.ndarray) -> torch.Tensor:\n        aug = self.aug_fn()\n        t   = aug(image=img)[\"image\"]           # (3, H, W)\n        if CFG.USE_SPECTRAL:\n            dct = torch.from_numpy(compute_dct_channel(img)).unsqueeze(0)\n            dct = (dct - dct.mean()) / (dct.std() + 1e-6)\n            t   = torch.cat([t, dct], dim=0)    # (4, H, W)\n        return t\n\n    def _make_freq_tensor(self, img: np.ndarray) -> torch.Tensor:\n        \"\"\"\n        Для FreqCNN: аугментируем пространственно (чтобы сеть не зависела от\n        положения артефактов), затем считаем FFT от аугментированного RGB.\n        \"\"\"\n        aug       = self.aug_fn()\n        img_aug   = aug(image=img)[\"image\"]     # (3, H, W) нормализованный тензор\n        # Денормализуем → numpy → FFT (FFT считается до нормализации ImageNet)\n        img_raw = np.array(Image.open(\n            resolve_path(self.df.iloc[0][\"path\"])\n        ).convert(\"RGB\")) # fallback — пересчитаем с оригинала\n        # Но лучше просто передать img_aug обратно в numpy:\n        # отменяем нормализацию ImageNet\n        mean_t = torch.tensor(_MEAN).view(3, 1, 1)\n        std_t  = torch.tensor(_STD).view(3, 1, 1)\n        img_denorm = (img_aug * std_t + mean_t).clamp(0, 1)\n        img_u8 = (img_denorm.permute(1, 2, 0).numpy() * 255).astype(np.uint8)\n\n        spec = compute_fft_spectrum(img_u8)          # (H, W, 2)\n        return torch.from_numpy(spec).permute(2, 0, 1)  # (2, H, W)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self._load(row[\"path\"])\n        if self.input_type == \"freq\":\n            x = self._make_freq_tensor(img)\n        else:\n            x = self._make_rgb_tensor(img)\n        if self.is_test:\n            return x, row[\"ID\"]\n        return x, int(row[\"y\"])\n\n\ndef make_loaders(train_df, val_df, test_df, input_type: str = \"rgb\"):\n    tr = SIADataset(train_df, build_train_aug, input_type=input_type)\n    vl = SIADataset(val_df,   build_val_aug,   input_type=input_type)\n    te = SIADataset(test_df,  build_val_aug,   is_test=True, input_type=input_type)\n    kw = dict(num_workers=CFG.NUM_WORKERS, pin_memory=True)\n    return (\n        DataLoader(tr, CFG.BATCH_SIZE, shuffle=True,  **kw),\n        DataLoader(vl, CFG.BATCH_SIZE, shuffle=False, **kw),\n        DataLoader(te, CFG.BATCH_SIZE, shuffle=False, **kw),\n    )\n\n# Загружаем dataloaders для RGB моделей (общие)\ntrain_loader, val_loader, test_loader = make_loaders(train_data, val_data, test_df, \"rgb\")\n# Загружаем dataloaders для FreqCNN (отдельные — другой тип входа)\ntrain_loader_f, val_loader_f, test_loader_f = make_loaders(train_data, val_data, test_df, \"freq\")\n\nprint(f\"RGB  batches → train:{len(train_loader)} val:{len(val_loader)} test:{len(test_loader)}\")\nprint(f\"Freq batches → train:{len(train_loader_f)} val:{len(val_loader_f)} test:{len(test_loader_f)}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 6.  ARCFACE HEAD\n# ══════════════════════════════════════════════════════════════════════════════\n\nclass ArcFaceHead(nn.Module):\n    def __init__(self, in_dim: int, n_cls: int, s: float = 30.0, m: float = 0.45):\n        super().__init__()\n        self.s, self.m = s, m\n        self.W     = nn.Parameter(torch.empty(n_cls, in_dim))\n        nn.init.xavier_uniform_(self.W)\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, emb: torch.Tensor,\n                labels: torch.Tensor | None = None) -> torch.Tensor:\n        cos = F.linear(F.normalize(emb), F.normalize(self.W))\n        if labels is None or not self.training:\n            return cos * self.s\n        sin     = torch.sqrt((1 - cos.pow(2)).clamp(1e-6))\n        cos_m   = cos * self.cos_m - sin * self.sin_m\n        cos_m   = torch.where(cos > self.th, cos_m, cos - self.mm)\n        one_hot = torch.zeros_like(cos).scatter_(1, labels.unsqueeze(1), 1.0)\n        return (one_hot * cos_m + (1 - one_hot) * cos) * self.s\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 7.  MODEL DEFINITIONS\n# ══════════════════════════════════════════════════════════════════════════════\n\n# ── 7a.  Generic RGB backbone (ConvNeXt / EfficientNet / DaViT) ───────────────\nclass AttributionModel(nn.Module):\n    \"\"\"\n    timm backbone → BN+Dropout neck → ArcFace head  +  CE head (shared emb).\n\n    DaViT notes\n    -----------\n    DaViT-Tiny использует попеременно window-attention и channel-attention.\n    Это делает его менее чувствительным к позиционным сдвигам (меньше\n    переобучение на специфику train), чем чистый Swin-T.\n    Опытным путём 6 эпох даёт лучший val_acc без деградации.\n    \"\"\"\n    def __init__(self, backbone_name: str, embed_dim: int):\n        super().__init__()\n        self.backbone = timm.create_model(\n            backbone_name,\n            pretrained=True,\n            num_classes=0,\n            in_chans=CFG.IN_CHANNELS,\n            drop_path_rate=CFG.DROP_PATH_RATE,\n        )\n        feat_dim = self.backbone.num_features\n        self.neck = nn.Sequential(\n            nn.Linear(feat_dim, embed_dim, bias=False),\n            nn.BatchNorm1d(embed_dim),\n            nn.Dropout(p=CFG.DROP_RATE),\n        )\n        self.arc_head = ArcFaceHead(embed_dim, CFG.NUM_CLASSES, CFG.ARC_S, CFG.ARC_M)\n        self.ce_head  = nn.Linear(embed_dim, CFG.NUM_CLASSES)\n\n    def forward(self, x: torch.Tensor,\n                labels: torch.Tensor | None = None):\n        feat  = self.backbone(x)\n        emb   = self.neck(feat)\n        arc_l = self.arc_head(emb, labels)\n        ce_l  = self.ce_head(emb)\n        return arc_l, ce_l, emb\n\n\n# ── 7b.  FreqCNN — специализированная сеть для FFT-спектра ────────────────────\nclass FreqCNN(nn.Module):\n    \"\"\"\n    Лёгкий (~1.2 M параметров) CNN, обученный ТОЛЬКО на 2-канальном\n    FFT-спектре изображения (log-magnitude + phase).\n\n    Архитектура:\n      5 × [Conv3×3 → BN → GELU → MaxPool2×2]\n         с постепенным ростом каналов: 2→32→64→128→256→256\n      Global Average Pool → Dropout → Linear neck → ArcFace + CE head\n\n    Почему отдельная сеть, а не просто 4-й канал?\n    ──────────────────────────────────────────────\n    • FFT — глобальная трансформация: каждый пиксель спектра кодирует\n      паттерн, распределённый по всему изображению. CNN по RGB обрабатывает\n      локальные патчи, поэтому он не может эффективно использовать FFT\n      напрямую — нужна отдельная ветка с другим receptive field.\n    • Phase несёт структурную информацию о фазовой согласованности,\n      которая у реальных и сгенерированных изображений принципиально разная.\n    • Маленькая сеть: быстро обучается (10 эп), не конкурирует по памяти.\n    \"\"\"\n    def __init__(self, embed_dim: int = 256):\n        super().__init__()\n\n        def _block(in_c, out_c, pool: bool = True):\n            layers = [\n                nn.Conv2d(in_c, out_c, kernel_size=3, padding=1, bias=False),\n                nn.BatchNorm2d(out_c),\n                nn.GELU(),\n            ]\n            if pool:\n                layers.append(nn.MaxPool2d(2, 2))\n            return nn.Sequential(*layers)\n\n        self.features = nn.Sequential(\n            _block(2,   32),   # 224 → 112\n            _block(32,  64),   # 112 → 56\n            _block(64,  128),  # 56  → 28\n            _block(128, 256),  # 28  → 14\n            _block(256, 256, pool=False),  # 14 → 14  (deeper without shrinking)\n            nn.AdaptiveAvgPool2d(1),       # → (B, 256, 1, 1)\n        )\n        self.neck = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(256, embed_dim, bias=False),\n            nn.BatchNorm1d(embed_dim),\n            nn.Dropout(p=CFG.DROP_RATE),\n        )\n        self.arc_head = ArcFaceHead(embed_dim, CFG.NUM_CLASSES, CFG.ARC_S, CFG.ARC_M)\n        self.ce_head  = nn.Linear(embed_dim, CFG.NUM_CLASSES)\n\n    def forward(self, x: torch.Tensor,\n                labels: torch.Tensor | None = None):\n        feat  = self.features(x)   # (B, 256, 1, 1) → flatten in neck\n        emb   = self.neck(feat)\n        arc_l = self.arc_head(emb, labels)\n        ce_l  = self.ce_head(emb)\n        return arc_l, ce_l, emb\n\n\ndef build_model(cfg_model: dict) -> nn.Module:\n    \"\"\"Factory: возвращает нужную модель по тегу.\"\"\"\n    if cfg_model[\"tag\"] == \"freq_cnn\":\n        return FreqCNN(embed_dim=cfg_model[\"embed_dim\"]).to(CFG.DEVICE)\n    return AttributionModel(cfg_model[\"name\"], cfg_model[\"embed_dim\"]).to(CFG.DEVICE)\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 8.  LOSS\n# ══════════════════════════════════════════════════════════════════════════════\n\ncounts        = train_df[\"y\"].value_counts().sort_index().values.astype(float)\nclass_weights = torch.tensor(\n    counts.sum() / (len(counts) * counts), dtype=torch.float32\n).to(CFG.DEVICE)\n\nce_criterion = nn.CrossEntropyLoss(\n    weight=class_weights, label_smoothing=CFG.LABEL_SMOOTH\n)\n\ndef combined_loss(arc_logits, ce_logits, labels):\n    arc_loss = F.cross_entropy(arc_logits, labels)\n    ce_loss  = ce_criterion(ce_logits, labels)\n    return CFG.ARCFACE_W * arc_loss + CFG.CE_W * ce_loss\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 9.  OPTIMIZER & SCHEDULER\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef build_optimizer_scheduler(model: nn.Module,\n                               n_batches: int,\n                               n_epochs: int):\n    optimizer    = torch.optim.AdamW(\n        model.parameters(), lr=CFG.LR, weight_decay=CFG.WEIGHT_DECAY\n    )\n    total_steps  = n_epochs * n_batches\n    warmup_steps = CFG.WARMUP_EPOCHS * n_batches\n\n    def lr_lambda(step: int) -> float:\n        if step < warmup_steps:\n            return float(step) / max(1, warmup_steps)\n        progress = float(step - warmup_steps) / max(1, total_steps - warmup_steps)\n        return max(CFG.LR_MIN / CFG.LR,\n                   0.5 * (1.0 + math.cos(math.pi * progress)))\n\n    return optimizer, torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 10.  TRAIN / EVAL\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef train_one_epoch(model, loader, optimizer, scheduler):\n    model.train()\n    total_loss = correct = total = 0\n    pbar = tqdm(loader, desc=\"  train\", leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(CFG.DEVICE), labels.to(CFG.DEVICE)\n        arc_l, ce_l, _ = model(imgs, labels)\n        loss = combined_loss(arc_l, ce_l, labels)\n        optimizer.zero_grad()\n        loss.backward()\n        nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        scheduler.step()\n        total_loss += loss.item() * len(labels)\n        correct    += (ce_l.argmax(1) == labels).sum().item()\n        total      += len(labels)\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\", acc=f\"{correct/total:.4f}\")\n    return total_loss / total, correct / total\n\n\n@torch.no_grad()\ndef evaluate(model, loader):\n    model.eval()\n    total_loss = correct = total = 0\n    for imgs, labels in tqdm(loader, desc=\"  val  \", leave=False):\n        imgs, labels = imgs.to(CFG.DEVICE), labels.to(CFG.DEVICE)\n        arc_l, ce_l, _ = model(imgs)\n        loss = ce_criterion(ce_l, labels)\n        total_loss += loss.item() * len(labels)\n        correct    += (ce_l.argmax(1) == labels).sum().item()\n        total      += len(labels)\n    return total_loss / total, correct / total\n\n\ndef train_model(cfg_model: dict) -> tuple[nn.Module, list]:\n    tag       = cfg_model[\"tag\"]\n    is_freq   = (cfg_model[\"type\"] == \"freq\")\n    n_epochs  = cfg_model.get(\"epochs_override\", CFG.NUM_EPOCHS)\n    tr_loader = train_loader_f if is_freq else train_loader\n    vl_loader = val_loader_f   if is_freq else val_loader\n\n    print(f\"\\n{'='*60}\")\n    print(f\"  Training : {tag}  ({'FreqCNN' if is_freq else cfg_model['name']})\")\n    print(f\"  Epochs   : {n_epochs}  |  Input: {'FFT 2ch' if is_freq else f'RGB+DCT {CFG.IN_CHANNELS}ch'}\")\n    print(f\"{'='*60}\")\n\n    model    = build_model(cfg_model)\n    n_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Params: {n_params:,}\")\n\n    optimizer, scheduler = build_optimizer_scheduler(model, len(tr_loader), n_epochs)\n    best_acc  = 0.0\n    ckpt_path = CFG.OUTPUT_DIR / f\"best_{tag}.pth\"\n    history   = []\n\n    for epoch in range(1, n_epochs + 1):\n        lr_now = optimizer.param_groups[0][\"lr\"]\n        print(f\"\\n  Epoch {epoch}/{n_epochs}  lr={lr_now:.2e}\")\n        tr_loss, tr_acc = train_one_epoch(model, tr_loader, optimizer, scheduler)\n        vl_loss, vl_acc = evaluate(model, vl_loader)\n        history.append(dict(epoch=epoch, tr_loss=tr_loss, tr_acc=tr_acc,\n                            vl_loss=vl_loss, vl_acc=vl_acc))\n        print(f\"  train loss={tr_loss:.4f} acc={tr_acc:.4f} | \"\n              f\"val loss={vl_loss:.4f} acc={vl_acc:.4f}\")\n        if vl_acc > best_acc:\n            best_acc = vl_acc\n            torch.save(model.state_dict(), ckpt_path)\n            print(f\"  ✓ saved ({best_acc:.4f})\")\n\n    print(f\"  Best val acc: {best_acc:.4f}\")\n    model.load_state_dict(torch.load(ckpt_path, map_location=CFG.DEVICE))\n    return model, history\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 11.  TTA INFERENCE\n# ══════════════════════════════════════════════════════════════════════════════\n\n@torch.no_grad()\ndef tta_predict(model: nn.Module, test_df: pd.DataFrame,\n                input_type: str = \"rgb\",\n                n_views: int = CFG.TTA_STEPS) -> tuple[np.ndarray, list]:\n    \"\"\"\n    Returns (N, 10) averaged softmax probs + list of IDs.\n    Each view uses a freshly sampled TTA augmentation.\n    \"\"\"\n    model.eval()\n    probs_sum  = None\n    stored_ids = None\n\n    for v in range(n_views):\n        aug_fn = build_tta_aug if input_type == \"rgb\" else build_val_aug\n        ds     = SIADataset(test_df, aug_fn, is_test=True, input_type=input_type)\n        dl     = DataLoader(ds, CFG.BATCH_SIZE, shuffle=False,\n                            num_workers=CFG.NUM_WORKERS, pin_memory=True)\n\n        view_probs, view_ids = [], []\n        for imgs, ids in tqdm(dl, desc=f\"  TTA {v+1}/{n_views}\", leave=False):\n            imgs = imgs.to(CFG.DEVICE)\n            _, ce_l, _ = model(imgs)\n            view_probs.append(F.softmax(ce_l, dim=1).cpu().numpy())\n            view_ids.extend(ids.numpy().tolist() if isinstance(ids, torch.Tensor)\n                            else list(ids))\n\n        batch = np.concatenate(view_probs, axis=0)\n        probs_sum   = batch if probs_sum is None else probs_sum + batch\n        stored_ids  = view_ids if stored_ids is None else stored_ids\n\n    return probs_sum / n_views, stored_ids\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 12.  TRAIN ALL MODELS\n# ══════════════════════════════════════════════════════════════════════════════\n\ntrained_models: list[tuple[str, nn.Module, str]] = []   # (tag, model, input_type)\nall_histories:  dict[str, list] = {}\n\nfor cfg_m in CFG.MODELS:\n\n    print(f\"\\nPreparing loaders for {cfg_m['tag']}\")\n\n    if cfg_m[\"type\"] == \"fft\":\n\n        train_loader, val_loader, test_loader = make_fft_loaders(\n            train_data,\n            val_data,\n            test_df\n        )\n\n        CFG.IN_CHANNELS = 3\n\n    else:\n\n        train_loader, val_loader, test_loader = make_loaders(\n            train_data,\n            val_data,\n            test_df\n        )\n\n        CFG.IN_CHANNELS = 4 if CFG.USE_SPECTRAL else 3\n\n    model, hist = train_model(cfg_m)\n\n    trained_models.append(\n        (cfg_m[\"tag\"], model, cfg_m[\"type\"])\n    )\n\n    all_histories[cfg_m[\"tag\"]] = hist\n\n\n# ── Training curves ────────────────────────────────────────────────────────────\nn_models = len(CFG.MODELS)\nfig, axes = plt.subplots(n_models, 2, figsize=(12, 4 * n_models), squeeze=False)\nfor row, (tag, _, _) in enumerate(trained_models):\n    h = pd.DataFrame(all_histories[tag])\n    axes[row][0].plot(h.epoch, h.tr_loss, label=\"train\")\n    axes[row][0].plot(h.epoch, h.vl_loss, label=\"val\")\n    axes[row][0].set_title(f\"{tag} – Loss\"); axes[row][0].legend()\n    axes[row][1].plot(h.epoch, h.tr_acc,  label=\"train\")\n    axes[row][1].plot(h.epoch, h.vl_acc,  label=\"val\")\n    axes[row][1].set_title(f\"{tag} – Accuracy\"); axes[row][1].legend()\nplt.tight_layout()\nplt.savefig(CFG.OUTPUT_DIR / \"training_curves.png\", dpi=120)\nplt.close()\nprint(f\"\\nTraining curves → {CFG.OUTPUT_DIR / 'training_curves.png'}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 13.  ENSEMBLE (soft-vote, TTA per model)\n# ══════════════════════════════════════════════════════════════════════════════\n\nprint(f\"\\nRunning TTA ensemble ({len(trained_models)} models × {CFG.TTA_STEPS} views) …\")\n\nensemble_probs: np.ndarray | None = None\nstored_ids: list | None = None\n\nfor tag, model, inp_type in trained_models:\n    print(f\"  [{tag}]  input={inp_type}\")\n    probs, ids = tta_predict(model, test_df, inp_type, CFG.TTA_STEPS)\n    if stored_ids is None:\n        stored_ids = ids\n    ensemble_probs = probs if ensemble_probs is None else ensemble_probs + probs\n\nensemble_probs /= len(trained_models)\nfinal_preds = ensemble_probs.argmax(axis=1).tolist()\nprint(f\"Ensemble done.  {len(final_preds)} predictions.\")\n\nif stored_ids is None:\n    stored_ids = test_df[\"ID\"].tolist()\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 14.  SUBMISSION\n# ══════════════════════════════════════════════════════════════════════════════\n\nclean_ids = [\n    int(x.item()) if isinstance(x, torch.Tensor) else int(x)\n    for x in stored_ids\n]\n\nsubmission = pd.DataFrame({\n    \"ID\": clean_ids,\n    \"TARGET\": final_preds\n})\n\nif CFG.SAMPLE_SUB.exists():\n    sample     = pd.read_csv(CFG.SAMPLE_SUB)\n    submission = sample[[\"ID\"]].merge(submission, on=\"ID\", how=\"left\")\n    missing    = submission[\"TARGET\"].isna().sum()\n    if missing:\n        print(f\"[WARN] {missing} missing → fill 0\")\n        submission[\"TARGET\"] = submission[\"TARGET\"].fillna(0)\n    submission[\"TARGET\"] = submission[\"TARGET\"].astype(int)\n\nsubmission.to_csv(CFG.SUBMISSION, index=False)\nprint(f\"\\n✓ Submission → {CFG.SUBMISSION}\")\nprint(submission.head(10))\nprint(\"\\nPrediction distribution:\\n\", submission[\"TARGET\"].value_counts().sort_index())\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 15.  VALIDATION ENSEMBLE ACCURACY\n# ══════════════════════════════════════════════════════════════════════════════\n\n@torch.no_grad()\ndef val_ensemble_acc() -> float:\n    all_probs: np.ndarray | None = None\n    all_labels: list = []\n\n    # одна проверка по val без TTA\n    for tag, model, inp_type in trained_models:\n        vl = val_loader_f if inp_type == \"freq\" else val_loader\n        model.eval()\n        batch_list: list[np.ndarray] = []\n        labs: list[int] = []\n        for imgs, labels in tqdm(vl, desc=f\"  val [{tag}]\", leave=False):\n            imgs = imgs.to(CFG.DEVICE)\n            _, ce_l, _ = model(imgs)\n            batch_list.append(F.softmax(ce_l, dim=1).cpu().numpy())\n            labs.extend(labels.numpy().tolist())\n        model_probs = np.concatenate(batch_list, axis=0)\n        all_probs   = model_probs if all_probs is None else all_probs + model_probs\n        if not all_labels:\n            all_labels = labs  # first model sets the labels\n\n    all_probs /= len(trained_models)\n    acc = (all_probs.argmax(1) == np.array(all_labels)).mean()\n    return float(acc)\n\nval_acc = val_ensemble_acc()\nprint(f\"\\nEnsemble val accuracy: {val_acc:.4f}\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# 16.  CLASS VISUALISATION\n# ══════════════════════════════════════════════════════════════════════════════\n\ndef visualize_two_classes(df: pd.DataFrame,\n                           class_a: int = 0,\n                           class_b: int = 1,\n                           n: int = 5):\n    fig, axes = plt.subplots(2, n, figsize=(3 * n, 7))\n    name_a = source_names.get(class_a, f\"Class {class_a}\")\n    name_b = source_names.get(class_b, f\"Class {class_b}\")\n    for row_i, (cls, name) in enumerate([(class_a, name_a), (class_b, name_b)]):\n        subset = df[df[\"y\"] == cls]\n        samples = subset.sample(\n            n=min(n, len(subset)),\n            random_state=CFG.SEED\n        )\n        for col_i, (_, rec) in enumerate(samples.iterrows()):\n            ax = axes[row_i][col_i]\n            try:\n                img = Image.open(resolve_path(rec[\"path\"])).convert(\"RGB\")\n            except Exception:\n                img = Image.new(\"RGB\", (224, 224), (200, 200, 200))\n            ax.imshow(img); ax.axis(\"off\")\n            if col_i == 0:\n                ax.set_ylabel(name, fontsize=10, fontweight=\"bold\",\n                              rotation=90, labelpad=6, va=\"center\")\n    fig.suptitle(f\"Samples: '{name_a}'  vs  '{name_b}'\", fontsize=13, y=1.01)\n    plt.tight_layout()\n    out = CFG.OUTPUT_DIR / \"class_samples.png\"\n    plt.savefig(out, bbox_inches=\"tight\", dpi=120); plt.close()\n    print(f\"Visualisation → {out}\")\n\nvisualize_two_classes(train_df, 0, 1)\nprint(\"\\n✓ All done.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-19T07:10:18.807956Z","iopub.execute_input":"2026-05-19T07:10:18.808607Z","iopub.status.idle":"2026-05-19T07:10:19.185306Z","shell.execute_reply.started":"2026-05-19T07:10:18.808571Z","shell.execute_reply":"2026-05-19T07:10:19.184316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# DLMMDD Workshop — Synthetic Source Attribution\n# FULL FIXED PIPELINE\n# =============================================================================\n\nimport os\nimport gc\nimport cv2\nimport math\nimport timm\nimport torch\nimport random\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport albumentations as A\nimport matplotlib.pyplot as plt\n\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nfrom pathlib import Path\nfrom scipy.fft import dctn\n\nfrom sklearn.model_selection import train_test_split\n\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom albumentations.pytorch import ToTensorV2\n\nwarnings.filterwarnings(\"ignore\")\n\n# =============================================================================\n# 1. CONFIG\n# =============================================================================\n\ndef find_data_root():\n    candidates = [\n        Path(\"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge\"),\n        Path(\"/kaggle/input/dlmmdd-workshop-synthetic-source-attribution-challenge\"),\n        Path(\"/kaggle/input\"),\n        Path(\"./data\"),\n    ]\n\n    for c in candidates:\n        if c.exists():\n            csvs = list(c.rglob(\"training.csv\"))\n            if len(csvs) > 0:\n                return csvs[0].parent\n\n    return Path(\"./data\")\n\n\nDATA_ROOT = find_data_root()\n\n\nclass CFG:\n    # =========================\n    # PATHS\n    # =========================\n    DATA_DIR = Path(DATA_ROOT)\n\n    TRAIN_CSV = DATA_DIR / \"training.csv\"\n    TEST_CSV = DATA_DIR / \"test.csv\"\n    SAMPLE_SUB = DATA_DIR / \"sample_submission.csv\"\n\n    OUTPUT_DIR = Path(\"/kaggle/working\")\n\n    # =========================\n    # TRAIN\n    # =========================\n    IMG_SIZE = 224\n    BATCH_SIZE = 32\n    NUM_WORKERS = 4\n\n    NUM_CLASSES = 10\n\n    NUM_EPOCHS = 6\n\n    LR = 3e-4\n    LR_MIN = 1e-6\n    WEIGHT_DECAY = 1e-2\n\n    VAL_SPLIT = 0.15\n\n    SEED = 42\n\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    # =========================\n    # REGULARIZATION\n    # =========================\n    DROP_RATE = 0.3\n    DROP_PATH_RATE = 0.1\n\n    # =========================\n    # ARCFACE\n    # =========================\n    ARC_S = 30.0\n    ARC_M = 0.45\n\n    # =========================\n    # LOSS\n    # =========================\n    ARCFACE_W = 0.5\n    CE_W = 0.5\n    LABEL_SMOOTH = 0.1\n\n    # =========================\n    # SPECTRAL\n    # =========================\n    USE_SPECTRAL = True\n    IN_CHANNELS = 4 if USE_SPECTRAL else 3\n\n    # =========================\n    # TTA\n    # =========================\n    TTA_STEPS = 5\n\n    # =========================\n    # MODELS\n    # =========================\n    MODELS = [\n        dict(\n            name=\"convnext_small\",\n            embed_dim=512,\n            tag=\"convnext\",\n            type=\"rgb\",\n            epochs=6,\n        ),\n\n        dict(\n            name=\"efficientnet_b2\",\n            embed_dim=512,\n            tag=\"efficientnet\",\n            type=\"rgb\",\n            epochs=6,\n        ),\n\n        dict(\n            name=\"davit_tiny\",\n            embed_dim=256,\n            tag=\"davit\",\n            type=\"rgb\",\n            epochs=4,\n        ),\n\n        dict(\n            name=\"freq_cnn\",\n            embed_dim=256,\n            tag=\"freq_cnn\",\n            type=\"freq\",\n            epochs=8,\n        ),\n    ]\n\n    SUBMISSION = OUTPUT_DIR / \"submission.csv\"\n\n\nCFG.OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(\"DATA:\", CFG.DATA_DIR)\nprint(\"DEVICE:\", CFG.DEVICE)\n\n\n# =============================================================================\n# 2. SEED\n# =============================================================================\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nseed_everything(CFG.SEED)\n\n\n# =============================================================================\n# 3. DATA\n# =============================================================================\n\ntrain_df = pd.read_csv(CFG.TRAIN_CSV)\ntest_df = pd.read_csv(CFG.TEST_CSV)\n\nprint(train_df.shape, test_df.shape)\n\ntrain_data, val_data = train_test_split(\n    train_df,\n    test_size=CFG.VAL_SPLIT,\n    stratify=train_df[\"y\"],\n    random_state=CFG.SEED,\n)\n\ntrain_data = train_data.reset_index(drop=True)\nval_data = val_data.reset_index(drop=True)\n\nprint(len(train_data), len(val_data))\n\n\n# =============================================================================\n# 4. FILE INDEX\n# =============================================================================\n\nprint(\"Building file index...\")\n\nFILE_INDEX = {}\n\nfor ext in [\"*.png\", \"*.jpg\", \"*.jpeg\", \"*.webp\"]:\n    for p in CFG.DATA_DIR.rglob(ext):\n        FILE_INDEX[p.name] = p\n\nprint(\"Indexed:\", len(FILE_INDEX))\n\n\ndef resolve_path(raw_path):\n\n    raw_path = str(raw_path)\n\n    p = Path(raw_path)\n\n    if p.exists():\n        return p\n\n    candidate = CFG.DATA_DIR / p\n\n    if candidate.exists():\n        return candidate\n\n    if p.name in FILE_INDEX:\n        return FILE_INDEX[p.name]\n\n    for part_idx in range(len(p.parts)):\n        candidate = CFG.DATA_DIR.joinpath(*p.parts[part_idx:])\n\n        if candidate.exists():\n            return candidate\n\n    return candidate\n\n\ntest_path = resolve_path(train_df.iloc[0][\"path\"])\n\nprint(test_path)\nprint(test_path.exists())\n\n\n# =============================================================================\n# 5. SPECTRAL FEATURES\n# =============================================================================\n\ndef compute_dct_channel(img):\n\n    img = cv2.resize(\n        img,\n        (CFG.IMG_SIZE, CFG.IMG_SIZE)\n    )\n\n    gray = (\n        0.299 * img[:, :, 0]\n        + 0.587 * img[:, :, 1]\n        + 0.114 * img[:, :, 2]\n    )\n\n    gray = gray.astype(np.float32)\n\n    dct = dctn(gray, norm=\"ortho\")\n\n    mag = np.log1p(np.abs(dct))\n\n    mag = (mag - mag.min()) / (\n        mag.max() - mag.min() + 1e-8\n    )\n\n    return mag.astype(np.float32)\n\n\ndef compute_fft_spectrum(img, size=224):\n\n    gray = (\n        0.299 * img[:, :, 0]\n        + 0.587 * img[:, :, 1]\n        + 0.114 * img[:, :, 2]\n    )\n\n    gray = cv2.resize(gray, (size, size))\n    gray = gray.astype(np.float32) / 255.0\n\n    fft = np.fft.fft2(gray)\n    fft = np.fft.fftshift(fft)\n\n    mag = np.log1p(np.abs(fft))\n    mag = (mag - mag.min()) / (mag.max() - mag.min() + 1e-8)\n\n    phase = np.angle(fft)\n    phase = (phase + np.pi) / (2 * np.pi)\n\n    spec = np.stack([mag, phase], axis=-1)\n\n    return spec.astype(np.float32)\n\n\n# =============================================================================\n# 6. AUGS\n# =============================================================================\n\nMEAN = (0.485, 0.456, 0.406)\nSTD = (0.229, 0.224, 0.225)\n\nPOST_OPS = [\n    A.ImageCompression(quality_range=(40, 85), p=1),\n    A.Downscale(scale_range=(0.5, 0.9), p=1),\n    A.Rotate(limit=15, p=1),\n    A.GaussianBlur(blur_limit=(3, 7), p=1),\n    A.RandomBrightnessContrast(0.3, 0.3, p=1),\n]\n\n\ndef random_post_ops():\n\n    k = random.randint(1, 3)\n\n    ops = random.sample(POST_OPS, k)\n\n    return A.Compose(ops)\n\n\ndef build_train_aug():\n\n    return A.Compose([\n\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n\n        A.HorizontalFlip(p=0.5),\n\n        A.ShiftScaleRotate(\n            shift_limit=0.05,\n            scale_limit=0.1,\n            rotate_limit=10,\n            p=0.5,\n        ),\n\n        A.OneOf([\n            A.GaussianBlur(blur_limit=(3, 5)),\n            A.MotionBlur(blur_limit=5),\n        ], p=0.3),\n\n        A.RandomBrightnessContrast(0.2, 0.2, p=0.4),\n\n        A.HueSaturationValue(10, 20, 10, p=0.3),\n\n        A.CoarseDropout(\n            num_holes_range=(1, 8),\n            hole_height_range=(8, 24),\n            hole_width_range=(8, 24),\n            p=0.3,\n        ),\n\n        A.Lambda(\n            image=lambda x, **k: random_post_ops()(image=x)[\"image\"],\n            p=0.7,\n        ),\n\n        A.Normalize(MEAN, STD),\n\n        ToTensorV2(),\n    ])\n\n\ndef build_val_aug():\n\n    return A.Compose([\n\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n\n        A.Normalize(MEAN, STD),\n\n        ToTensorV2(),\n    ])\n\n\n# =============================================================================\n# 7. DATASET\n# =============================================================================\n\nclass SIADataset(Dataset):\n\n    def __init__(\n        self,\n        df,\n        aug_fn,\n        is_test=False,\n        input_type=\"rgb\",\n    ):\n\n        self.df = df.reset_index(drop=True)\n\n        self.aug_fn = aug_fn\n\n        self.is_test = is_test\n\n        self.input_type = input_type\n\n    def __len__(self):\n        return len(self.df)\n\n    def load_image(self, path):\n\n        path = resolve_path(path)\n\n        try:\n            img = Image.open(path).convert(\"RGB\")\n            img = np.array(img)\n\n        except:\n            img = np.zeros((224, 224, 3), dtype=np.uint8)\n\n        return img\n\n    def make_rgb(self, img):\n    \n        aug = self.aug_fn()\n    \n        tensor = aug(image=img)[\"image\"]   # (3,224,224)\n    \n        if CFG.USE_SPECTRAL:\n    \n            # resize DCT to IMG_SIZE\n            dct_channel = compute_dct_channel(img)\n    \n            dct_channel = cv2.resize(\n                dct_channel,\n                (CFG.IMG_SIZE, CFG.IMG_SIZE)\n            )\n    \n            dct_channel = torch.tensor(\n                dct_channel,\n                dtype=torch.float32\n            ).unsqueeze(0)   # (1,224,224)\n    \n            tensor = torch.cat(\n                [tensor, dct_channel],\n                dim=0\n            )   # (4,224,224)\n    \n        return tensor\n\n    def make_freq(self, img):\n\n        spec = compute_fft_spectrum(img)\n\n        tensor = torch.tensor(spec).permute(2, 0, 1)\n\n        return tensor\n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n\n        img = self.load_image(row[\"path\"])\n\n        if self.input_type == \"freq\":\n            x = self.make_freq(img)\n        else:\n            x = self.make_rgb(img)\n\n        if self.is_test:\n\n            sample_id = row[\"ID\"]\n\n            if torch.is_tensor(sample_id):\n                sample_id = sample_id.item()\n\n            return x, int(sample_id)\n\n        y = int(row[\"y\"])\n\n        return x, y\n\n\n# =============================================================================\n# 8. LOADERS\n# =============================================================================\n\ndef make_loaders(input_type):\n\n    train_ds = SIADataset(\n        train_data,\n        build_train_aug,\n        input_type=input_type,\n    )\n\n    val_ds = SIADataset(\n        val_data,\n        build_val_aug,\n        input_type=input_type,\n    )\n\n    test_ds = SIADataset(\n        test_df,\n        build_val_aug,\n        is_test=True,\n        input_type=input_type,\n    )\n\n    loader_args = dict(\n        batch_size=CFG.BATCH_SIZE,\n        num_workers=CFG.NUM_WORKERS,\n        pin_memory=True,\n    )\n\n    train_loader = DataLoader(\n        train_ds,\n        shuffle=True,\n        **loader_args,\n    )\n\n    val_loader = DataLoader(\n        val_ds,\n        shuffle=False,\n        **loader_args,\n    )\n\n    test_loader = DataLoader(\n        test_ds,\n        shuffle=False,\n        **loader_args,\n    )\n\n    return train_loader, val_loader, test_loader\n\n\ntrain_loader_rgb, val_loader_rgb, test_loader_rgb = make_loaders(\"rgb\")\n\ntrain_loader_freq, val_loader_freq, test_loader_freq = make_loaders(\"freq\")\n\n\n# =============================================================================\n# 9. ARCFACE\n# =============================================================================\n\nclass ArcFaceHead(nn.Module):\n\n    def __init__(self, in_features, out_features):\n\n        super().__init__()\n\n        self.W = nn.Parameter(torch.FloatTensor(out_features, in_features))\n\n        nn.init.xavier_uniform_(self.W)\n\n    def forward(self, x):\n\n        x = F.normalize(x)\n\n        W = F.normalize(self.W)\n\n        return F.linear(x, W)\n\n\n# =============================================================================\n# 10. RGB MODEL\n# =============================================================================\n\nclass AttributionModel(nn.Module):\n\n    def __init__(self, backbone_name, embed_dim):\n\n        super().__init__()\n\n        self.backbone = timm.create_model(\n            backbone_name,\n            pretrained=True,\n            num_classes=0,\n            in_chans=3,\n            drop_path_rate=CFG.DROP_PATH_RATE,\n        )\n\n        feat_dim = self.backbone.num_features\n\n        self.neck = nn.Sequential(\n\n            nn.Linear(feat_dim, embed_dim),\n\n            nn.BatchNorm1d(embed_dim),\n\n            nn.Dropout(CFG.DROP_RATE),\n        )\n\n        self.arc = ArcFaceHead(embed_dim, CFG.NUM_CLASSES)\n\n        self.fc = nn.Linear(embed_dim, CFG.NUM_CLASSES)\n\n    def forward(self, x):\n\n        feat = self.backbone(x)\n\n        emb = self.neck(feat)\n\n        logits = self.fc(emb)\n\n        return logits\n\n\n# =============================================================================\n# 11. FREQ CNN\n# =============================================================================\n\nclass FreqCNN(nn.Module):\n\n    def __init__(self, embed_dim=256):\n\n        super().__init__()\n\n        self.features = nn.Sequential(\n\n            nn.Conv2d(2, 32, 3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.GELU(),\n            nn.MaxPool2d(2),\n\n            nn.Conv2d(32, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.GELU(),\n            nn.MaxPool2d(2),\n\n            nn.Conv2d(64, 128, 3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.GELU(),\n            nn.MaxPool2d(2),\n\n            nn.Conv2d(128, 256, 3, padding=1),\n            nn.BatchNorm2d(256),\n            nn.GELU(),\n\n            nn.AdaptiveAvgPool2d(1),\n        )\n\n        self.head = nn.Sequential(\n\n            nn.Flatten(),\n\n            nn.Linear(256, embed_dim),\n\n            nn.BatchNorm1d(embed_dim),\n\n            nn.Dropout(CFG.DROP_RATE),\n\n            nn.Linear(embed_dim, CFG.NUM_CLASSES),\n        )\n\n    def forward(self, x):\n\n        x = self.features(x)\n\n        x = self.head(x)\n\n        return x\n\n\n# =============================================================================\n# 12. BUILD MODEL\n# =============================================================================\n\ndef build_model(cfg_model):\n\n    if cfg_model[\"type\"] == \"freq\":\n\n        model = FreqCNN(cfg_model[\"embed_dim\"])\n\n    else:\n\n        model = AttributionModel(\n            cfg_model[\"name\"],\n            cfg_model[\"embed_dim\"],\n        )\n\n    return model.to(CFG.DEVICE)\n\n\n# =============================================================================\n# 13. LOSS\n# =============================================================================\n\nclass_counts = train_df[\"y\"].value_counts().sort_index().values\n\nweights = class_counts.sum() / (\n    len(class_counts) * class_counts\n)\n\nweights = torch.tensor(\n    weights,\n    dtype=torch.float32,\n).to(CFG.DEVICE)\n\ncriterion = nn.CrossEntropyLoss(\n    weight=weights,\n    label_smoothing=CFG.LABEL_SMOOTH,\n)\n\n\n# =============================================================================\n# 14. TRAIN\n# =============================================================================\n\ndef train_epoch(model, loader, optimizer):\n\n    model.train()\n\n    total_loss = 0\n    correct = 0\n    total = 0\n\n    bar = tqdm(loader)\n\n    for images, labels in bar:\n\n        images = images.to(CFG.DEVICE)\n        labels = labels.to(CFG.DEVICE)\n\n        optimizer.zero_grad()\n\n        logits = model(images)\n\n        loss = criterion(logits, labels)\n\n        loss.backward()\n\n        nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n\n        optimizer.step()\n\n        total_loss += loss.item() * labels.size(0)\n\n        preds = logits.argmax(1)\n\n        correct += (preds == labels).sum().item()\n\n        total += labels.size(0)\n\n        bar.set_postfix(\n            loss=f\"{loss.item():.4f}\",\n            acc=f\"{correct/total:.4f}\",\n        )\n\n    return total_loss / total, correct / total\n\n\n@torch.no_grad()\ndef valid_epoch(model, loader):\n\n    model.eval()\n\n    total_loss = 0\n    correct = 0\n    total = 0\n\n    for images, labels in tqdm(loader):\n\n        images = images.to(CFG.DEVICE)\n        labels = labels.to(CFG.DEVICE)\n\n        logits = model(images)\n\n        loss = criterion(logits, labels)\n\n        total_loss += loss.item() * labels.size(0)\n\n        preds = logits.argmax(1)\n\n        correct += (preds == labels).sum().item()\n\n        total += labels.size(0)\n\n    return total_loss / total, correct / total\n\n\n# =============================================================================\n# 15. TRAIN ALL MODELS\n# =============================================================================\n\ntrained_models = []\n\nfor cfg_model in CFG.MODELS:\n\n    print(\"\\n\" + \"=\" * 60)\n    print(cfg_model[\"tag\"])\n    print(\"=\" * 60)\n\n    if cfg_model[\"type\"] == \"freq\":\n\n        train_loader = train_loader_freq\n        val_loader = val_loader_freq\n\n    else:\n\n        train_loader = train_loader_rgb\n        val_loader = val_loader_rgb\n\n    model = build_model(cfg_model)\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=CFG.LR,\n        weight_decay=CFG.WEIGHT_DECAY,\n    )\n\n    best_acc = 0\n\n    save_path = CFG.OUTPUT_DIR / f'{cfg_model[\"tag\"]}.pth'\n\n    for epoch in range(cfg_model[\"epochs\"]):\n\n        print(f\"\\nEpoch {epoch+1}\")\n\n        tr_loss, tr_acc = train_epoch(\n            model,\n            train_loader,\n            optimizer,\n        )\n\n        vl_loss, vl_acc = valid_epoch(\n            model,\n            val_loader,\n        )\n\n        print(\n            f\"train={tr_acc:.4f} \"\n            f\"val={vl_acc:.4f}\"\n        )\n\n        if vl_acc > best_acc:\n\n            best_acc = vl_acc\n\n            torch.save(\n                model.state_dict(),\n                save_path,\n            )\n\n            print(\"saved\")\n\n    model.load_state_dict(\n        torch.load(save_path, map_location=CFG.DEVICE)\n    )\n\n    trained_models.append(\n        (\n            cfg_model[\"tag\"],\n            cfg_model[\"type\"],\n            model,\n        )\n    )\n\n    gc.collect()\n    torch.cuda.empty_cache()\n\n\n# =============================================================================\n# 16. TTA\n# =============================================================================\n\n@torch.no_grad()\ndef predict_tta(model, loader):\n\n    model.eval()\n\n    probs = []\n    ids_all = []\n\n    for images, ids in tqdm(loader):\n\n        images = images.to(CFG.DEVICE)\n\n        logits = model(images)\n\n        prob = F.softmax(logits, dim=1)\n\n        probs.append(prob.cpu().numpy())\n\n        if torch.is_tensor(ids):\n            ids = ids.cpu().numpy().tolist()\n\n        ids = [int(x) for x in ids]\n\n        ids_all.extend(ids)\n\n    probs = np.concatenate(probs)\n\n    return probs, ids_all\n\n\n# =============================================================================\n# 17. ENSEMBLE\n# =============================================================================\n\nensemble_probs = None\nstored_ids = None\n\nfor tag, input_type, model in trained_models:\n\n    print(\"Predict:\", tag)\n\n    if input_type == \"freq\":\n        test_loader = test_loader_freq\n    else:\n        test_loader = test_loader_rgb\n\n    probs_sum = None\n\n    for _ in range(CFG.TTA_STEPS):\n\n        probs, ids = predict_tta(model, test_loader)\n\n        if probs_sum is None:\n            probs_sum = probs\n        else:\n            probs_sum += probs\n\n    probs_sum /= CFG.TTA_STEPS\n\n    if ensemble_probs is None:\n        ensemble_probs = probs_sum\n    else:\n        ensemble_probs += probs_sum\n\n    if stored_ids is None:\n        stored_ids = ids\n\nensemble_probs /= len(trained_models)\n\npreds = ensemble_probs.argmax(1)\n\n\n# =============================================================================\n# 18. SUBMISSION\n# =============================================================================\n\nstored_ids = [int(x) for x in stored_ids]\npreds = [int(x) for x in preds]\n\nsubmission = pd.DataFrame({\n    \"ID\": stored_ids,\n    \"TARGET\": preds,\n})\n\nsubmission[\"ID\"] = submission[\"ID\"].astype(int)\nsubmission[\"TARGET\"] = submission[\"TARGET\"].astype(int)\n\nsubmission.to_csv(\n    CFG.SUBMISSION,\n    index=False,\n)\n\nprint(\"\\nSUBMISSION SAVED:\")\nprint(CFG.SUBMISSION)\n\nprint(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-21T13:59:16.139640Z","iopub.execute_input":"2026-05-21T13:59:16.139950Z","iopub.status.idle":"2026-05-21T14:00:19.184451Z","shell.execute_reply.started":"2026-05-21T13:59:16.139927Z","shell.execute_reply":"2026-05-21T14:00:19.182610Z"}},"outputs":[],"execution_count":null}]}