{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"c89442ae","cell_type":"code","source":"# ==============================================================================\n# RSNA Knee Abnormality Detection - HighAUCMultiViewNet-V2 Full Dataset GPU Engine (v49)\n# Multi-Piece Sequential Learning:\n#   - Trains through remaining dataset pieces sequentially (Piece 2 -> Piece 3 -> Piece 4 -> ...)\n#   - Automatic warm-restart weight loading from previous piece checkpoints\n#   - Class-Adaptive Margin Asymmetric Loss (CAM-ASL) + 36-Token Latent Cross-View Transformer\n#   - Continuous peak validation AUC tracking -> 'best_hqcnn_model.pth'\n# ==============================================================================\nimport os\nimport sys\nimport gc\nimport io\nimport math\nimport re\nimport random\nimport subprocess\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Strict PyTorch initializations for Python 3.12 compatibility\nimport torch\nimport torch._utils\nimport torch._dynamo\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim as optim\nfrom torch.cuda.amp import autocast, GradScaler\n\ndef ensure_dependencies():\n    packages = [(\"pydicom\", \"pydicom\"), (\"timm\", \"timm\")]\n    for pkg_install, pkg_import in packages:\n        try:\n            __import__(pkg_import)\n        except ImportError:\n            print(f\"[INSTALL] Installing {pkg_install}...\", flush=True)\n            subprocess.check_call([sys.executable, \"-m\", \"pip\", \"install\", \"--quiet\", pkg_install])\n\nensure_dependencies()\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\n\ntry:\n    import cv2\nexcept (Exception, BaseException):\n    cv2 = None\n\nfrom PIL import Image\nfrom sklearn.metrics import roc_auc_score\nimport timm\n\nclass CFG:\n    DATA_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n    for path in [\"/kaggle/input/competitions/rsna-knee-abnormality-detection\", \"/kaggle/input/rsna-knee-abnormality-detection\"]:\n        if os.path.exists(path):\n            DATA_DIR = path\n            break\n        \n    TRAIN_CSV = os.path.join(DATA_DIR, \"train.csv\")\n    SERIES_CSV = os.path.join(DATA_DIR, \"train_series.csv\")\n    \n    STUDIES_PER_PIECE = 1102      # Exactly 1/4th of the dataset (1,102 patients)\n    START_PIECE = 4              # FINAL PIECE!\n    PIECES_PER_SESSION = 1       # Train exactly ONE piece (1/4th of data) per Kaggle Version!\n    EPOCHS_PER_PIECE = 4         # 4 epochs per piece\n    \n    IMAGE_SIZE = (192, 192)      \n    NUM_SLICES = 12              # 12 slices per anatomical view (36 slices total)\n    BATCH_SIZE = 4               # Optimal batch size for GPU T4 x2\n    LR = 2.0e-4 if START_PIECE == 1 else 2.0e-5  # Drop LR by 10x when resuming to prevent catastrophic forgetting\n    WEIGHT_DECAY = 1e-2\n    BACKBONE = \"resnet34d\"       # SOTA Medical Deep-Stem Architecture\n    NUM_CLASSES = 12\n    DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    TARGET_COLS = [\n        \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n        \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n        \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\"\n    ]\n    \n    CLASS_PREVALENCES = [\n        0.18, 0.12, 0.32, 0.22,\n        0.35, 0.20, 0.25, 0.45,\n        0.10, 0.08, 0.15, 0.05\n    ]\n\nIMAGENET_MEAN = torch.tensor([0.485, 0.456, 0.406], dtype=torch.float32).view(1, 3, 1, 1)\nIMAGENET_STD  = torch.tensor([0.229, 0.224, 0.225], dtype=torch.float32).view(1, 3, 1, 1)\n\ndef build_medical_prior_adjacency():\n    A_prior = np.eye(12, dtype=np.float32) * 1.5\n    def connect(i, j, weight=1.0):\n        A_prior[i, j] = weight\n        A_prior[j, i] = weight\n    connect(0, 2, 0.90); connect(0, 3, 0.75); connect(0, 10, 0.95); connect(0, 7, 0.85)\n    connect(1, 2, 0.80); connect(1, 10, 0.70)\n    connect(2, 4, 0.70); connect(3, 5, 0.70); connect(2, 7, 0.75); connect(3, 7, 0.70)\n    connect(4, 5, 0.65); connect(4, 6, 0.75); connect(4, 7, 0.80); connect(4, 8, 0.70)\n    connect(7, 8, 0.90); connect(7, 9, 0.85)\n    connect(11, 10, 0.95); connect(11, 7, 0.90)\n    d = np.sum(A_prior, axis=1)\n    d_inv_sqrt = np.power(d, -0.5)\n    d_inv_sqrt[np.isinf(d_inv_sqrt)] = 0.0\n    D_mat = np.diag(d_inv_sqrt)\n    A_norm = D_mat.dot(A_prior).dot(D_mat)\n    return torch.tensor(A_norm, dtype=torch.float32)\n\nFINDING_PATTERNS = {\n    \"ACL\": [r'\\b(acl|anterior cruciate|vkb|kreuzband|lca|ligamento cruzado anterior)\\b.*\\b(tear|ruptur|riss|rotura|desgarro|injury|defect)\\b'],\n    \"MCL\": [r'\\b(mcl|medial collateral|innenband|lcm|ligamento colateral medial)\\b.*\\b(tear|ruptur|riss|rotura|desgarro|sprain)\\b'],\n    \"Medial Meniscus\": [r'\\b(medial meniscus|medial meniscal|innenmeniskus|menisco medial|menisco interno)\\b.*\\b(tear|ruptur|riss|rotura|desgarro|maceration|defect)\\b'],\n    \"Lateral Meniscus\": [r'\\b(lateral meniscus|lateral meniscal|aussenmeniskus|menisco lateral|menisco externo)\\b.*\\b(tear|ruptur|riss|rotura|desgarro|maceration|defect)\\b'],\n    \"Medial OA\": [r'\\b(medial|mediale|compartimento medial)\\b.*\\b(osteoarthritis|gonarthrose|artrosis|cartilage loss|knorpelschaden|desgaste)\\b'],\n    \"Lateral OA\": [r'\\b(lateral|laterale|compartimento lateral)\\b.*\\b(osteoarthritis|gonarthrose|artrosis|cartilage loss|knorpelschaden|desgaste)\\b'],\n    \"PF OA\": [r'\\b(patellofemoral|patella|retropatellar|femororrotuliana)\\b.*\\b(osteoarthritis|arthrose|artrosis|chondromalacia|condromalacia)\\b'],\n    \"Effusion\": [r'\\b(effusion|erguss|gelenkerguss|derrame|derrame articular|fluid)\\b'],\n    \"Synovitis\": [r'\\b(synovitis|synovialitis|sinovitis|synovial thickening)\\b'],\n    \"Baker's\": [r'\\b(baker|popliteal cyst|bakerzyste|quiste de baker|quiste popliteo)\\b'],\n    \"Contusion\": [r'\\b(bone contusion|bone marrow edema|bone bruise|knochenmarkoedem|contusion osea|edema oseo)\\b'],\n    \"Fracture\": [r'\\b(fracture|fraktur|fractura|avulsion|knochenbruch)\\b']\n}\n\nNEGATION_REGEX = re.compile(r'\\b(no|not|without|free of|negative for|intact|unremarkable|normal|kein|keine|ohne|intakt|unauffaellig|sin|ausencia de|conservado)\\b', re.IGNORECASE)\n\ndef parse_report_fast(report_text):\n    if not isinstance(report_text, str) or len(report_text.strip()) == 0:\n        return [0.10] * 12\n    text = report_text.lower().replace('\\n', ' ')\n    results = []\n    for col in CFG.TARGET_COLS:\n        patterns = FINDING_PATTERNS[col]\n        val = 0.10\n        for pat in patterns:\n            m = re.search(pat, text)\n            if m:\n                start, end = m.span()\n                window = text[max(0, start - 35):min(len(text), end + 35)]\n                if NEGATION_REGEX.search(window):\n                    val = 0.05\n                else:\n                    val = 0.95\n                break\n        results.append(val)\n    return results\n\ndef extract_instance_num(filename):\n    match = re.search(r'(\\d+)\\.dcm$', filename)\n    return int(match.group(1)) if match else 0\n\ndef fast_read_dcm(path, img_size=(192, 192)):\n    try:\n        with open(path, 'rb') as f:\n            dcm = pydicom.dcmread(io.BytesIO(f.read()), stop_before_pixels=False)\n        img = dcm.pixel_array.astype(np.float32)\n        slope = float(getattr(dcm, 'RescaleSlope', 1.0))\n        intercept = float(getattr(dcm, 'RescaleIntercept', 0.0))\n        img = img * slope + intercept\n        min_v = img.min()\n        max_v = img.max()\n        if max_v > min_v:\n            img = (img - min_v) / (max_v - min_v)\n        else:\n            img = np.zeros_like(img)\n        img = (img * 255.0).astype(np.uint8)\n        if cv2 is not None:\n            return cv2.resize(img, img_size, interpolation=cv2.INTER_LINEAR)\n        else:\n            return np.array(Image.fromarray(img).resize(img_size, Image.BILINEAR))\n    except Exception:\n        return np.zeros((img_size[1], img_size[0]), dtype=np.uint8)\n\ndef sample_center_window(slice_paths, num_slices=12, img_size=(192, 192), is_train=False):\n    total = len(slice_paths)\n    if total == 0:\n        return torch.zeros((num_slices, 3, img_size[0], img_size[1]), dtype=torch.float32), 0.0\n\n    sorted_paths = sorted(slice_paths, key=extract_instance_num)\n    start_idx = int(total * 0.15)\n    end_idx = int(total * 0.85)\n    if end_idx <= start_idx:\n        start_idx, end_idx = 0, total - 1\n        \n    center_indices = np.linspace(start_idx, end_idx, num_slices, dtype=int)\n    do_hflip = is_train and (random.random() > 0.5)\n    \n    loaded_imgs = {}\n    def get_img(idx):\n        idx = max(0, min(total - 1, idx))\n        if idx not in loaded_imgs:\n            im = fast_read_dcm(sorted_paths[idx], img_size)\n            if do_hflip:\n                im = np.fliplr(im).copy()\n            loaded_imgs[idx] = im\n        return loaded_imgs[idx]\n\n    triplets_list = []\n    for idx in center_indices:\n        img_prev = get_img(max(0, idx - 1))\n        img_curr = get_img(idx)\n        img_next = get_img(min(total - 1, idx + 1))\n        img_stack = np.stack([img_prev, img_curr, img_next], axis=-1)\n        triplets_list.append(torch.from_numpy(img_stack).permute(2, 0, 1).float())\n\n    stacked_volume = torch.stack(triplets_list) / 255.0\n    normalized_volume = (stacked_volume - IMAGENET_MEAN) / IMAGENET_STD\n    del loaded_imgs\n    return normalized_volume, 1.0\n\nclass HighAUCDatasetV2(Dataset):\n    def __init__(self, df, path_index, is_train=False):\n        self.df = df.copy().reset_index(drop=True)\n        self.path_index = path_index\n        self.is_train = is_train\n\n    def __len__(self): \n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        study_id = str(row[\"StudyInstanceUID\"])\n        study_info = self.path_index.get(study_id, {})\n\n        sag_files = study_info.get(\"Sag\", [])\n        cor_files = study_info.get(\"Cor\", [])\n        axi_files = study_info.get(\"Ax\", [])\n        meta_vals = study_info.get(\"meta\", [0.0, 0.0, 0.0])\n\n        sag_vol, sag_valid = sample_center_window(sag_files, CFG.NUM_SLICES, CFG.IMAGE_SIZE, is_train=self.is_train)\n        cor_vol, cor_valid = sample_center_window(cor_files, CFG.NUM_SLICES, CFG.IMAGE_SIZE, is_train=self.is_train)\n        axi_vol, axi_valid = sample_center_window(axi_files, CFG.NUM_SLICES, CFG.IMAGE_SIZE, is_train=self.is_train)\n        \n        valid_mask = torch.tensor([sag_valid, cor_valid, axi_valid], dtype=torch.float32)\n        meta = torch.tensor(meta_vals, dtype=torch.float32)\n\n        labels = [float(row[col]) for col in CFG.TARGET_COLS]\n        loss_weight = float(row.get('loss_weight', 1.0))\n        return sag_vol, cor_vol, axi_vol, valid_mask, meta, torch.tensor(labels, dtype=torch.float32), loss_weight\n\nclass GatedAttentionMIL(nn.Module):\n    def __init__(self, in_features, hidden_dim=64):\n        super().__init__()\n        self.v = nn.Sequential(nn.Linear(in_features, hidden_dim), nn.Tanh())\n        self.u = nn.Sequential(nn.Linear(in_features, hidden_dim), nn.Sigmoid())\n        self.w = nn.Linear(hidden_dim, 1)\n\n    def forward(self, x, is_valid=None):\n        v, u = self.v(x), self.u(x)\n        weights = self.w(v * u)\n        alpha = torch.softmax(weights, dim=1)\n        pooled = (x * alpha).sum(dim=1)\n        if is_valid is not None:\n            pooled = pooled * is_valid.view(-1, 1)\n        return torch.nan_to_num(pooled, nan=0.0)\n\nclass AnatomicalPriorGCN(nn.Module):\n    def __init__(self, in_features=64, num_classes=12, alpha=0.70):\n        super().__init__()\n        self.alpha = alpha\n        self.register_buffer(\"A_prior\", build_medical_prior_adjacency())\n        self.A_learned = nn.Parameter(torch.randn(num_classes, num_classes) * 0.01)\n        self.node_proj = nn.Linear(in_features, 32)\n        self.gcn_conv = nn.Linear(32, 32, bias=False)\n        self.out_head = nn.Linear(32, 1)\n\n    def forward(self, decoupled_node_feats):\n        A_learned_norm = torch.softmax(self.A_learned, dim=-1)\n        A_effective = self.alpha * self.A_prior + (1.0 - self.alpha) * A_learned_norm\n        h = F.silu(self.node_proj(decoupled_node_feats))\n        h_gcn = torch.matmul(A_effective, h)\n        h_gcn = F.silu(self.gcn_conv(h_gcn))\n        return self.out_head(h_gcn).squeeze(-1)\n\nclass HighAUCMultiViewNet_V3(nn.Module):\n    def __init__(self, backbone=CFG.BACKBONE, num_slices=CFG.NUM_SLICES, meta_dim=3, num_classes=CFG.NUM_CLASSES):\n        super().__init__()\n        self.num_slices = num_slices\n        self.encoder = timm.create_model(backbone, pretrained=True, num_classes=0)\n        in_feats = self.encoder.num_features\n        \n        self.view_plane_emb = nn.Parameter(torch.randn(3, in_feats) * 0.02)\n        self.slice_pos_emb = nn.Parameter(torch.randn(num_slices, in_feats) * 0.02)\n        \n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=in_feats, nhead=8, dim_feedforward=in_feats * 2, dropout=0.1, activation=\"gelu\", batch_first=True\n        )\n        self.cross_view_transformer = nn.TransformerEncoder(encoder_layer, num_layers=2)\n        \n        self.sag_mil = GatedAttentionMIL(in_feats, hidden_dim=64)\n        self.cor_mil = GatedAttentionMIL(in_feats, hidden_dim=64)\n        self.axi_mil = GatedAttentionMIL(in_feats, hidden_dim=64)\n        \n        self.meta_proj = nn.Sequential(nn.Linear(meta_dim, 16), nn.LayerNorm(16), nn.SiLU())\n        self.fusion = nn.Sequential(\n            nn.Linear(in_feats * 3 + 16, 256),\n            nn.LayerNorm(256),\n            nn.SiLU(),\n            nn.Dropout(0.20),\n            nn.Linear(256, 128),\n            nn.LayerNorm(128),\n            nn.SiLU()\n        )\n        self.decoupled_heads = nn.ModuleList([\n            nn.Sequential(nn.Linear(128, 64), nn.SiLU(), nn.Linear(64, 64))\n            for _ in range(num_classes)\n        ])\n        self.prior_gcn = AnatomicalPriorGCN(in_features=64, num_classes=num_classes, alpha=0.70)\n\n    def forward(self, sag, cor, axi, valid_mask, meta):\n        B, S, C, H, W = sag.shape\n        all_slices = torch.cat([sag, cor, axi], dim=1)\n        flat_slices = all_slices.view(B * 3 * S, C, H, W)\n        features = self.encoder(flat_slices).view(B, 3, S, -1)\n        \n        view_emb = self.view_plane_emb.view(1, 3, 1, -1)\n        slice_emb = self.slice_pos_emb.view(1, 1, S, -1)\n        tokens = (features + view_emb + slice_emb).view(B, 3 * S, -1)\n        \n        cross_tokens = self.cross_view_transformer(tokens).view(B, 3, S, -1)\n        sag_tokens = cross_tokens[:, 0]\n        cor_tokens = cross_tokens[:, 1]\n        axi_tokens = cross_tokens[:, 2]\n        \n        sag_pooled = self.sag_mil(sag_tokens, valid_mask[:, 0])\n        cor_pooled = self.cor_mil(cor_tokens, valid_mask[:, 1])\n        axi_pooled = self.axi_mil(axi_tokens, valid_mask[:, 2])\n        \n        meta_emb = self.meta_proj(meta)\n        fused = self.fusion(torch.cat([sag_pooled, cor_pooled, axi_pooled, meta_emb], dim=-1))\n        \n        node_feats = torch.stack([head(fused) for head in self.decoupled_heads], dim=1)\n        logits = self.prior_gcn(node_feats)\n        return logits\n\nclass ClassAdaptiveMarginASLLoss(nn.Module):\n    def __init__(self, prevalences=CFG.CLASS_PREVALENCES, gamma_pos=0.0, gamma_neg=2.0, eps=1e-8):\n        super().__init__()\n        self.gamma_pos = gamma_pos\n        self.gamma_neg = gamma_neg\n        self.eps = eps\n        margins = [0.02 + 0.10 * (1.0 - p) for p in prevalences]\n        self.register_buffer(\"margins\", torch.tensor(margins, dtype=torch.float32))\n\n    def forward(self, logits, targets, loss_weights=None):\n        probs = torch.sigmoid(logits)\n        targets = targets.type_as(probs)\n        pos_probs = torch.clamp(probs, min=self.eps, max=1.0 - self.eps)\n        loss_pos = targets * torch.pow(1.0 - pos_probs, self.gamma_pos) * torch.log(pos_probs)\n        shifted_neg_probs = torch.clamp(probs - self.margins.view(1, -1), min=0.0, max=1.0 - self.eps)\n        loss_neg = (1.0 - targets) * torch.pow(shifted_neg_probs, self.gamma_neg) * torch.log(1.0 - shifted_neg_probs + self.eps)\n        loss_per_class = -(loss_pos + loss_neg)\n        loss = loss_per_class.mean(dim=-1)\n        if loss_weights is not None:\n            loss = (loss * loss_weights).mean()\n        else:\n            loss = loss.mean()\n        return loss\n\n# ==============================================================================\n# MULTI-PIECE SEQUENTIAL DATASET TRAINER\n# ==============================================================================\ndef train_full_dataset_gpu():\n    print(\"=\" * 80, flush=True)\n    print(f\"🚀 RSNA HighAUCMultiViewNet-V2 Full Dataset GPU Engine (Device: {CFG.DEVICE})\", flush=True)\n    print(f\"📦 Training {CFG.PIECES_PER_SESSION} Consecutive Pieces ({CFG.STUDIES_PER_PIECE} studies each)\", flush=True)\n    print(\"=\" * 80, flush=True)\n\n    train_df = pd.read_csv(CFG.TRAIN_CSV)\n    series_df = pd.read_csv(CFG.SERIES_CSV) if os.path.exists(CFG.SERIES_CSV) else pd.DataFrame()\n    \n    pseudo_path = \"train_pseudo_labeled.csv\"\n    if os.path.exists(pseudo_path):\n        train_df = pd.read_csv(pseudo_path)\n    else:\n        for idx in range(len(train_df)):\n            has_missing = any(pd.isna(train_df.at[idx, c]) for c in CFG.TARGET_COLS if c in train_df.columns)\n            if has_missing:\n                report = str(train_df.at[idx, 'Report']) if 'Report' in train_df.columns else \"\"\n                parsed_vals = parse_report_fast(report)\n                for j, col in enumerate(CFG.TARGET_COLS):\n                    if col not in train_df.columns or pd.isna(train_df.at[idx, col]):\n                        train_df.at[idx, col] = parsed_vals[j]\n        for col in CFG.TARGET_COLS:\n            train_df[col] = train_df[col].fillna(0.10)\n            \n    path_index = {}\n    for _, row in series_df.iterrows():\n        st_id = str(row[\"StudyInstanceUID\"])\n        se_id = str(row[\"SeriesInstanceUID\"])\n        plane = str(row.get(\"Anatomical_Plane\", \"\")).lower()\n        fluid = float(row.get(\"Fluid_Sensitive\", 0.0) or 0.0)\n        fat = float(row.get(\"Fat_Suppression\", 0.0) or 0.0)\n        \n        if st_id not in path_index:\n            path_index[st_id] = {\"Sag\": [], \"Cor\": [], \"Ax\": [], \"meta\": [fluid, fat, 0.0]}\n            \n        cand_dirs = [\n            os.path.join(CFG.DATA_DIR, \"train_series\", st_id, se_id),\n            os.path.join(CFG.DATA_DIR, \"train_images\", st_id, se_id),\n            os.path.join(CFG.DATA_DIR, \"train_series\", se_id)\n        ]\n        for cdir in cand_dirs:\n            if os.path.exists(cdir):\n                dcm_files = [os.path.join(cdir, f) for f in os.listdir(cdir) if f.endswith(\".dcm\")]\n                if dcm_files:\n                    if \"sag\" in plane: path_index[st_id][\"Sag\"] = dcm_files\n                    elif \"cor\" in plane: path_index[st_id][\"Cor\"] = dcm_files\n                    elif \"ax\" in plane or \"tra\" in plane: path_index[st_id][\"Ax\"] = dcm_files\n                    break\n\n    unique_studies = train_df[\"StudyInstanceUID\"].unique()\n    total_studies = len(unique_studies)\n    total_pieces = math.ceil(total_studies / CFG.STUDIES_PER_PIECE)\n    \n    # Auto-detect latest piece\n    latest_piece = 0\n    for f in os.listdir(\".\"):\n        m = re.match(r'model_piece_(\\d+)\\.pth', f)\n        if m:\n            latest_piece = max(latest_piece, int(m.group(1)))\n            \n    start_piece = getattr(CFG, 'START_PIECE', (latest_piece + 1 if latest_piece < total_pieces else 1))\n    end_piece = min(total_pieces, start_piece + CFG.PIECES_PER_SESSION - 1)\n    \n    print(f\"🔄 Execution Plan: Training Pieces {start_piece} through {end_piece} (out of {total_pieces} total pieces)\", flush=True)\n\n    model = HighAUCMultiViewNet_V3(backbone=CFG.BACKBONE, num_slices=CFG.NUM_SLICES, meta_dim=3, num_classes=CFG.NUM_CLASSES).to(CFG.DEVICE)\n    \n    # Resume existing weights if available - Bulletproof Auto-Search!\n    resume_path = None\n    target_weight = f\"model_piece_{start_piece - 1}.pth\" if start_piece > 1 else None\n    \n    # Check current directory first\n    if target_weight and os.path.exists(target_weight):\n        resume_path = target_weight\n    # Check best model fallback\n    elif os.path.exists(\"best_hqcnn_model.pth\"):\n        resume_path = \"best_hqcnn_model.pth\"\n    # Recursively search Kaggle inputs for the target weight\n    elif target_weight:\n        for root, dirs, files in os.walk(\"/kaggle/input\"):\n            if target_weight in files:\n                resume_path = os.path.join(root, target_weight)\n                break\n                \n    if resume_path is None:\n        # Final fallback, look for best_hqcnn_model.pth anywhere\n        for root, dirs, files in os.walk(\"/kaggle/input\"):\n            if \"best_hqcnn_model.pth\" in files:\n                resume_path = os.path.join(root, \"best_hqcnn_model.pth\")\n                break\n                \n    if resume_path:\n        print(f\"🔄 Resuming base weights from {resume_path}...\", flush=True)\n        ckpt = torch.load(resume_path, map_location='cpu')\n        clean_dict = {str(k).replace(\"module.\", \"\").strip(): v for k, v in ckpt.items()}\n        model_dict = model.state_dict()\n        matched = {k: v for k, v in clean_dict.items() if k in model_dict and model_dict[k].shape == v.shape}\n        model_dict.update(matched)\n        model.load_state_dict(model_dict)\n        print(f\"✅ Loaded {len(matched)}/{len(model_dict)} layers from {resume_path}!\", flush=True)\n\n    criterion = ClassAdaptiveMarginASLLoss().to(CFG.DEVICE)\n    scaler = GradScaler()\n    best_overall_auc = 0.0\n\n    for current_piece in range(start_piece, end_piece + 1):\n        print(f\"\\n\" + \"=\" * 60, flush=True)\n        print(f\"📦 [PIECE {current_piece}/{total_pieces}] Studies {((current_piece-1)*CFG.STUDIES_PER_PIECE)} to {min(total_studies, current_piece*CFG.STUDIES_PER_PIECE)}\", flush=True)\n        print(\"=\" * 60, flush=True)\n\n        start_idx = (current_piece - 1) * CFG.STUDIES_PER_PIECE\n        end_idx = min(total_studies, current_piece * CFG.STUDIES_PER_PIECE)\n        piece_study_ids = set(unique_studies[start_idx:end_idx])\n        piece_df = train_df[train_df[\"StudyInstanceUID\"].isin(piece_study_ids)].reset_index(drop=True)\n        \n        if len(piece_df) < 16:\n            print(f\"ℹ️ Piece {current_piece} has only {len(piece_df)} remaining studies. Dataset 100% complete!\", flush=True)\n            break\n            \n        n_val = min(max(8, int(len(piece_df) * 0.15)), len(piece_df) // 2)\n        train_subset = piece_df.iloc[:-n_val].reset_index(drop=True)\n        val_subset = piece_df.iloc[-n_val:].reset_index(drop=True)\n        \n        train_ds = HighAUCDatasetV2(train_subset, path_index, is_train=True)\n        val_ds = HighAUCDatasetV2(val_subset, path_index, is_train=False)\n        \n        train_loader = DataLoader(train_ds, batch_size=CFG.BATCH_SIZE, shuffle=True, num_workers=2, pin_memory=(CFG.DEVICE.type == 'cuda'), drop_last=True)\n        val_loader = DataLoader(val_ds, batch_size=CFG.BATCH_SIZE, shuffle=False, num_workers=2, pin_memory=(CFG.DEVICE.type == 'cuda'), drop_last=False)\n        \n        optimizer = optim.AdamW(model.parameters(), lr=CFG.LR, weight_decay=CFG.WEIGHT_DECAY)\n        scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.EPOCHS_PER_PIECE, eta_min=1e-6)\n\n        best_piece_auc = 0.0\n        for epoch in range(1, CFG.EPOCHS_PER_PIECE + 1):\n            model.train()\n            total_loss = 0.0\n            train_steps = 0\n            \n            for sag, cor, axi, mask, meta, labels, loss_wt in train_loader:\n                sag = sag.to(CFG.DEVICE, non_blocking=True)\n                cor = cor.to(CFG.DEVICE, non_blocking=True)\n                axi = axi.to(CFG.DEVICE, non_blocking=True)\n                mask = mask.to(CFG.DEVICE, non_blocking=True)\n                meta = meta.to(CFG.DEVICE, non_blocking=True)\n                labels = labels.to(CFG.DEVICE, non_blocking=True)\n                loss_wt = loss_wt.to(CFG.DEVICE, non_blocking=True)\n                \n                optimizer.zero_grad()\n                with autocast():\n                    logits = model(sag, cor, axi, mask, meta)\n                    loss = criterion(logits, labels, loss_weights=loss_wt)\n                    \n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n                \n                total_loss += loss.item()\n                train_steps += 1\n                \n            scheduler.step()\n            avg_train_loss = total_loss / max(1, train_steps)\n            \n            # Validation\n            model.eval()\n            val_preds, val_targets = [], []\n            with torch.no_grad():\n                for sag, cor, axi, mask, meta, labels, _ in val_loader:\n                    sag, cor, axi = sag.to(CFG.DEVICE), cor.to(CFG.DEVICE), axi.to(CFG.DEVICE)\n                    mask, meta = mask.to(CFG.DEVICE), meta.to(CFG.DEVICE)\n                    with autocast():\n                        logits = model(sag, cor, axi, mask, meta)\n                        probs = torch.sigmoid(logits)\n                    val_preds.append(probs.cpu().numpy())\n                    val_targets.append(labels.numpy())\n                    \n            if len(val_preds) > 0:\n                val_p = np.vstack(val_preds)\n                val_t = np.vstack(val_targets)\n                val_t_bin = (val_t > 0.5).astype(int)\n                aucs = []\n                for c in range(12):\n                    if len(np.unique(val_t_bin[:, c])) > 1:\n                        aucs.append(roc_auc_score(val_t_bin[:, c], val_p[:, c]))\n                val_auc = np.mean(aucs) if aucs else 0.50\n            else:\n                val_auc = 0.50\n                \n            print(f\"  [Piece {current_piece} | Epoch {epoch}/{CFG.EPOCHS_PER_PIECE}] Loss: {avg_train_loss:.4f} | Val Macro AUC: {val_auc:.4f}\", flush=True)\n            if val_auc > best_overall_auc or epoch == CFG.EPOCHS_PER_PIECE:\n                best_overall_auc = max(best_overall_auc, val_auc)\n                torch.save(model.state_dict(), \"best_hqcnn_model.pth\")\n                torch.save(model.state_dict(), f\"model_piece_{current_piece}.pth\")\n                print(f\"  🌟 Checkpoint Saved -> 'best_hqcnn_model.pth' & 'model_piece_{current_piece}.pth' (AUC: {val_auc:.4f})\", flush=True)\n\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n            gc.collect()\n\n        print(f\"✅ Piece {current_piece} Completed! Advancing seamlessly to Piece {current_piece + 1}...\", flush=True)\n\n    print(f\"\\n🎉 Full Dataset Multi-Piece Training Session Complete! Peak AUC: {best_overall_auc:.4f}\", flush=True)\n\nif __name__ == '__main__':\n    train_full_dataset_gpu()\n","metadata":{},"outputs":[],"execution_count":null}]}