{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":14535717,"sourceType":"datasetVersion","datasetId":9009659}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport timm\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import Dataset, ConcatDataset\nfrom fastai.vision.all import *\nimport albumentations as A\nfrom sklearn.metrics import accuracy_score, f1_score, roc_auc_score\n\n# --- CONFIGURATION ---\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nSEED = 101\nPATCH_H, PATCH_W = 512, 512  \nAXIAL_SIZE = 384             \npatch_size = 64              \nLmax = 15                    \nANGLE = 30\nLR_MAX = 5e-6\nBS = 10\nEPOCHS = 3\n\nLEVELS = ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\nLABELS_MAP = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n\ndef seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(SEED)\n\n# --- AUGMENTATIONS & LOADING ---\ntrain_aug = A.Compose([\n    A.RandomBrightnessContrast(brightness_limit=(-0.2, 0.2), contrast_limit=(-0.2, 0.2), p=0.5),\n    A.OneOf([\n        A.MotionBlur(blur_limit=5),\n        A.MedianBlur(blur_limit=5),\n        A.GaussianBlur(blur_limit=5),\n        A.GaussNoise(var_limit=(5.0, 30.0)),\n    ], p=0.5),\n    A.OneOf([\n        A.OpticalDistortion(distort_limit=1.0),\n        A.GridDistortion(num_steps=5, distort_limit=1.),\n        A.ElasticTransform(alpha=3),\n    ], p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, border_mode=0, p=0.5),\n    A.CoarseDropout(max_holes=8, max_height=32, max_width=32, min_holes=1, min_height=8, min_width=8, p=0.5),\n])\n\ndef load_axial_slice(path, x, y):\n    if pd.isna(path) or not isinstance(path, str) or not os.path.exists(path):\n        return torch.zeros(1, AXIAL_SIZE, AXIAL_SIZE)\n    try:\n        dcm = pydicom.dcmread(path)\n        img = dcm.pixel_array.astype(np.float32)\n        \n        # FIX: Handle NaNs/Infs in raw DICOM data\n        img = np.nan_to_num(img)\n        \n        img = img - img.min()\n        v_max = np.percentile(img, 99)\n        img = (img / v_max if v_max > 0 else img / (img.max() + 1e-6))\n        img = np.clip(img, 0, 1)\n        img = (img * 255).astype(np.uint8)\n\n        h, w = img.shape\n        cx, cy = (int(x), int(y)) if (pd.notna(x) and pd.notna(y)) else (w//2, h//2)\n        pad_h, pad_w = int(0.15 * h), int(0.15 * w)\n        \n        crop = img[max(0, cy-pad_h):min(h, cy+pad_h), max(0, cx-pad_w):min(w, cx+pad_w)]\n        if crop.size == 0: return torch.zeros(1, AXIAL_SIZE, AXIAL_SIZE)\n        \n        img_resized = cv2.resize(crop, (AXIAL_SIZE, AXIAL_SIZE), interpolation=cv2.INTER_LINEAR)\n        return torch.from_numpy(img_resized).unsqueeze(0).float() / 255.0\n    except Exception:\n        return torch.zeros(1, AXIAL_SIZE, AXIAL_SIZE)\n\nclass DualView_Spinal_Dataset(Dataset):\n    def __init__(self, sot_df, label_df=None, VALID=False, augment=False, P=patch_size):\n        self.sot_df = sot_df.copy()\n        self.sot_df['study_id'] = self.sot_df['study_id'].astype(str)\n        self.VALID = VALID\n        self.augment = augment \n        self.P = P\n        self.sag_resize = torchvision.transforms.Resize((PATCH_H, PATCH_W), antialias=True)\n        self.label_lookup = {}\n        if label_df is not None:\n            self.label_lookup = label_df.copy().set_index('study_id').to_dict('index')\n        self.study_groups = self.sot_df.groupby('study_id')\n        self.study_ids = list(self.study_groups.groups.keys())\n\n    def __len__(self): return len(self.study_ids)\n\n    def __getitem__(self, index):\n        study_id = self.study_ids[index]\n        group = self.study_groups.get_group(study_id).set_index('level')\n        sag_imgs = torch.zeros(5, Lmax, 2*self.P, 2*self.P)\n        sag_masks = torch.ones(5, Lmax).bool() \n        ax_imgs = torch.zeros(5, 2, 1, AXIAL_SIZE, AXIAL_SIZE)\n        targets = []\n        lbl_row = self.label_lookup.get(study_id, {})\n        for idx, level in enumerate(LEVELS):\n            val = lbl_row.get(f\"spinal_canal_stenosis_{level.replace('/', '_').lower()}\", 'UNK')\n            targets.append(LABELS_MAP.get(val, -100))\n            if level not in group.index: continue\n            row = group.loc[level]\n            ax_imgs[idx, 0] = load_axial_slice(row['ax_0_path'], row['ax_0_x'], row['ax_0_y'])\n            ax_imgs[idx, 1] = load_axial_slice(row['ax_1_path'], row['ax_1_x'], row['ax_1_y'])\n            if pd.notna(row['sag_path']) and os.path.exists(row['sag_path']):\n                try:\n                    self._load_sagittal_volume(row, sag_imgs, sag_masks, idx)\n                except Exception: pass\n        if self.augment:\n            sag_imgs, ax_imgs = self._apply_augmentations(sag_imgs, ax_imgs)\n        if not self.VALID and not self.augment:\n            B, D, H, W = sag_imgs.shape\n            angle = random.uniform(-ANGLE, ANGLE)\n            sag_imgs = torchvision.transforms.functional.rotate(sag_imgs.view(-1, H, W), angle).view(B, D, H, W)\n        c = self.P // 2\n        sag_imgs = sag_imgs[:, :, c:c+self.P, c:c+self.P]\n        return [sag_imgs, sag_masks, ax_imgs], torch.tensor(targets, dtype=torch.long)\n\n    def _apply_augmentations(self, sag_imgs, ax_imgs):\n        for i in range(5):\n            vol = sag_imgs[i] \n            vol_np = (vol.permute(1, 2, 0).numpy() * 255).astype(np.uint8)\n            aug_vol = train_aug(image=vol_np)['image']\n            sag_imgs[i] = torch.from_numpy(aug_vol).permute(2, 0, 1).float() / 255.0\n        for i in range(5):\n            for j in range(2):\n                img = ax_imgs[i, j, 0] \n                img_np = (img.numpy() * 255).astype(np.uint8)                \n                aug_img = train_aug(image=img_np)['image']                \n                ax_imgs[i, j, 0] = torch.from_numpy(aug_img).float() / 255.0\n        return sag_imgs, ax_imgs\n\n    def _load_sagittal_volume(self, row, sag_imgs, sag_masks, idx):\n        dcm_ref = pydicom.dcmread(row['sag_path'])\n        img_ref = dcm_ref.pixel_array.astype(np.float32)\n        H_orig, W_orig = img_ref.shape\n        cx, cy = (row['sag_x'], row['sag_y']) if pd.notna(row['sag_x']) else (W_orig/2, H_orig/2)\n        transpose = False\n        if H_orig > W_orig:\n            cy -= (H_orig - W_orig) // 2\n            H_orig, transpose = W_orig, True\n        elif H_orig < W_orig:\n            cx -= (W_orig - H_orig) // 2\n            W_orig = H_orig\n        if not self.VALID:\n            sc_y = int(cy * PATCH_H / H_orig + self.P) + int(random.gauss(0, 5))\n            sc_x = int(cx * PATCH_W / W_orig + self.P) + int(random.gauss(0, 5))\n        else:\n            sc_y = int(cy * PATCH_H / H_orig + self.P)\n            sc_x = int(cx * PATCH_W / W_orig + self.P)\n        folder = os.path.dirname(row['sag_path'])\n        try:\n            center_inst = int(os.path.basename(row['sag_path']).split('.')[0])\n        except:\n            return \n        if not self.VALID:\n            center_inst = center_inst + random.randint(-2, 2)\n        for k in range(Lmax):\n            inst = center_inst - Lmax // 2 + k\n            fpath = os.path.join(folder, f\"{inst}.dcm\")\n            if os.path.exists(fpath):\n                try:\n                    px = pydicom.dcmread(fpath).pixel_array.astype(np.float32)\n                    v_max = np.quantile(px, 0.99)\n                    px = (px - px.min()) / (v_max + 1e-6)\n                    px = torch.tensor(np.clip(px, 0, 1))\n                    if transpose: \n                        diff = px.shape[0] - px.shape[1]\n                        start = diff // 2\n                        px = px[start : start + px.shape[1], :]\n                    elif px.shape[0] < px.shape[1]: \n                        diff = px.shape[1] - px.shape[0]\n                        start = diff // 2\n                        px = px[:, start : start + px.shape[0]]\n                    px = self.sag_resize(px.unsqueeze(0))\n                    px = nn.functional.pad(px, [self.P]*4, 'reflect').squeeze(0)\n                    if sc_y-self.P >= 0 and sc_x-self.P >= 0:\n                        if sc_y+self.P <= px.shape[0] and sc_x+self.P <= px.shape[1]:\n                            sag_imgs[idx, k] = px[sc_y-self.P:sc_y+self.P, sc_x-self.P:sc_x+self.P]\n                            sag_masks[idx, k] = False\n                except:\n                    continue\n\n# --- ABLATION: Simple CNN Definition ---\nclass SimpleCNN(nn.Module):\n    def __init__(self, in_chans=1, out_dim=512):\n        super().__init__()\n        self.features = nn.Sequential(\n            nn.Conv2d(in_chans, 32, 3, padding=1, bias=False), \n            nn.BatchNorm2d(32), \n            nn.ReLU(inplace=True), \n            nn.MaxPool2d(2),\n            \n            nn.Conv2d(32, 64, 3, padding=1, bias=False), \n            nn.BatchNorm2d(64), \n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2),\n            \n            nn.Conv2d(64, 128, 3, padding=1, bias=False), \n            nn.BatchNorm2d(128), \n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2),\n            \n            nn.Conv2d(128, 256, 3, padding=1, bias=False), \n            nn.BatchNorm2d(256), \n            nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool2d((1, 1))\n        )\n        self.fc = nn.Linear(256, out_dim)\n        # FIX: LayerNorm is crucial before feeding into Transformer to match scale\n        self.ln = nn.LayerNorm(out_dim) \n        \n        self._init_weights()\n\n    def _init_weights(self):\n        # FIX: Proper initialization prevents NaN start\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n            elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm, nn.LayerNorm)):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.Linear):\n                nn.init.normal_(m.weight, 0, 0.01)\n                nn.init.constant_(m.bias, 0)\n\n    def forward(self, x):\n        x = self.features(x)\n        x = x.flatten(1)\n        x = self.fc(x)\n        return self.ln(x) # Output is now normalized\n\nclass SpatialDropout(nn.Module):\n    def __init__(self, drop=0.1):\n        super(SpatialDropout, self).__init__()\n        self.drop = drop\n        \n    def forward(self, inputs, noise_shape=None):\n        outputs = inputs.clone()\n        if noise_shape is None:\n            noise_shape = (inputs.shape[0], inputs.shape[1], *([1] * (inputs.dim() - 2)))\n        \n        if not self.training or self.drop == 0:\n            return inputs\n        else:\n            noises = self._make_noises(inputs, noise_shape)\n            if self.drop == 1:\n                noises.fill_(0.0)\n            else:\n                noises.bernoulli_(1 - self.drop).div_(1 - self.drop)\n            noises = noises.expand_as(inputs)    \n            outputs.mul_(noises)\n            return outputs\n            \n    def _make_noises(self, inputs, noise_shape):\n        return torch.zeros(noise_shape, dtype=inputs.dtype, device=inputs.device)\n\n# --- MODEL: DualView_ViT with Simple Encoders ---\nclass DualView_ViT(nn.Module):\n    def __init__(self, axial_weight_path=None, dim=512, depth=6, head_size=64):\n        super().__init__()\n        \n        # Uses the stabilized SimpleCNN\n        self.sag_encoder = SimpleCNN(in_chans=1, out_dim=dim)\n        self.ax_encoder = SimpleCNN(in_chans=1, out_dim=dim)\n\n        self.register_buffer('slices_enc', self._get_sinusoid(Lmax, dim))\n        self.register_buffer('pos_enc', self._get_sinusoid(5, dim))\n        \n        encoder_layer = nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim, dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True)\n        self.sag_slice_transformer = nn.TransformerEncoder(encoder_layer, depth)\n        self.level_transformer = nn.TransformerEncoder(encoder_layer, depth)\n        \n        self.cross_attn = nn.MultiheadAttention(dim, dim//head_size, batch_first=True)\n        self.norm_sag = nn.LayerNorm(dim)\n        self.norm_ax = nn.LayerNorm(dim)\n        \n        self.s_dropout = SpatialDropout(0.1)\n        self.proj_out = nn.Linear(dim, 3)\n        \n        # Initialize the projection head safely\n        nn.init.xavier_uniform_(self.proj_out.weight)\n        nn.init.zeros_(self.proj_out.bias)\n\n    def _get_sinusoid(self, n_pos, d_model):\n        position = torch.arange(0, n_pos, dtype=torch.float).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n        pe = torch.zeros(n_pos, d_model)\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        return pe\n\n    def forward(self, x):\n        sag_imgs, sag_masks, ax_imgs = x\n        B = sag_imgs.shape[0]\n\n        # Input Check for debugging (Optional, removes NaNs if they appear mid-batch)\n        if torch.isnan(sag_imgs).any() or torch.isnan(ax_imgs).any():\n             sag_imgs = torch.nan_to_num(sag_imgs)\n             ax_imgs = torch.nan_to_num(ax_imgs)\n\n        s = self.sag_encoder(sag_imgs.view(-1, 1, patch_size, patch_size))\n        s = s.view(B*5, Lmax, -1) + self.slices_enc\n        \n        flat_mask = sag_masks.view(B*5, Lmax)\n        s = self.sag_slice_transformer(s, src_key_padding_mask=flat_mask)\n        s = s.masked_fill(flat_mask.unsqueeze(-1), 0)\n        s_emb = s.sum(1) / (~flat_mask).sum(1).unsqueeze(-1).clamp(min=1)\n\n        a = self.ax_encoder(ax_imgs.view(-1, 1, AXIAL_SIZE, AXIAL_SIZE))\n        a = a.view(B*5, 2, -1)\n\n        q = self.norm_sag(s_emb.unsqueeze(1))\n        kv = self.norm_ax(a)\n        attn_out, _ = self.cross_attn(q, kv, kv)\n        fused = s_emb + attn_out.squeeze(1)\n\n        lvl_mask = (sag_masks.sum(2) == Lmax)\n        out = self.level_transformer(fused.view(B, 5, -1) + self.pos_enc, src_key_padding_mask=lvl_mask)\n        \n        out = out.permute(0, 2, 1) \n        out = self.s_dropout(out)\n        out = out.permute(0, 2, 1) \n        \n        return self.proj_out(out)\n\n# --- METRICS & LOSS ---\ndef myLoss(preds, target):\n    weights = torch.tensor([1., 2., 4.], device=preds.device)\n    return nn.CrossEntropyLoss(weight=weights, ignore_index=-100)(preds.view(-1, 3), target.view(-1))\n\ndef calculate_level_metrics(preds, targs):\n    \"\"\"\n    Calculates metrics per level.\n    preds: (N, 5, 3)\n    targs: (N, 5)\n    Returns a pandas DataFrame with results.\n    \"\"\"\n    results = []\n    \n    for i, level_name in enumerate(LEVELS):\n        p_sub = preds[:, i, :] # (N, 3)\n        t_sub = targs[:, i]    # (N,)\n\n        # Filter valid\n        mask = t_sub != -100\n        if not mask.any():\n            continue\n\n        p_valid = p_sub[mask]\n        t_valid = t_sub[mask]\n        \n        # Softmax for AUC\n        probs_valid = torch.softmax(p_valid, dim=1).cpu().numpy()\n        preds_cls = p_valid.argmax(dim=1).cpu().numpy()\n        t_np = t_valid.cpu().numpy()\n\n        # 1. Accuracy\n        acc = accuracy_score(t_np, preds_cls)\n\n        # 2. Macro F1\n        f1 = f1_score(t_np, preds_cls, average='macro')\n\n        # 3. Multiclass AUROC (One-vs-Rest)\n        try:\n            auroc = roc_auc_score(t_np, probs_valid, multi_class='ovr', average='macro')\n        except:\n            auroc = np.nan\n\n        # 4. Binary AUROC (Normal/Mild [0] vs Mod/Severe [1+2])\n        # Binarize targets: 0 -> 0, {1,2} -> 1\n        t_bin = (t_np > 0).astype(int) \n        # Probability of \"Disease\" (Class 1 + Class 2)\n        prob_bin = probs_valid[:, 1] + probs_valid[:, 2]\n        \n        # Only calculate if both classes are present\n        if len(np.unique(t_bin)) > 1:\n            try:\n                bin_auroc = roc_auc_score(t_bin, prob_bin)\n            except:\n                bin_auroc = np.nan\n        else:\n            bin_auroc = np.nan\n\n        results.append({\n            'Level': level_name,\n            'Acc': acc,\n            'Macro_F1': f1,\n            'Multi_AUC': auroc,\n            'Binary_AUC': bin_auroc\n        })\n        \n    return pd.DataFrame(results).round(4)\n\ndef calculate_global_metrics(preds, targs):\n    mask = targs != -100\n    if not mask.any(): return 0.0, 0.0, 0.0\n    p = preds[mask]\n    t = targs[mask]\n    pred_cls = p.argmax(dim=1)\n    acc = accuracy_score(t.cpu().numpy(), pred_cls.cpu().numpy())\n    f1 = f1_score(t.cpu().numpy(), pred_cls.cpu().numpy(), average='macro')\n    try:\n        probs = torch.softmax(p, dim=1).cpu().numpy()\n        roc = roc_auc_score(t.cpu().numpy(), probs, multi_class='ovr', average='macro')\n    except:\n        roc = 0.0 \n    return acc, f1, roc\n\ndef nt(nmin, nmax, tcur, tmax):\n    return (nmax - .5*(nmax-nmin)*(1+np.cos(tcur*np.pi/tmax))).astype(np.float32)\n\nclass AlphaScheduler(Callback):\n    def before_batch(self):\n        total_iter = self.learn.n_epoch * len(self.learn.dls.train)\n        current_iter = self.learn.train_iter\n        if total_iter == 0: total_iter = 1\n        alpha = torch.as_tensor(nt(0.25, 1, current_iter, total_iter))\n        self.learn.alpha_val = alpha \n\n# --- PIPELINE ---\ndef run_pipeline(sot_df, train_labels_df, folds=[5], sag_folder=None, ax_path=None):\n    sot_df = sot_df.copy()\n    train_labels_df = train_labels_df.copy()\n    \n    sot_df['study_id'] = sot_df['study_id'].astype(str)\n    train_labels_df['study_id'] = train_labels_df['study_id'].astype(str)\n    \n    if set(sot_df.study_id).isdisjoint(set(train_labels_df.study_id)):\n        print(\"CRITICAL WARNING: No matching study_ids found between SOT and Labels.\")\n\n    for f in folds:\n        print(f\"\\n================ STARTING FOLD {f} (ABLATION STUDY: Simple CNN) ================\")\n        seed_everything(SEED)\n        \n        model = DualView_ViT(axial_weight_path=None).to(device)\n        print(\"--> ABLATION INFO: Using SimpleCNN encoders. Pretrained weights are NOT loaded.\")\n        \n        t_ids = train_labels_df[train_labels_df.fold != f].study_id.unique()\n        v_ids = train_labels_df[train_labels_df.fold == f].study_id.unique()\n        \n        tds_clean = DualView_Spinal_Dataset(sot_df[sot_df.study_id.isin(t_ids)], train_labels_df, VALID=False, augment=False)\n        tds_aug = DualView_Spinal_Dataset(sot_df[sot_df.study_id.isin(t_ids)], train_labels_df, VALID=False, augment=True)\n        tds_combined = ConcatDataset([tds_clean, tds_aug])\n        vds = DualView_Spinal_Dataset(sot_df[sot_df.study_id.isin(v_ids)], train_labels_df, VALID=True, augment=False)\n        \n        tdl = torch.utils.data.DataLoader(tds_combined, batch_size=BS, shuffle=True, drop_last=True, num_workers=4, pin_memory=True)\n        vdl = torch.utils.data.DataLoader(vds, batch_size=BS, shuffle=False, num_workers=4, pin_memory=True)\n        dls = DataLoaders(tdl, vdl, device=device)\n\n        callbacks = [GradientClip(3.0), AlphaScheduler()]\n        \n        learn = Learner(dls, model, loss_func=myLoss, cbs=callbacks)\n        \n        print(\"--> Training...\")\n        try:\n            learn.fit_one_cycle(EPOCHS, lr_max=LR_MAX, wd=0.05, pct_start=0.02)\n        except Exception as e:\n            print(f\"ERROR during training Fold {f}: {e}\")\n            continue \n        \n        save_name = f'DualView_ViT_SimpleCNN_Fold{f}.pth'\n        torch.save(model.state_dict(), save_name)\n        print(f\"--> Model saved: {save_name}\")\n        \n        print(f\"--> Calculating Metrics for Fold {f}...\")\n        try:\n            # preds: (N, 5, 3), targs: (N, 5)\n            preds, targs = learn.get_preds(dl=vdl) \n            \n            # Global Metrics\n            preds_flat = preds.view(-1, 3)\n            targs_flat = targs.view(-1)\n            g_acc, g_f1, g_roc = calculate_global_metrics(preds_flat, targs_flat)\n            \n            weights = torch.tensor([1., 2., 4.], device='cpu') \n            final_loss = nn.CrossEntropyLoss(weight=weights, ignore_index=-100)(preds_flat, targs_flat).item()\n            \n            print(f\"\\n[FOLD {f} GLOBAL REPORT]\")\n            print(f\"  Loss : {final_loss:.4f}\")\n            print(f\"  Acc  : {g_acc:.4f}\")\n            print(f\"  F1   : {g_f1:.4f}\")\n            print(f\"  AUC  : {g_roc:.4f}\\n\")\n            \n            print(f\"[FOLD {f} LEVEL-WISE REPORT]\")\n            level_df = calculate_level_metrics(preds, targs)\n            print(level_df.to_string(index=False))\n            print(\"\\n\")\n            \n        except Exception as e:\n            print(f\"Error calculating metrics: {e}\")\n            import traceback\n            traceback.print_exc()\n\n        del model, learn, dls, tdl, vdl, tds_clean, tds_aug, tds_combined, vds\n        gc.collect()\n        torch.cuda.empty_cache()","metadata":{"_uuid":"f72b9b9a-f9d2-4e5d-b8c7-07b286aaac1a","_cell_guid":"d0ac0c2b-0a20-44d4-992e-9914ae6fb345","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-08T13:29:35.420915Z","iopub.execute_input":"2026-02-08T13:29:35.421311Z","iopub.status.idle":"2026-02-08T13:29:40.328679Z","shell.execute_reply.started":"2026-02-08T13:29:35.421274Z","shell.execute_reply":"2026-02-08T13:29:40.327451Z"}},"outputs":[{"name":"stderr","text":"/opt/conda/lib/python3.10/site-packages/albumentations/__init__.py:13: UserWarning: A new version of Albumentations is available: 2.0.8 (you have 1.4.17). Upgrade using: pip install -U albumentations. To disable automatic update checks, set the environment variable NO_ALBUMENTATIONS_UPDATE to 1.\n  check_for_updates()\n/opt/conda/lib/python3.10/site-packages/pydantic/main.py:212: UserWarning: blur_limit and sigma_limit minimum value can not be both equal to 0. blur_limit minimum value changed to 3.\n  validated_self = self.__pydantic_validator__.validate_python(data, self_instance=self)\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"# --- PREP & RUN ---\ndf_coors = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv\")\nlabels = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\")\ntrain_desc = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\")\n\ndf = df_coors.merge(train_desc, on=['study_id', 'series_id'], how='left')\ndf = df[df['series_description'] != 'Sagittal T1']\n\ndata = dict()   \nTRAIN_PATH = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images'\n\nprint(\"Parsing dataset structure...\")\nfor i in range(df.shape[0]):\n    study_id = str(df.iloc[i]['study_id'])\n    series_id = str(df.iloc[i]['series_id'])\n    instance_number = str(df.iloc[i]['instance_number'])\n    level = str(df.iloc[i]['level'])\n    x, y = df.iloc[i]['x'], df.iloc[i]['y']\n    \n    desc_val = str(df.iloc[i]['series_description'])\n    desc = desc_val.split()[0].lower() if len(desc_val.split()) > 0 else \"unknown\"\n\n    if study_id not in data:\n        data[study_id] = dict()\n\n    if level not in data[study_id]:\n        data[study_id][level] = {'sagittal': {}, 'axial': []}\n        \n    full_path = os.path.join(TRAIN_PATH, str(study_id), str(series_id), str(instance_number) + '.dcm')\n    \n    if desc == 'axial':\n        data[study_id][level]['axial'].append({'x': x, 'y': y, 'path': full_path})\n    else:\n        data[study_id][level]['sagittal'] = {'x': x, 'y': y, 'path': full_path}\n\nrows = []\nfor study_id, levels_dict in data.items():\n    for level, views in levels_dict.items():\n        row = {'study_id': study_id, 'level': level}\n        if 'sagittal' in views and views['sagittal']:\n            row['sag_x'] = views['sagittal']['x']\n            row['sag_y'] = views['sagittal']['y']\n            row['sag_path'] = views['sagittal']['path']\n        if 'axial' in views:\n            for idx, ax_slice in enumerate(views['axial']):\n                if idx > 1: break \n                row[f'ax_{idx}_x'] = ax_slice['x']\n                row[f'ax_{idx}_y'] = ax_slice['y']\n                row[f'ax_{idx}_path'] = ax_slice['path']\n        rows.append(row)\n\nsot_df = pd.DataFrame(rows)\nspinal = [\n    'spinal_canal_stenosis_l1_l2', 'spinal_canal_stenosis_l2_l3',\n    'spinal_canal_stenosis_l3_l4', 'spinal_canal_stenosis_l4_l5',\n    'spinal_canal_stenosis_l5_s1'\n]\n\nn_folds = 5\nlabels[\"fold\"] = (np.arange(len(labels)) % n_folds) + 1\nlabels = labels[['study_id','fold']+spinal][labels[spinal].isna().sum(1) < len(spinal)].reset_index(drop=True)\n\nprint(\"Data preparation complete.\")\n\nrun_pipeline(\n    sot_df, \n    labels, \n    sag_folder=None, \n    ax_path=None     \n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T13:29:42.237647Z","iopub.execute_input":"2026-02-08T13:29:42.238138Z","iopub.status.idle":"2026-02-08T13:33:51.358081Z","shell.execute_reply.started":"2026-02-08T13:29:42.238107Z","shell.execute_reply":"2026-02-08T13:33:51.356155Z"}},"outputs":[{"name":"stdout","text":"Parsing dataset structure...\nData preparation complete.\n\n================ STARTING FOLD 5 (ABLATION STUDY: Simple CNN) ================\n","output_type":"stream"},{"name":"stderr","text":"/opt/conda/lib/python3.10/site-packages/torch/nn/modules/transformer.py:307: UserWarning: enable_nested_tensor is True, but self.use_nested_tensor is False because encoder_layer.norm_first was True\n  warnings.warn(f\"enable_nested_tensor is True, but self.use_nested_tensor is False because {why_not_sparsity_fast_path}\")\n","output_type":"stream"},{"name":"stdout","text":"--> ABLATION INFO: Using SimpleCNN encoders. Pretrained weights are NOT loaded.\n--> Training...\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"<IPython.core.display.HTML object>","text/html":"\n<style>\n    /* Turns off some styling */\n    progress {\n        /* gets rid of default border in Firefox and Opera. */\n        border: none;\n        /* Needs to be in here for Safari polyfill so background images work as expected. */\n        background-size: auto;\n    }\n    progress:not([value]), progress:not([value])::-webkit-progress-bar {\n        background: repeating-linear-gradient(45deg, #7e7e7e, #7e7e7e 10px, #5c5c5c 10px, #5c5c5c 20px);\n    }\n    .progress-bar-interrupted, .progress-bar-interrupted::-webkit-progress-bar {\n        background: #F44336;\n    }\n</style>\n"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"<IPython.core.display.HTML object>","text/html":"\n    <div>\n      <progress value='0' class='' max='3' style='width:300px; height:20px; vertical-align: middle;'></progress>\n      0.00% [0/3 00:00&lt;?]\n    </div>\n    \n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: left;\">\n      <th>epoch</th>\n      <th>train_loss</th>\n      <th>valid_loss</th>\n      <th>time</th>\n    </tr>\n  </thead>\n  <tbody>\n  </tbody>\n</table><p>\n\n    <div>\n      <progress value='31' class='' max='315' style='width:300px; height:20px; vertical-align: middle;'></progress>\n      9.84% [31/315 03:50&lt;35:15 nan]\n    </div>\n    "},"metadata":{}},{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mKeyboardInterrupt\u001b[0m                         Traceback (most recent call last)","Cell \u001b[0;32mIn[2], line 65\u001b[0m\n\u001b[1;32m     61\u001b[0m labels \u001b[38;5;241m=\u001b[39m labels[[\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mstudy_id\u001b[39m\u001b[38;5;124m'\u001b[39m,\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mfold\u001b[39m\u001b[38;5;124m'\u001b[39m]\u001b[38;5;241m+\u001b[39mspinal][labels[spinal]\u001b[38;5;241m.\u001b[39misna()\u001b[38;5;241m.\u001b[39msum(\u001b[38;5;241m1\u001b[39m) \u001b[38;5;241m<\u001b[39m \u001b[38;5;28mlen\u001b[39m(spinal)]\u001b[38;5;241m.\u001b[39mreset_index(drop\u001b[38;5;241m=\u001b[39m\u001b[38;5;28;01mTrue\u001b[39;00m)\n\u001b[1;32m     63\u001b[0m \u001b[38;5;28mprint\u001b[39m(\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mData preparation complete.\u001b[39m\u001b[38;5;124m\"\u001b[39m)\n\u001b[0;32m---> 65\u001b[0m \u001b[43mrun_pipeline\u001b[49m\u001b[43m(\u001b[49m\n\u001b[1;32m     66\u001b[0m \u001b[43m    \u001b[49m\u001b[43msot_df\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\n\u001b[1;32m     67\u001b[0m \u001b[43m    \u001b[49m\u001b[43mlabels\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\n\u001b[1;32m     68\u001b[0m \u001b[43m    \u001b[49m\u001b[43msag_folder\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[38;5;28;43;01mNone\u001b[39;49;00m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\n\u001b[1;32m     69\u001b[0m \u001b[43m    \u001b[49m\u001b[43max_path\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[38;5;28;43;01mNone\u001b[39;49;00m\u001b[43m     \u001b[49m\n\u001b[1;32m     70\u001b[0m \u001b[43m)\u001b[49m\n","Cell \u001b[0;32mIn[1], line 474\u001b[0m, in \u001b[0;36mrun_pipeline\u001b[0;34m(sot_df, train_labels_df, folds, sag_folder, ax_path)\u001b[0m\n\u001b[1;32m    472\u001b[0m \u001b[38;5;28mprint\u001b[39m(\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124m--> Training...\u001b[39m\u001b[38;5;124m\"\u001b[39m)\n\u001b[1;32m    473\u001b[0m \u001b[38;5;28;01mtry\u001b[39;00m:\n\u001b[0;32m--> 474\u001b[0m     \u001b[43mlearn\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mfit_one_cycle\u001b[49m\u001b[43m(\u001b[49m\u001b[43mEPOCHS\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mlr_max\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mLR_MAX\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mwd\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[38;5;241;43m0.05\u001b[39;49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mpct_start\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[38;5;241;43m0.02\u001b[39;49m\u001b[43m)\u001b[49m\n\u001b[1;32m    475\u001b[0m \u001b[38;5;28;01mexcept\u001b[39;00m \u001b[38;5;167;01mException\u001b[39;00m \u001b[38;5;28;01mas\u001b[39;00m e:\n\u001b[1;32m    476\u001b[0m     \u001b[38;5;28mprint\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mERROR during training Fold \u001b[39m\u001b[38;5;132;01m{\u001b[39;00mf\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00me\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m\"\u001b[39m)\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/fastai/callback/schedule.py:121\u001b[0m, in \u001b[0;36mfit_one_cycle\u001b[0;34m(self, n_epoch, lr_max, div, div_final, pct_start, wd, moms, cbs, reset_opt, start_epoch)\u001b[0m\n\u001b[1;32m    118\u001b[0m lr_max \u001b[38;5;241m=\u001b[39m np\u001b[38;5;241m.\u001b[39marray([h[\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mlr\u001b[39m\u001b[38;5;124m'\u001b[39m] \u001b[38;5;28;01mfor\u001b[39;00m h \u001b[38;5;129;01min\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mopt\u001b[38;5;241m.\u001b[39mhypers])\n\u001b[1;32m    119\u001b[0m scheds \u001b[38;5;241m=\u001b[39m {\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mlr\u001b[39m\u001b[38;5;124m'\u001b[39m: combined_cos(pct_start, lr_max\u001b[38;5;241m/\u001b[39mdiv, lr_max, lr_max\u001b[38;5;241m/\u001b[39mdiv_final),\n\u001b[1;32m    120\u001b[0m           \u001b[38;5;124m'\u001b[39m\u001b[38;5;124mmom\u001b[39m\u001b[38;5;124m'\u001b[39m: combined_cos(pct_start, \u001b[38;5;241m*\u001b[39m(\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mmoms \u001b[38;5;28;01mif\u001b[39;00m moms \u001b[38;5;129;01mis\u001b[39;00m \u001b[38;5;28;01mNone\u001b[39;00m \u001b[38;5;28;01melse\u001b[39;00m moms))}\n\u001b[0;32m--> 121\u001b[0m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mfit\u001b[49m\u001b[43m(\u001b[49m\u001b[43mn_epoch\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mcbs\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mParamScheduler\u001b[49m\u001b[43m(\u001b[49m\u001b[43mscheds\u001b[49m\u001b[43m)\u001b[49m\u001b[38;5;241;43m+\u001b[39;49m\u001b[43mL\u001b[49m\u001b[43m(\u001b[49m\u001b[43mcbs\u001b[49m\u001b[43m)\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mreset_opt\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mreset_opt\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mwd\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mwd\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mstart_epoch\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mstart_epoch\u001b[49m\u001b[43m)\u001b[49m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/fastai/learner.py:266\u001b[0m, in \u001b[0;36mLearner.fit\u001b[0;34m(self, n_epoch, lr, wd, cbs, reset_opt, start_epoch)\u001b[0m\n\u001b[1;32m    264\u001b[0m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mopt\u001b[38;5;241m.\u001b[39mset_hypers(lr\u001b[38;5;241m=\u001b[39m\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mlr \u001b[38;5;28;01mif\u001b[39;00m lr \u001b[38;5;129;01mis\u001b[39;00m \u001b[38;5;28;01mNone\u001b[39;00m \u001b[38;5;28;01melse\u001b[39;00m lr)\n\u001b[1;32m    265\u001b[0m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mn_epoch \u001b[38;5;241m=\u001b[39m n_epoch\n\u001b[0;32m--> 266\u001b[0m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_with_events\u001b[49m\u001b[43m(\u001b[49m\u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_do_fit\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[38;5;124;43mfit\u001b[39;49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mCancelFitException\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_end_cleanup\u001b[49m\u001b[43m)\u001b[49m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/fastai/learner.py:201\u001b[0m, in \u001b[0;36mLearner._with_events\u001b[0;34m(self, f, event_type, ex, final)\u001b[0m\n\u001b[1;32m    200\u001b[0m \u001b[38;5;28;01mdef\u001b[39;00m \u001b[38;5;21m_with_events\u001b[39m(\u001b[38;5;28mself\u001b[39m, f, event_type, ex, final\u001b[38;5;241m=\u001b[39mnoop):\n\u001b[0;32m--> 201\u001b[0m     \u001b[38;5;28;01mtry\u001b[39;00m: \u001b[38;5;28mself\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mbefore_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mevent_type\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m);  \u001b[43mf\u001b[49m\u001b[43m(\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m    202\u001b[0m     \u001b[38;5;28;01mexcept\u001b[39;00m ex: \u001b[38;5;28mself\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mafter_cancel_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mevent_type\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m)\n\u001b[1;32m    203\u001b[0m     \u001b[38;5;28mself\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mafter_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mevent_type\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m);  final()\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/fastai/learner.py:255\u001b[0m, in \u001b[0;36mLearner._do_fit\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    253\u001b[0m \u001b[38;5;28;01mfor\u001b[39;00m epoch \u001b[38;5;129;01min\u001b[39;00m \u001b[38;5;28mrange\u001b[39m(\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mn_epoch):\n\u001b[1;32m    254\u001b[0m     \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mepoch\u001b[38;5;241m=\u001b[39mepoch\n\u001b[0;32m--> 255\u001b[0m     \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_with_events\u001b[49m\u001b[43m(\u001b[49m\u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_do_epoch\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[38;5;124;43mepoch\u001b[39;49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mCancelEpochException\u001b[49m\u001b[43m)\u001b[49m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/fastai/learner.py:201\u001b[0m, in \u001b[0;36mLearner._with_events\u001b[0;34m(self, f, event_type, ex, final)\u001b[0m\n\u001b[1;32m    200\u001b[0m \u001b[38;5;28;01mdef\u001b[39;00m \u001b[38;5;21m_with_events\u001b[39m(\u001b[38;5;28mself\u001b[39m, f, event_type, ex, final\u001b[38;5;241m=\u001b[39mnoop):\n\u001b[0;32m--> 201\u001b[0m     \u001b[38;5;28;01mtry\u001b[39;00m: \u001b[38;5;28mself\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mbefore_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mevent_type\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m);  \u001b[43mf\u001b[49m\u001b[43m(\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m    202\u001b[0m     \u001b[38;5;28;01mexcept\u001b[39;00m ex: \u001b[38;5;28mself\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mafter_cancel_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mevent_type\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m)\n\u001b[1;32m    203\u001b[0m     \u001b[38;5;28mself\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mafter_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mevent_type\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m);  final()\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/fastai/learner.py:249\u001b[0m, in \u001b[0;36mLearner._do_epoch\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    248\u001b[0m \u001b[38;5;28;01mdef\u001b[39;00m \u001b[38;5;21m_do_epoch\u001b[39m(\u001b[38;5;28mself\u001b[39m):\n\u001b[0;32m--> 249\u001b[0m     \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_do_epoch_train\u001b[49m\u001b[43m(\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m    250\u001b[0m     \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_do_epoch_validate()\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/fastai/learner.py:241\u001b[0m, in \u001b[0;36mLearner._do_epoch_train\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    239\u001b[0m \u001b[38;5;28;01mdef\u001b[39;00m \u001b[38;5;21m_do_epoch_train\u001b[39m(\u001b[38;5;28mself\u001b[39m):\n\u001b[1;32m    240\u001b[0m     \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdl \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdls\u001b[38;5;241m.\u001b[39mtrain\n\u001b[0;32m--> 241\u001b[0m     \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_with_events\u001b[49m\u001b[43m(\u001b[49m\u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mall_batches\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[38;5;124;43mtrain\u001b[39;49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mCancelTrainException\u001b[49m\u001b[43m)\u001b[49m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/fastai/learner.py:201\u001b[0m, in \u001b[0;36mLearner._with_events\u001b[0;34m(self, f, event_type, ex, final)\u001b[0m\n\u001b[1;32m    200\u001b[0m \u001b[38;5;28;01mdef\u001b[39;00m \u001b[38;5;21m_with_events\u001b[39m(\u001b[38;5;28mself\u001b[39m, f, event_type, ex, final\u001b[38;5;241m=\u001b[39mnoop):\n\u001b[0;32m--> 201\u001b[0m     \u001b[38;5;28;01mtry\u001b[39;00m: \u001b[38;5;28mself\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mbefore_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mevent_type\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m);  \u001b[43mf\u001b[49m\u001b[43m(\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m    202\u001b[0m     \u001b[38;5;28;01mexcept\u001b[39;00m ex: \u001b[38;5;28mself\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mafter_cancel_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mevent_type\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m)\n\u001b[1;32m    203\u001b[0m     \u001b[38;5;28mself\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mafter_\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mevent_type\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m);  final()\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/fastai/learner.py:207\u001b[0m, in \u001b[0;36mLearner.all_batches\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    205\u001b[0m \u001b[38;5;28;01mdef\u001b[39;00m \u001b[38;5;21mall_batches\u001b[39m(\u001b[38;5;28mself\u001b[39m):\n\u001b[1;32m    206\u001b[0m     \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mn_iter \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mlen\u001b[39m(\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdl)\n\u001b[0;32m--> 207\u001b[0m     \u001b[38;5;28;01mfor\u001b[39;00m o \u001b[38;5;129;01min\u001b[39;00m \u001b[38;5;28menumerate\u001b[39m(\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdl): \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mone_batch(\u001b[38;5;241m*\u001b[39mo)\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/torch/utils/data/dataloader.py:630\u001b[0m, in \u001b[0;36m_BaseDataLoaderIter.__next__\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    627\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_sampler_iter \u001b[38;5;129;01mis\u001b[39;00m \u001b[38;5;28;01mNone\u001b[39;00m:\n\u001b[1;32m    628\u001b[0m     \u001b[38;5;66;03m# TODO(https://github.com/pytorch/pytorch/issues/76750)\u001b[39;00m\n\u001b[1;32m    629\u001b[0m     \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_reset()  \u001b[38;5;66;03m# type: ignore[call-arg]\u001b[39;00m\n\u001b[0;32m--> 630\u001b[0m data \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_next_data\u001b[49m\u001b[43m(\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m    631\u001b[0m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_num_yielded \u001b[38;5;241m+\u001b[39m\u001b[38;5;241m=\u001b[39m \u001b[38;5;241m1\u001b[39m\n\u001b[1;32m    632\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_dataset_kind \u001b[38;5;241m==\u001b[39m _DatasetKind\u001b[38;5;241m.\u001b[39mIterable \u001b[38;5;129;01mand\u001b[39;00m \\\n\u001b[1;32m    633\u001b[0m         \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_IterableDataset_len_called \u001b[38;5;129;01mis\u001b[39;00m \u001b[38;5;129;01mnot\u001b[39;00m \u001b[38;5;28;01mNone\u001b[39;00m \u001b[38;5;129;01mand\u001b[39;00m \\\n\u001b[1;32m    634\u001b[0m         \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_num_yielded \u001b[38;5;241m>\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_IterableDataset_len_called:\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/torch/utils/data/dataloader.py:1327\u001b[0m, in \u001b[0;36m_MultiProcessingDataLoaderIter._next_data\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m   1324\u001b[0m     \u001b[38;5;28;01mreturn\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_process_data(data)\n\u001b[1;32m   1326\u001b[0m \u001b[38;5;28;01massert\u001b[39;00m \u001b[38;5;129;01mnot\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_shutdown \u001b[38;5;129;01mand\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_tasks_outstanding \u001b[38;5;241m>\u001b[39m \u001b[38;5;241m0\u001b[39m\n\u001b[0;32m-> 1327\u001b[0m idx, data \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_get_data\u001b[49m\u001b[43m(\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m   1328\u001b[0m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_tasks_outstanding \u001b[38;5;241m-\u001b[39m\u001b[38;5;241m=\u001b[39m \u001b[38;5;241m1\u001b[39m\n\u001b[1;32m   1329\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_dataset_kind \u001b[38;5;241m==\u001b[39m _DatasetKind\u001b[38;5;241m.\u001b[39mIterable:\n\u001b[1;32m   1330\u001b[0m     \u001b[38;5;66;03m# Check for _IterableDatasetStopIteration\u001b[39;00m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/torch/utils/data/dataloader.py:1283\u001b[0m, in \u001b[0;36m_MultiProcessingDataLoaderIter._get_data\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m   1281\u001b[0m \u001b[38;5;28;01melif\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_pin_memory:\n\u001b[1;32m   1282\u001b[0m     \u001b[38;5;28;01mwhile\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_pin_memory_thread\u001b[38;5;241m.\u001b[39mis_alive():\n\u001b[0;32m-> 1283\u001b[0m         success, data \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_try_get_data\u001b[49m\u001b[43m(\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m   1284\u001b[0m         \u001b[38;5;28;01mif\u001b[39;00m success:\n\u001b[1;32m   1285\u001b[0m             \u001b[38;5;28;01mreturn\u001b[39;00m data\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/torch/utils/data/dataloader.py:1131\u001b[0m, in \u001b[0;36m_MultiProcessingDataLoaderIter._try_get_data\u001b[0;34m(self, timeout)\u001b[0m\n\u001b[1;32m   1118\u001b[0m \u001b[38;5;28;01mdef\u001b[39;00m \u001b[38;5;21m_try_get_data\u001b[39m(\u001b[38;5;28mself\u001b[39m, timeout\u001b[38;5;241m=\u001b[39m_utils\u001b[38;5;241m.\u001b[39mMP_STATUS_CHECK_INTERVAL):\n\u001b[1;32m   1119\u001b[0m     \u001b[38;5;66;03m# Tries to fetch data from `self._data_queue` once for a given timeout.\u001b[39;00m\n\u001b[1;32m   1120\u001b[0m     \u001b[38;5;66;03m# This can also be used as inner loop of fetching without timeout, with\u001b[39;00m\n\u001b[0;32m   (...)\u001b[0m\n\u001b[1;32m   1128\u001b[0m     \u001b[38;5;66;03m# Returns a 2-tuple:\u001b[39;00m\n\u001b[1;32m   1129\u001b[0m     \u001b[38;5;66;03m#   (bool: whether successfully get data, any: data if successful else None)\u001b[39;00m\n\u001b[1;32m   1130\u001b[0m     \u001b[38;5;28;01mtry\u001b[39;00m:\n\u001b[0;32m-> 1131\u001b[0m         data \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_data_queue\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mget\u001b[49m\u001b[43m(\u001b[49m\u001b[43mtimeout\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mtimeout\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m   1132\u001b[0m         \u001b[38;5;28;01mreturn\u001b[39;00m (\u001b[38;5;28;01mTrue\u001b[39;00m, data)\n\u001b[1;32m   1133\u001b[0m     \u001b[38;5;28;01mexcept\u001b[39;00m \u001b[38;5;167;01mException\u001b[39;00m \u001b[38;5;28;01mas\u001b[39;00m e:\n\u001b[1;32m   1134\u001b[0m         \u001b[38;5;66;03m# At timeout and error, we manually check whether any worker has\u001b[39;00m\n\u001b[1;32m   1135\u001b[0m         \u001b[38;5;66;03m# failed. Note that this is the only mechanism for Windows to detect\u001b[39;00m\n\u001b[1;32m   1136\u001b[0m         \u001b[38;5;66;03m# worker failures.\u001b[39;00m\n","File \u001b[0;32m/opt/conda/lib/python3.10/queue.py:180\u001b[0m, in \u001b[0;36mQueue.get\u001b[0;34m(self, block, timeout)\u001b[0m\n\u001b[1;32m    178\u001b[0m         \u001b[38;5;28;01mif\u001b[39;00m remaining \u001b[38;5;241m<\u001b[39m\u001b[38;5;241m=\u001b[39m \u001b[38;5;241m0.0\u001b[39m:\n\u001b[1;32m    179\u001b[0m             \u001b[38;5;28;01mraise\u001b[39;00m Empty\n\u001b[0;32m--> 180\u001b[0m         \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mnot_empty\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mwait\u001b[49m\u001b[43m(\u001b[49m\u001b[43mremaining\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m    181\u001b[0m item \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_get()\n\u001b[1;32m    182\u001b[0m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mnot_full\u001b[38;5;241m.\u001b[39mnotify()\n","File \u001b[0;32m/opt/conda/lib/python3.10/threading.py:324\u001b[0m, in \u001b[0;36mCondition.wait\u001b[0;34m(self, timeout)\u001b[0m\n\u001b[1;32m    322\u001b[0m \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[1;32m    323\u001b[0m     \u001b[38;5;28;01mif\u001b[39;00m timeout \u001b[38;5;241m>\u001b[39m \u001b[38;5;241m0\u001b[39m:\n\u001b[0;32m--> 324\u001b[0m         gotit \u001b[38;5;241m=\u001b[39m \u001b[43mwaiter\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43macquire\u001b[49m\u001b[43m(\u001b[49m\u001b[38;5;28;43;01mTrue\u001b[39;49;00m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mtimeout\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m    325\u001b[0m     \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[1;32m    326\u001b[0m         gotit \u001b[38;5;241m=\u001b[39m waiter\u001b[38;5;241m.\u001b[39macquire(\u001b[38;5;28;01mFalse\u001b[39;00m)\n","\u001b[0;31mKeyboardInterrupt\u001b[0m: "],"ename":"KeyboardInterrupt","evalue":"","output_type":"error"}],"execution_count":2},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}