{"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":"gpu","dataSources":[{"sourceId":20270,"databundleVersionId":1222630,"sourceType":"competition"},{"sourceId":14487666,"sourceType":"datasetVersion","datasetId":9253460}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# <center> Computer Vision - Train model </center>","metadata":{}},{"cell_type":"markdown","source":"## Import","metadata":{}},{"cell_type":"code","source":"import os, json\nfrom pathlib import Path\nfrom dataclasses import dataclass\nfrom typing import Optional, List, Dict, Tuple\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.amp import autocast, GradScaler\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom tqdm.auto import tqdm\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom sklearn.metrics import (\n    roc_auc_score, average_precision_score,\n    accuracy_score, f1_score, precision_score, recall_score,\n    roc_curve\n)\n\nimport timm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-13T16:55:51.963569Z","iopub.execute_input":"2026-01-13T16:55:51.964317Z","iopub.status.idle":"2026-01-13T16:55:51.970992Z","shell.execute_reply.started":"2026-01-13T16:55:51.964279Z","shell.execute_reply":"2026-01-13T16:55:51.970182Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"@dataclass\nclass Config:\n    cache_root: str = Path(\"/kaggle/input/melanoma2020/siim_isic_cache\")\n    fp_dir: str | None = None # if None, auto-pick first fp_* folder\n\n    backbone: str = \"resnet18\"\n    pretrained: bool = True\n\n    image_size = 384\n    batch_size = 32\n    num_workers = 4\n\n    epochs: int = 5\n    lr: float = 3e-4\n    wd: float = 1e-2\n\n    # meta token encoder\n    meta_d_model: int = 256\n    meta_heads: int = 8\n    meta_layers: int = 2\n    meta_dropout: float = 0.1\n    meta_use_cls: bool = True\n    meta_out_dim: int = 256\n\n    # fusion head\n    fusion_hidden: int = 256\n    fusion_dropout: float = 0.2\n\n    # augmentation\n    brightness_contrast_p: float = 0.2\n\n    # training\n    amp: bool = True\n    seed: int = 42\n    device: str = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\ndef set_seed(seed: int) -> None:\n    import random\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\ncfg = Config()\nset_seed(cfg.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-13T16:54:15.813368Z","iopub.execute_input":"2026-01-13T16:54:15.814135Z","iopub.status.idle":"2026-01-13T16:54:15.824533Z","shell.execute_reply.started":"2026-01-13T16:54:15.814102Z","shell.execute_reply":"2026-01-13T16:54:15.823826Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class MelanomaDataset(Dataset):\n    def __init__(\n        self,\n        imgs: np.ndarray,\n        ys: np.ndarray,\n        df_csv: pd.DataFrame,\n        enc_state: Dict,\n        indices: np.ndarray,\n        transform=None,\n    ):\n        self.imgs = imgs\n        self.ys = ys\n        self.df = df_csv.reset_index(drop=True)\n        self.indices = np.asarray(indices, dtype=np.int64)\n        self.transform = transform\n\n        self.cat_cols = list(enc_state[\"cat_cols\"])\n        self.num_cols = list(enc_state[\"num_cols\"])\n\n        # categorical mappings + ensure UNK\n        self.cat_levels = {}\n        self.cat_maps = []\n        self.cat_cardinalities = []\n        for c in self.cat_cols:\n            levels = list(enc_state[\"cat_levels\"][c])\n            if \"__UNK__\" not in levels:\n                levels = levels + [\"__UNK__\"]\n            self.cat_levels[c] = levels\n            m = {v: i for i, v in enumerate(levels)}\n            self.cat_maps.append(m)\n            self.cat_cardinalities.append(len(levels))\n\n        # numeric stats\n        self.age_median = float(enc_state.get(\"age_median\", 0.0))\n        self.age_mean = float(enc_state.get(\"age_mean\", 0.0))\n        self.age_std = float(enc_state.get(\"age_std\", 1.0) + 1e-8)\n\n    def __len__(self) -> int:\n        return len(self.indices)\n\n    def _encode_cat(self, row) -> np.ndarray:\n        ids = np.zeros((len(self.cat_cols),), dtype=np.int64)\n        for j, c in enumerate(self.cat_cols):\n            v = row[c]\n            if pd.isna(v):\n                v = \"__UNK__\"\n            v = str(v)\n            m = self.cat_maps[j]\n            ids[j] = m.get(v, m[\"__UNK__\"])\n        return ids\n\n    def _encode_cont(self, row) -> np.ndarray:\n        vals = np.zeros((len(self.num_cols),), dtype=np.float32)\n        for j, c in enumerate(self.num_cols):\n            v = row[c]\n            if pd.isna(v):\n                v = self.age_median\n            v = float(v)\n            # z-score (train stats)\n            v = (v - self.age_mean) / self.age_std\n            vals[j] = v\n        return vals\n\n    def __getitem__(self, k: int):\n        i = int(self.indices[k])\n\n        img = self.imgs[i]  # uint8 HWC\n        y = float(self.ys[i])\n\n        row = self.df.iloc[i]\n        x_cat = self._encode_cat(row)   # (n_cat,)\n        x_cont = self._encode_cont(row) # (n_cont,)\n\n        if self.transform is not None:\n            img = self.transform(image=img)[\"image\"]  # torch tensor CHW float\n\n        return (\n            img,\n            torch.from_numpy(x_cat).long(),\n            torch.from_numpy(x_cont).float(),\n            torch.tensor(y, dtype=torch.float32),\n        )","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-01-13T16:53:35.303543Z","iopub.execute_input":"2026-01-13T16:53:35.303827Z","iopub.status.idle":"2026-01-13T16:53:35.314863Z","shell.execute_reply.started":"2026-01-13T16:53:35.303803Z","shell.execute_reply":"2026-01-13T16:53:35.314208Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cfg.fp_dir is None:\n    fps = sorted([p for p in cfg.cache_root.glob(\"fp_*\") if p.is_dir()])\n    if len(fps) == 0:\n        raise FileNotFoundError(f\"No fp_* folder found in {cfg.cache_root}\")\n    cache_dir = fps[0]\nelse:\n    cache_dir = cfg.cache_root / cfg.fp_dir\n\nimgs = np.load(cache_dir / \"images_uint8.npy\", mmap_mode=\"r\")\nys = np.load(cache_dir / \"targets_uint8.npy\", mmap_mode=\"r\")\n\nwith open(cache_dir / \"splits.json\", \"r\") as f:\n    splits = json.load(f)\ntrain_idx = np.array(splits[\"train_idx\"], dtype=np.int64)\nval_idx = np.array(splits[\"val_idx\"], dtype=np.int64)\n\nwith open(cache_dir / \"metadata_encoder.json\", \"r\") as f:\n    enc_state = json.load(f)\n\ndf = pd.read_csv('/kaggle/input/siim-isic-melanoma-classification/train.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-13T16:53:40.645035Z","iopub.execute_input":"2026-01-13T16:53:40.645294Z","iopub.status.idle":"2026-01-13T16:53:40.776377Z","shell.execute_reply.started":"2026-01-13T16:53:40.645272Z","shell.execute_reply":"2026-01-13T16:53:40.775806Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMAGENET_MEAN = (0.485, 0.456, 0.406)\nIMAGENET_STD = (0.229, 0.224, 0.225)\n\ntrain_tf = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.RandomBrightnessContrast(p=cfg.brightness_contrast_p),\n    A.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ToTensorV2(),\n])\n\nval_tf = A.Compose([\n    A.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ToTensorV2(),\n])\n\ntrain_ds = MelanomaDataset(imgs, ys, df, enc_state, train_idx, transform=train_tf)\nval_ds = MelanomaDataset(imgs, ys, df, enc_state, val_idx, transform=val_tf)\n\ntrain_loader = DataLoader(train_ds, batch_size=cfg.batch_size, shuffle=True, persistent_workers=True,\n                          num_workers=cfg.num_workers, pin_memory=True, drop_last=True)\nval_loader = DataLoader(val_ds, batch_size=cfg.batch_size, shuffle=False, persistent_workers=True,\n                        num_workers=cfg.num_workers, pin_memory=True)\n\nlen(train_loader), len(val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-13T16:53:45.514318Z","iopub.execute_input":"2026-01-13T16:53:45.514727Z","iopub.status.idle":"2026-01-13T16:53:45.530971Z","shell.execute_reply.started":"2026-01-13T16:53:45.514692Z","shell.execute_reply":"2026-01-13T16:53:45.530327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch_sample = next(iter(train_loader))\nbatch_sample[0].shape, batch_sample[1].shape, batch_sample[2].shape, batch_sample[3].shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-13T16:53:47.083370Z","iopub.execute_input":"2026-01-13T16:53:47.084010Z","iopub.status.idle":"2026-01-13T16:53:48.127708Z","shell.execute_reply.started":"2026-01-13T16:53:47.083981Z","shell.execute_reply":"2026-01-13T16:53:48.126573Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class MetaEncoder(nn.Module):\n    def __init__(\n        self,\n        cat_cardinalities: List[int],\n        n_cont: int,\n        d_model: int = 256,\n        n_heads: int = 8,\n        n_layers: int = 2,\n        ff_mult: int = 4,\n        dropout: float = 0.1,\n        use_cls_token: bool = True,\n        out_dim: Optional[int] = None,\n    ):\n        super().__init__()\n        if d_model % n_heads != 0:\n            raise ValueError(\"d_model must be divisible by n_heads\")\n\n        self.n_cat = len(cat_cardinalities)\n        self.n_cont = int(n_cont)\n        self.d_model = int(d_model)\n        self.use_cls = bool(use_cls_token)\n\n        self.cat_embeds = nn.ModuleList([\n            nn.Embedding(int(v), d_model) for v in cat_cardinalities\n        ])\n\n        self.cont_projs = nn.ModuleList([\n            nn.Sequential(\n                nn.Linear(1, d_model),\n                nn.LayerNorm(d_model),\n            )\n            for _ in range(self.n_cont)\n        ])\n\n        n_tokens = self.n_cat + self.n_cont + (1 if self.use_cls else 0)\n        self.pos_embed = nn.Parameter(torch.zeros(1, n_tokens, d_model))\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n        if self.use_cls:\n            self.cls = nn.Parameter(torch.zeros(1, 1, d_model))\n            nn.init.trunc_normal_(self.cls, std=0.02)\n\n        self.drop = nn.Dropout(dropout)\n        layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=n_heads,\n            dim_feedforward=d_model * ff_mult,\n            dropout=dropout,\n            activation=\"gelu\",\n            batch_first=True,\n            norm_first=True,\n        )\n        self.encoder = nn.TransformerEncoder(layer, num_layers=n_layers)\n        self.norm = nn.LayerNorm(d_model)\n\n        self.out_dim = int(out_dim) if out_dim is not None else d_model\n        self.proj = nn.Identity() if out_dim is None else nn.Linear(d_model, self.out_dim)\n\n    def forward(self, x_cat, x_cont):\n        toks = []\n\n        if self.n_cat > 0:\n            if x_cat is None:\n                raise ValueError(\"x_cat required because n_cat > 0\")\n            if x_cat.dim() != 2 or x_cat.size(1) != self.n_cat:\n                raise ValueError(f\"x_cat must be (B, {self.n_cat}) got {tuple(x_cat.shape)}\")\n            x_cat = x_cat.long()\n            for j, emb in enumerate(self.cat_embeds):\n                toks.append(emb(x_cat[:, j]))\n\n        if self.n_cont > 0:\n            if x_cont is None:\n                raise ValueError(\"x_cont required because n_cont > 0\")\n            if x_cont.dim() != 2 or x_cont.size(1) != self.n_cont:\n                raise ValueError(f\"x_cont must be (B, {self.n_cont}) got {tuple(x_cont.shape)}\")\n            x_cont = x_cont.float()\n            for j, proj in enumerate(self.cont_projs):\n                toks.append(proj(x_cont[:, j:j+1]))\n\n        if len(toks) == 0:\n            raise ValueError(\"No metadata fields provided\")\n\n        x = torch.stack(toks, dim=1)  # (B, T, d_model)\n\n        if self.use_cls:\n            cls = self.cls.expand(x.size(0), -1, -1)\n            x = torch.cat([cls, x], dim=1)\n\n        if x.size(1) != self.pos_embed.size(1):\n            raise ValueError(f\"Token count mismatch: {x.size(1)} vs pos_embed {self.pos_embed.size(1)}\")\n\n        x = self.drop(x + self.pos_embed)\n        x = self.encoder(x)\n        x = self.norm(x)\n\n        pooled = x[:, 0, :] if self.use_cls else x.mean(dim=1)\n        return self.proj(pooled)","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-01-13T16:53:56.510573Z","iopub.execute_input":"2026-01-13T16:53:56.510960Z","iopub.status.idle":"2026-01-13T16:53:56.526113Z","shell.execute_reply.started":"2026-01-13T16:53:56.510915Z","shell.execute_reply":"2026-01-13T16:53:56.525478Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MelaModel(nn.Module):\n    def __init__(\n        self,\n        image_encoder: nn.Module,\n        image_dim: int,\n        meta_encoder: nn.Module,\n        meta_out_dim: int,\n        fusion_hidden: int = 256,\n        fusion_dropout: float = 0.2,\n    ):\n        super().__init__()\n        self.image_encoder = image_encoder\n        self.image_dim = int(image_dim)\n        self.meta_encoder = meta_encoder\n        self.meta_out_dim = int(meta_out_dim)\n\n        self.fusion_head = nn.Sequential(\n            nn.Linear(self.image_dim + self.meta_out_dim, fusion_hidden),\n            nn.ReLU(inplace=True),\n            nn.Dropout(fusion_dropout),\n            nn.Linear(fusion_hidden, 1),\n        )\n\n    def forward(self, x_img: torch.Tensor, x_cat: torch.Tensor, x_cont: torch.Tensor) -> torch.Tensor:\n        f_img = self.image_encoder(x_img)\n        if f_img.dim() == 3:\n            f_img = f_img[:, 0, :]\n        elif f_img.dim() != 2:\n            raise ValueError(f\"Unexpected image encoder output: {tuple(f_img.shape)}\")\n        if f_img.size(-1) != self.image_dim:\n            raise ValueError(f\"image_dim mismatch: expected {self.image_dim}, got {f_img.size(-1)}\")\n\n        f_meta = self.meta_encoder(x_cat, x_cont)\n        fused = torch.cat([f_img, f_meta], dim=1)\n        return self.fusion_head(fused).squeeze(1)  # logits","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-01-13T16:53:56.726442Z","iopub.execute_input":"2026-01-13T16:53:56.726675Z","iopub.status.idle":"2026-01-13T16:53:56.733333Z","shell.execute_reply.started":"2026-01-13T16:53:56.726651Z","shell.execute_reply":"2026-01-13T16:53:56.732572Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def count_parameters(model):\n    total = sum(p.numel() for p in model.parameters())\n    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\n    print(f\"Total params      : {total:,}\")\n    print(f\"Trainable params  : {trainable:,}\")\n    print(f\"Non-trainable     : {total - trainable:,}\")\n    print(f\"Trainable percent: {100 * trainable / total:.2f}%\")\n\n    return total, trainable","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-13T16:53:58.354317Z","iopub.execute_input":"2026-01-13T16:53:58.354625Z","iopub.status.idle":"2026-01-13T16:53:58.359291Z","shell.execute_reply.started":"2026-01-13T16:53:58.354596Z","shell.execute_reply":"2026-01-13T16:53:58.358665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_encoder = timm.create_model(cfg.backbone, pretrained=cfg.pretrained, num_classes=0)\nimage_dim = image_encoder.num_features\n\nmeta_encoder = MetaEncoder(\n    cat_cardinalities=train_ds.cat_cardinalities,\n    n_cont=len(train_ds.num_cols),\n    d_model=cfg.meta_d_model,\n    n_heads=cfg.meta_heads,\n    n_layers=cfg.meta_layers,\n    dropout=cfg.meta_dropout,\n    use_cls_token=cfg.meta_use_cls,\n    out_dim=cfg.meta_out_dim,\n)\n\nmodel = MelaModel(\n    image_encoder=image_encoder,\n    image_dim=image_dim,\n    meta_encoder=meta_encoder,\n    meta_out_dim=cfg.meta_out_dim,\n    fusion_hidden=cfg.fusion_hidden,\n    fusion_dropout=cfg.fusion_dropout,\n).to(cfg.device)\n\n_ = count_parameters(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-13T17:12:39.046792Z","iopub.execute_input":"2026-01-13T17:12:39.047118Z","iopub.status.idle":"2026-01-13T17:12:39.404075Z","shell.execute_reply.started":"2026-01-13T17:12:39.047091Z","shell.execute_reply":"2026-01-13T17:12:39.403346Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"def tpr_at_fpr(y_true: np.ndarray, y_prob: np.ndarray, fpr_target: float = 0.10) -> float:\n    y_true = np.asarray(y_true).astype(int)\n    y_prob = np.asarray(y_prob).astype(float)\n    if len(np.unique(y_true)) < 2:\n        return float(\"nan\")\n    fpr, tpr, _ = roc_curve(y_true, y_prob)\n    idx = np.searchsorted(fpr, fpr_target, side=\"left\")\n    if idx >= len(tpr):\n        return float(tpr[-1])\n    return float(tpr[idx])\n\ndef expected_calibration_error(\n    y_true: np.ndarray,\n    y_prob: np.ndarray,\n    n_bins: int = 15\n) -> float:\n    y_true = np.asarray(y_true).astype(int)\n    y_prob = np.asarray(y_prob).astype(float)\n\n    bins = np.linspace(0.0, 1.0, n_bins + 1)\n    ece = 0.0\n    n = len(y_prob)\n    if n == 0:\n        return float(\"nan\")\n\n    for b0, b1 in zip(bins[:-1], bins[1:]):\n        if b1 < 1.0:\n            mask = (y_prob >= b0) & (y_prob < b1)\n        else:\n            mask = (y_prob >= b0) & (y_prob <= b1)\n\n        m = mask.sum()\n        if m == 0:\n            continue\n\n        conf = y_prob[mask].mean()\n        acc = y_true[mask].mean()\n        ece += (m / n) * abs(acc - conf)\n\n    return float(ece)\n\ndef compute_all_metrics(\n    y_true: np.ndarray,\n    y_prob: np.ndarray,\n    y_pred: np.ndarray,\n) -> Dict[str, float]:\n    out = {}\n    \n    out[\"acc\"] = float(accuracy_score(y_true, y_pred))\n    out[\"f1\"] = float(f1_score(y_true, y_pred, zero_division=0))\n    out[\"prec\"] = float(precision_score(y_true, y_pred, zero_division=0))\n    out[\"recall\"] = float(recall_score(y_true, y_pred, zero_division=0))\n\n    # ranking metrics\n    if len(np.unique(y_true)) == 2:\n        out[\"roc_auc\"] = float(roc_auc_score(y_true, y_prob))\n        out[\"pr_auc\"] = float(average_precision_score(y_true, y_prob))\n        out[\"tpr_at_10fpr\"] = tpr_at_fpr(y_true, y_prob, fpr_target=0.10)\n    else:\n        out[\"roc_auc\"] = float(\"nan\")\n        out[\"pr_auc\"] = float(\"nan\")\n        out[\"tpr_at_10fpr\"] = float(\"nan\")\n\n    # calibration\n    out[\"ece\"] = expected_calibration_error(y_true, y_prob, n_bins=15)\n    return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-13T17:12:39.405343Z","iopub.execute_input":"2026-01-13T17:12:39.405628Z","iopub.status.idle":"2026-01-13T17:12:39.415142Z","shell.execute_reply.started":"2026-01-13T17:12:39.405604Z","shell.execute_reply":"2026-01-13T17:12:39.414422Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef evaluate(model, loader, device: str, threshold: float = 0.5) -> Dict[str, float]:\n    model.eval()\n    all_y, all_p = [], []\n    total_loss = 0.0\n    n = 0\n\n    for x_img, x_cat, x_cont, y in loader:\n        x_img = x_img.to(device, non_blocking=True)\n        x_cat = x_cat.to(device, non_blocking=True)\n        x_cont = x_cont.to(device, non_blocking=True)\n        y = y.to(device, non_blocking=True)\n\n        logits = model(x_img, x_cat, x_cont)\n        probs = torch.sigmoid(logits)\n\n        all_y.append(y.detach().cpu().numpy())\n        all_p.append(probs.detach().cpu().numpy())\n\n    y_true = np.concatenate(all_y)\n    y_prob = np.concatenate(all_p)\n    y_pred = (y_prob >= threshold).astype(int)\n    \n    metrics = compute_all_metrics(y_true, y_prob, y_pred)\n    metrics[\"mean_prob\"] = float(np.mean(y_prob))\n    \n    return metrics\n\ndef train_one_epoch(model, loader, optimizer, scaler, criterion, device: str, amp: bool) -> float:\n    model.train()\n    total_loss = 0.0\n    n = 0\n\n    pbar = tqdm(loader, desc=\"Training...\", leave=False)\n    for x_img, x_cat, x_cont, y in pbar:\n        x_img = x_img.to(device, non_blocking=True)\n        x_cat = x_cat.to(device, non_blocking=True)\n        x_cont = x_cont.to(device, non_blocking=True)\n        y = y.to(device, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with autocast(\"cuda\", enabled=(amp and device.startswith(\"cuda\"))):\n            logits = model(x_img, x_cat, x_cont)\n            loss = criterion(logits, y)\n\n        if amp and device.startswith(\"cuda\"):\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            loss.backward()\n            optimizer.step()\n\n        bs = y.size(0)\n        total_loss += float(loss.detach().cpu()) * bs\n        n += bs\n\n    return total_loss / max(n, 1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-13T17:12:39.741268Z","iopub.execute_input":"2026-01-13T17:12:39.741784Z","iopub.status.idle":"2026-01-13T17:12:39.752001Z","shell.execute_reply.started":"2026-01-13T17:12:39.741733Z","shell.execute_reply":"2026-01-13T17:12:39.751254Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_train = ys[train_idx].astype(np.float32)\npos = float(y_train.sum())\nneg = float(len(y_train) - pos)\npos_weight = torch.tensor([neg / (pos + 1e-8)], device=cfg.device, dtype=torch.float32)\n\ncriterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\noptimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.wd)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=cfg.epochs)\nscaler = GradScaler(\"cuda\", enabled=(cfg.amp and cfg.device.startswith(\"cuda\")))\n\nbest_roc = -1.0\nbest_path = \"best_model.pt\"\n\nfor epoch in range(1, cfg.epochs + 1):\n    train_loss = train_one_epoch(model, train_loader, optimizer, scaler, criterion, cfg.device, cfg.amp)\n    val_metrics = evaluate(model, val_loader, cfg.device)\n\n    scheduler.step()\n\n    lr_now = optimizer.param_groups[0][\"lr\"]\n    print(\n        f\"Epoch {epoch:02d}/{cfg.epochs} | lr={lr_now:.2e} | train_loss={train_loss:.4f} | \"\n        f\"acc={val_metrics['acc']:.4f} f1={val_metrics['f1']:.4f} \"\n        f\"prec={val_metrics['prec']:.4f} rec={val_metrics['recall']:.4f} \"\n        f\"roc={val_metrics['roc_auc']:.4f} pr={val_metrics['pr_auc']:.4f} \"\n        f\"tpr@10fpr={val_metrics['tpr_at_10fpr']:.4f} ece={val_metrics['ece']:.4f}\"\n    )\n\n    if val_metrics['roc_auc'] == val_metrics['roc_auc'] and val_metrics['roc_auc'] > best_roc:\n        best_roc = val_metrics['roc_auc']\n        torch.save(\n            {\n                \"model\": model.state_dict(),\n                \"cfg\": cfg.__dict__,\n                \"cat_cardinalities\": train_ds.cat_cardinalities,\n                \"num_cols\": train_ds.num_cols,\n                \"best_roc\": best_roc,\n            },\n            best_path\n        )\n        print(f\"  Saved best to: {best_path} (best_roc={best_roc:.4f})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-13T17:12:40.155114Z","iopub.execute_input":"2026-01-13T17:12:40.155578Z","iopub.status.idle":"2026-01-13T17:22:58.486266Z","shell.execute_reply.started":"2026-01-13T17:12:40.155551Z","shell.execute_reply":"2026-01-13T17:22:58.485570Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---","metadata":{}}]}