{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":99552,"databundleVersionId":13441085,"sourceType":"competition"}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ---------------------------\n#  Import libraries\n# ---------------------------\nimport os\nfrom collections import OrderedDict\nfrom typing import List, Tuple, Dict, Optional\nfrom concurrent.futures import ProcessPoolExecutor, as_completed\n\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models.detection import MaskRCNN\nfrom torchvision.ops.feature_pyramid_network import FeaturePyramidNetwork, LastLevelMaxPool\nfrom torchvision.transforms.functional import normalize\nimport timm\nimport pydicom\n\n# ---------------------------\n#  Hard-coded dataset references\n# ---------------------------\nSERIES_ROOT = \"/kaggle/input/rsna-intracranial-aneurysm-detection/series/\"\nTRAIN_CSV   = \"/kaggle/input/rsna-intracranial-aneurysm-detection/train.csv\"\n\n# ---------------------------\n#  Config (assume CUDA + AMP)\n# ---------------------------\nDEVICE = torch.device(\"cuda\")\ntorch.backends.cudnn.benchmark = True\ntorch.set_float32_matmul_precision(\"high\")\n\nBACKBONE = \"efficientvit_b1\"\nOUT_INDICES = (0, 1, 2, 3)\nFPN_OUT = 256\nNUM_CLASSES = 2\nBATCH_SIZE = 2\nEPOCHS = 20\nLR = 5e-4\nWEIGHT_DECAY = 1e-4\nPRINT_FREQ = 50\nLOG_AVG = 128\nAMP = True\nCPU_COUNT = max(1, os.cpu_count() or 1)\n\n# ---------------------------\n#  Slice selection policy\n# ---------------------------\nSLICES_PER_SERIES = 3\n\n# ---------------------------\n#  EfficientViT backbone + FPN\n# ---------------------------\nclass EfficientViTBackboneWithFPN(nn.Module):\n    def __init__(self, model_name: str, out_indices: Tuple[int, ...], out_channels: int):\n        super().__init__()\n        self.body = timm.create_model(\n            model_name,\n            features_only=True,\n            out_indices=out_indices,\n            pretrained=True\n        )\n        in_channels_list = self.body.feature_info.channels()\n        self.fpn = FeaturePyramidNetwork(\n            in_channels_list=in_channels_list,\n            out_channels=out_channels,\n            extra_blocks=LastLevelMaxPool()\n        )\n        self.out_channels = out_channels\n\n    def forward(self, x: torch.Tensor) -> OrderedDict:\n        feats = self.body(x)\n        feat_dict = OrderedDict({str(i): f for i, f in enumerate(feats)})\n        return self.fpn(feat_dict)\n\n# ---------------------------\n#  Build Mask R-CNN with EfficientViT+FPN\n# ---------------------------\ndef build_model() -> MaskRCNN:\n    backbone = EfficientViTBackboneWithFPN(BACKBONE, OUT_INDICES, FPN_OUT)\n    model = MaskRCNN(backbone, num_classes=NUM_CLASSES)\n    return model\n\n# ---------------------------\n#  DICOM sorting helpers\n# ---------------------------\ndef _safe_get_instance_number(d) -> float:\n    value = getattr(d, \"InstanceNumber\", None)\n    return float(value) if value is not None else 0.0\n\ndef _safe_get_z(d) -> Optional[float]:\n    ipp = getattr(d, \"ImagePositionPatient\", None)\n    if ipp is None:\n        return None\n    try:\n        return float(ipp[2])\n    except Exception:\n        return None\n\ndef _sort_dcm_paths_by_z(dcm_paths: List[str]) -> List[str]:\n    try:\n        mets = []\n        for p in dcm_paths:\n            d = pydicom.dcmread(p, stop_before_pixels=True, force=True)\n            z = _safe_get_z(d)\n            key = z if z is not None else _safe_get_instance_number(d)\n            mets.append((key, p))\n        mets.sort(key=lambda t: t[0])\n        return [p for _, p in mets]\n    except Exception:\n        return sorted(dcm_paths)\n\n# ---------------------------\n#  Per-series indexing helper (returns multiple well-placed slices)\n# ---------------------------\ndef _pick_well_placed_indices(n: int, k: int) -> List[int]:\n    if n <= 0:\n        return []\n    if n < 3 or k == 1:\n        return [n // 2]\n    if k == 3:\n        return [n // 4, n // 2, (3 * n) // 4]\n    if k == 5:\n        return [n // 6, n // 3, n // 2, (2 * n) // 3, (5 * n) // 6]\n    step = max(1, n // (k + 1))\n    centers = [step * (i + 1) for i in range(k)]\n    centers = [min(n - 1, max(0, c)) for c in centers]\n    return sorted(list(dict.fromkeys(centers)))\n\ndef _index_one_series(args: Tuple[str, str, int, int]) -> Optional[List[Tuple[str, int]]]:\n    series_root, sid, flag, k = args\n    series_dir = os.path.join(series_root, sid)\n    if not os.path.isdir(series_dir):\n        return None\n\n    dcm_paths = []\n    for r, _, files in os.walk(series_dir):\n        for f in files:\n            if f.lower().endswith(\".dcm\"):\n                dcm_paths.append(os.path.join(r, f))\n    if not dcm_paths:\n        return None\n\n    try:\n        sorted_paths = _sort_dcm_paths_by_z(dcm_paths)\n    except Exception:\n        sorted_paths = sorted(dcm_paths)\n\n    n = len(sorted_paths)\n    idxs = _pick_well_placed_indices(n, k)\n    picks = [sorted_paths[i] for i in idxs]\n    return [(p, int(flag)) for p in picks]\n\n# ---------------------------\n#  RSNA CTA slice dataset with parallel indexing\n# ---------------------------\nclass RSNACTAMaskDataset(Dataset):\n    def __init__(self, series_root: str, csv_path: str, slices_per_series: int = SLICES_PER_SERIES):\n        self.root = series_root\n        self.df = pd.read_csv(csv_path)\n        self.df = self.df[self.df[\"Modality\"].astype(str).str.upper() == \"CTA\"]\n        self.series_ids = self.df[\"SeriesInstanceUID\"].astype(str).tolist()\n        self.ap_flags = self.df[\"Aneurysm Present\"].astype(float).fillna(0.0).astype(int).tolist()\n        self.series_to_flag = dict(zip(self.series_ids, self.ap_flags))\n        self.k = max(1, int(slices_per_series))\n\n        print(f\"[Index] Starting parallel indexing with {CPU_COUNT} processes over {len(self.series_ids)} series...\")\n        self.samples = self._parallel_index_slices(self.series_ids, self.series_to_flag, self.k)\n        print(f\"[Index] Completed. Indexed {len(self.samples)} slice paths from {len(self.series_ids)} series.\")\n\n    def _parallel_index_slices(self, series_ids: List[str], flag_map: Dict[str, int], k: int) -> List[Tuple[str, int]]:\n        tasks = [(self.root, sid, flag_map.get(sid, 0), k) for sid in series_ids]\n        out: List[Tuple[str, int]] = []\n        with ProcessPoolExecutor(max_workers=CPU_COUNT) as ex:\n            futures = [ex.submit(_index_one_series, t) for t in tasks]\n            done = 0\n            for fut in as_completed(futures):\n                res = fut.result()\n                if res is not None:\n                    out.extend(res)\n                done += 1\n                if done % 256 == 0:\n                    print(f\"[Index] processed {done}/{len(futures)} series...\")\n        return out\n\n    def __len__(self) -> int:\n        return len(self.samples)\n\n    def _read_dicom_image(self, path: str) -> torch.Tensor:\n        dcm = pydicom.dcmread(path, force=True)\n        arr = dcm.pixel_array.astype(\"float32\")\n        m = float(arr.mean())\n        s = float(arr.std()) + 1e-6\n        arr = (arr - m) / s\n        t = torch.from_numpy(arr)\n        if t.ndim == 2:\n            t = t.unsqueeze(0).repeat(3, 1, 1)\n        elif t.ndim == 3 and t.shape[0] != 3:\n            t = t.permute(2, 0, 1)\n            if t.shape[0] == 1:\n                t = t.repeat(3, 1, 1)\n        return normalize(t, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n\n    def __getitem__(self, idx: int):\n        dcm_path, ap = self.samples[idx]\n        image = self._read_dicom_image(dcm_path)\n        H = int(image.shape[1])\n        W = int(image.shape[2])\n\n        if ap > 0:\n            cx = W // 2\n            cy = H // 2\n            rw = max(16, int(W * 0.6))\n            rh = max(16, int(H * 0.6))\n            x1 = max(0, cx - rw // 2)\n            y1 = max(0, cy - rh // 2)\n            x2 = min(W - 1, x1 + rw)\n            y2 = min(H - 1, y1 + rh)\n            boxes = torch.tensor([[float(x1), float(y1), float(x2), float(y2)]], dtype=torch.float32)\n            labels = torch.tensor([1], dtype=torch.int64)\n            masks = torch.zeros((1, H, W), dtype=torch.uint8)\n            masks[0, y1:y2, x1:x2] = 1\n        else:\n            boxes = torch.zeros((0, 4), dtype=torch.float32)\n            labels = torch.zeros((0,), dtype=torch.int64)\n            masks = torch.zeros((0, H, W), dtype=torch.uint8)\n\n        areas = (boxes[:, 2] - boxes[:, 0]) * (boxes[:, 3] - boxes[:, 1]) if boxes.numel() else torch.zeros((0,), dtype=torch.float32)\n        target = {\n            \"boxes\": boxes,\n            \"labels\": labels,\n            \"masks\": masks,\n            \"image_id\": torch.tensor([idx], dtype=torch.int64),\n            \"area\": areas,\n            \"iscrowd\": torch.zeros((labels.shape[0],), dtype=torch.int64),\n        }\n        return image, target\n\n# ---------------------------\n#  Data utilities\n# ---------------------------\ndef collate_fn(batch):\n    return tuple(zip(*batch))\n\n# ---------------------------\n#  One training epoch\n# ---------------------------\ndef train_one_epoch(model, loader, optimizer, scaler, epoch):\n    model.train()\n    running_loss = 0.0\n    for i, (images, targets) in enumerate(loader):\n        images = [img.to(DEVICE, non_blocking=True) for img in images]\n        targets = [{k: v.to(DEVICE, non_blocking=True) for k, v in t.items()} for t in targets]\n        optimizer.zero_grad(set_to_none=True)\n        with torch.amp.autocast(device_type=\"cuda\", enabled=AMP):\n            loss_dict = model(images, targets)\n            losses = sum(loss for loss in loss_dict.values())\n        scaler.scale(losses).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        running_loss += losses.item()\n        if (i + 1) % PRINT_FREQ == 0:\n            print(f\"[Train] epoch={epoch} iter={i+1}/{len(loader)} loss={losses.item():.4f}\")\n        if (i + 1) % LOG_AVG == 0:\n            avg_loss = running_loss / LOG_AVG\n            print(f\"[Train] epoch={epoch} iter={i+1} avg_loss(last {LOG_AVG})={avg_loss:.4f}\")\n            running_loss = 0.0\n\n# ---------------------------\n#  Evaluation using train set (compute loss)\n# ---------------------------\n@torch.no_grad()\ndef evaluate_loss(model, loader, epoch):\n    model.train()\n    total = 0.0\n    count = 0\n    for i, (images, targets) in enumerate(loader):\n        images = [img.to(DEVICE, non_blocking=True) for img in images]\n        targets = [{k: v.to(DEVICE, non_blocking=True) for k, v in t.items()} for t in targets]\n        with torch.amp.autocast(device_type=\"cuda\", enabled=AMP):\n            loss_dict = model(images, targets)\n            loss = sum(loss for loss in loss_dict.values())\n        total += float(loss.item())\n        count += 1\n        if (i + 1) % PRINT_FREQ == 0:\n            print(f\"[Eval] epoch={epoch} iter={i+1}/{len(loader)} loss={loss.item():.4f}\")\n    return total / max(1, count)\n\n# ---------------------------\n#  Main training loop\n# ---------------------------\ndef main():\n    print(f\"[Setup] Using {CPU_COUNT} CPU workers for dataset indexing and DataLoader.\")\n    dataset = RSNACTAMaskDataset(series_root=SERIES_ROOT, csv_path=TRAIN_CSV, slices_per_series=SLICES_PER_SERIES)\n    print(f\"[Data] Loaded dataset with {len(dataset)} slice samples\")\n\n    train_loader = DataLoader(\n        dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=True,\n        num_workers=CPU_COUNT,\n        pin_memory=True,\n        collate_fn=collate_fn,\n        persistent_workers=False\n    )\n\n    val_loader = DataLoader(\n        dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        num_workers=CPU_COUNT,\n        pin_memory=True,\n        collate_fn=collate_fn,\n        persistent_workers=False\n    )\n\n    print(\"[Model] Building EfficientViT+FPN Mask R-CNN...\")\n    model = build_model().to(DEVICE)\n    print(\"[Model] Built. Starting training...\")\n\n    params = [p for p in model.parameters() if p.requires_grad]\n    optimizer = torch.optim.AdamW(params, lr=LR, weight_decay=WEIGHT_DECAY)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n    scaler = torch.cuda.amp.GradScaler(enabled=AMP)\n\n    best = float(\"inf\")\n    for epoch in range(1, EPOCHS + 1):\n        print(f\"\\n=== Epoch {epoch}/{EPOCHS} ===\")\n        train_one_epoch(model, train_loader, optimizer, scaler, epoch)\n        val_loss = evaluate_loss(model, val_loader, epoch)\n        print(f\"[Eval] epoch={epoch} avg_val_loss={val_loss:.4f}\")\n        scheduler.step()\n        if val_loss < best:\n            best = val_loss\n            torch.save(model.state_dict(), \"maskrcnn_efficientvit_best.pth\")\n            print(f\"[Checkpoint] Saved new best model at epoch={epoch} val_loss={val_loss:.4f}\")\n\n# ---------------------------\n#  Entrypoint\n# ---------------------------\nif __name__ == \"__main__\":\n    main()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}