{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"jupytext":{"cell_metadata_filter":"-all","main_language":"python","notebook_metadata_filter":"-all"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":99552,"databundleVersionId":13190393,"sourceType":"competition"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nimport pydicom\nimport torch.nn.functional as F\nfrom multiprocessing import Pool, cpu_count\nfrom collections import defaultdict\nimport time\nimport cv2\nfrom sklearn.model_selection import StratifiedShuffleSplit\nfrom pydicom.errors import InvalidDicomError\nfrom pydicom import dcmread\nimport torchvision.models as models\nimport random","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T13:02:46.073277Z","iopub.execute_input":"2025-08-07T13:02:46.073612Z","iopub.status.idle":"2025-08-07T13:02:46.078707Z","shell.execute_reply.started":"2025-08-07T13:02:46.073593Z","shell.execute_reply":"2025-08-07T13:02:46.077869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q pylibjpeg[all]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T13:02:46.079727Z","iopub.execute_input":"2025-08-07T13:02:46.079923Z","iopub.status.idle":"2025-08-07T13:02:49.239498Z","shell.execute_reply.started":"2025-08-07T13:02:46.079908Z","shell.execute_reply":"2025-08-07T13:02:49.238377Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**1) Data storage and Test train split. **\n\nCreate a list of tuples containing the file path followed by the modality vector.","metadata":{}},{"cell_type":"code","source":"# 1. Map modalities to one-hot encoded vectors\ndef get_modality_vector(modality, modality_list=None):\n    # Predefined or discovered modality list\n    if modality_list is None:\n        modality_list = ['CT', 'MR']  # Extend this as needed\n    one_hot = [0] * len(modality_list)\n    if modality in modality_list:\n        index = modality_list.index(modality)\n        one_hot[index] = 1\n    return one_hot\n\n# 2. Scan directories and collect (path, modality_vector) pairs\ndef collect_dicom_paths_and_modalities(root_path, modality_list=None):\n    path_modality_pairs = []\n    seen_modalities = set()\n\n    for dirpath, dirnames, filenames in os.walk(root_path):\n        # Skip root, we want only leaf folders with DICOMs\n        if dirpath == root_path or not filenames:\n            continue\n\n        # Try to find a .dcm file\n        dcm_file = next((f for f in filenames if f.lower().endswith('.dcm')), None)\n        if not dcm_file:\n            continue\n\n        try:\n            dcm_path = os.path.join(dirpath, dcm_file)\n            ds = pydicom.dcmread(dcm_path, stop_before_pixels=True)\n            modality = ds.get(\"Modality\", \"Unknown\")\n\n            if modality_list is None:\n                seen_modalities.add(modality)\n\n            vector = get_modality_vector(modality, modality_list)\n            path_modality_pairs.append((dirpath, vector, modality))\n\n        except Exception as e:\n            print(f\"Skipping {dirpath} due to error: {e}\")\n            continue\n\n    # If no modality_list was given, return all discovered ones too\n    if modality_list is None:\n        return path_modality_pairs, sorted(list(seen_modalities))\n    else:\n        return path_modality_pairs\n\n# 3. Stratified train-test split\ndef stratified_split(pairs, test_size=0.3, random_state=42):\n    paths = [x[0] for x in pairs]\n    modalities = [x[2] for x in pairs]  # original modality name, not one-hot\n\n    splitter = StratifiedShuffleSplit(n_splits=1, test_size=test_size, random_state=random_state)\n    for train_idx, test_idx in splitter.split(paths, modalities):\n        train_data = [pairs[i] for i in train_idx]\n        test_data = [pairs[i] for i in test_idx]\n        return train_data, test_data\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T13:02:49.241072Z","iopub.execute_input":"2025-08-07T13:02:49.241322Z","iopub.status.idle":"2025-08-07T13:02:49.250659Z","shell.execute_reply.started":"2025-08-07T13:02:49.241298Z","shell.execute_reply":"2025-08-07T13:02:49.249761Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"2) Dataset class.\n\nThis here creates a dataset class that takes in the path and extracts, resizes the slice.","metadata":{}},{"cell_type":"code","source":"def safe_read_dicom(path):\n    try:\n        ds = dcmread(path)\n        arr = ds.pixel_array  # This triggers decoding\n        return path\n    except Exception as e:\n        print(f\"⚠️ Skipping {path}: {e}\")\n        return None  # Signal an unreadable image\n\ndef load_dicom_as_grayscale_tensor(dcm_path):\n    \"\"\"\n    Loads any DICOM file, converts to grayscale if needed, takes middle slice if 3D,\n    and returns a tensor of shape [1, 224, 224]\n    \"\"\"\n    # Load DICOM\n    ds = pydicom.dcmread(dcm_path)\n\n    # Get pixel array\n    img = ds.pixel_array.astype(np.float32)\n\n    # Handle 4D or RGB DICOMs\n    if img.ndim == 4:\n        # E.g., (D, H, W, 3) → take middle slice and convert to grayscale\n        z = img.shape[0] // 2\n        img = img[z]\n    \n    if img.ndim == 3:\n        if img.shape[-1] == 3:\n            # Color image (H, W, 3) or (D, H, W, 3)\n            img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n        else:\n            # 3D grayscale (D, H, W)\n            z = img.shape[0] // 2\n            img = img[z]  # Middle slice\n\n    elif img.ndim == 2:\n        # Already 2D grayscale\n        pass\n    else:\n        raise ValueError(f\"Unsupported image shape: {img.shape}\")\n\n    # Normalize to [0, 1]\n    img = (img - np.min(img)) / (np.max(img) - np.min(img) + 1e-5)\n\n    # Resize to (224, 224)\n    img = cv2.resize(img, (224, 224))\n\n    # Convert to PyTorch tensor with shape [1, 256, 256]\n    tensor = torch.tensor(img, dtype=torch.float32).unsqueeze(0)\n\n    return tensor\n\n\ndef get_random_middle_slice_path(path):\n    # Collect (full_path, instance_number) pairs\n    dicom_info = []\n    for fname in os.listdir(path):\n        if fname.lower().endswith('.dcm'):\n            full_path = os.path.join(path, fname)\n            try:\n                dcm = pydicom.dcmread(full_path, stop_before_pixels=True)\n                instance = int(dcm.InstanceNumber)\n                dicom_info.append((full_path, instance))\n            except Exception as e:\n                print(f\"Skipping file {fname}: {e}\")\n    \n    if len(dicom_info) < 10:\n        dicom_info.sort(key=lambda x: x[1])\n        mid_idx = len(dicom_info) // 2\n        return dicom_info[mid_idx][0]\n\n    # Sort by InstanceNumber\n    dicom_info.sort(key=lambda x: x[1])\n\n    # Extract the middle 10%\n    n = len(dicom_info)\n    lower = int(n * 0.45)\n    upper = int(n * 0.55)\n    middle_slices = dicom_info[lower:upper]\n\n    # Choose a random one and return the path\n    random_slice_path = random.choice(middle_slices)[0]\n    return random_slice_path\n\n\nclass DicomSliceDataset(Dataset):\n    def __init__(self, data_pairs, transform=None):\n        \"\"\"\n        data_pairs: list of tuples like (path, one_hot_modality_vector, modality_name)\n        \"\"\"\n        self.data_pairs = data_pairs\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data_pairs)\n\n    def __getitem__(self, idx):\n        path, label_vector, _ = self.data_pairs[idx]\n\n\n        # Load middle slice\n        dcm_path = get_random_middle_slice_path(path)\n        ######################################################################################\n        dcm_path = safe_read_dicom(dcm_path)\n        \n        if dcm_path is None:\n            # Skip or fallback: return a blank tensor or raise StopIteration\n            label_dim = 2\n            return torch.zeros(1, 224, 224), torch.zeros(label_dim)\n\n        ######################################################################################\n        \n\n        img_tensor = load_dicom_as_grayscale_tensor(dcm_path)\n\n        # Convert label to tensor\n        label_tensor = torch.tensor(label_vector, dtype=torch.float32)\n\n        return img_tensor, label_tensor\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T13:02:49.2514Z","iopub.execute_input":"2025-08-07T13:02:49.2517Z","iopub.status.idle":"2025-08-07T13:02:49.271373Z","shell.execute_reply.started":"2025-08-07T13:02:49.251674Z","shell.execute_reply":"2025-08-07T13:02:49.270694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ModalityClassifierCNN(nn.Module):\n    def __init__(self, num_classes):\n        super(ModalityClassifierCNN, self).__init__()\n        self.conv1 = nn.Conv2d(1, 16, 3, padding=1)\n        self.conv2 = nn.Conv2d(16, 32, 3, padding=1)\n        self.conv3 = nn.Conv2d(32, 64, 3, padding=1)\n        self.pool = nn.MaxPool2d(2, 2)\n        self.fc1 = nn.Linear(64 * 28 * 28, 128)\n        self.fc2 = nn.Linear(128, num_classes)\n\n    def forward(self, x):\n        x = self.pool(F.relu(self.conv1(x)))  # -> [16, 128, 128]\n        x = self.pool(F.relu(self.conv2(x)))  # -> [32, 64, 64]\n        x = self.pool(F.relu(self.conv3(x)))  # -> [64, 32, 32]\n        x = x.view(-1, 64 * 28 * 28)          # Flatten\n        x = F.relu(self.fc1(x))\n        x = self.fc2(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T13:02:49.272705Z","iopub.execute_input":"2025-08-07T13:02:49.272924Z","iopub.status.idle":"2025-08-07T13:02:49.289681Z","shell.execute_reply.started":"2025-08-07T13:02:49.272909Z","shell.execute_reply":"2025-08-07T13:02:49.288937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Split train/test\n# Step 1 & 2: Get data\nroot_dir = \"/kaggle/input/rsna-intracranial-aneurysm-detection/series\"\npairs, modality_list = collect_dicom_paths_and_modalities(root_dir)\n\n# Step 3: Stratified split\ntrain_data, test_data = stratified_split(pairs)\ntrain_dataset = DicomSliceDataset(train_data)\ntest_dataset = DicomSliceDataset(test_data)\n\n\ntrain_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)\ntest_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T13:02:49.290416Z","iopub.execute_input":"2025-08-07T13:02:49.290645Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nnum_classes = len(train_data[0][1])  # Length of one-hot vector\nmodel = ModalityClassifierCNN(num_classes).to(device)\n\ncriterion = nn.BCEWithLogitsLoss()  # Better for multi-label or one-hot vectors\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate(model, loader, device):\n    model.eval()\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for images, labels in loader:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            preds = torch.argmax(outputs, dim=1)\n            targets = torch.argmax(labels, dim=1)\n            correct += (preds == targets).sum().item()\n            total += labels.size(0)\n    print(f\"Accuracy: {correct / total * 100:.2f}%\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train(model, loader, criterion, optimizer, device, epochs=10):\n    model.train()\n    for epoch in range(epochs):\n        total_loss = 0\n        for images, labels in loader:\n            images, labels = images.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(images)\n\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n\n            total_loss += loss.item()\n\n        avg_loss = total_loss / len(loader)\n        print(f\"Epoch {epoch+1}/{epochs}, Loss: {avg_loss:.4f}\")\n        print(f\"Epoch {epoch+1}/{epochs}, Test \")\n        evaluate(model, test_loader, device)\n        if epoch % 5 == 0:  # Save every 5 epochs\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'loss': loss,\n            }, f'/kaggle/working/checkpoint_epoch_{epoch}.pth')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate(model, loader, device):\n    model.eval()\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for images, labels in loader:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            preds = torch.argmax(outputs, dim=1)\n            targets = torch.argmax(labels, dim=1)\n            correct += (preds == targets).sum().item()\n            total += labels.size(0)\n    print(f\"Accuracy: {correct / total * 100:.2f}%\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train(model, train_loader, criterion, optimizer, device, epochs=10)\ntorch.save(model.state_dict(), '/kaggle/working/final_model_weights.pth')\nprint(\"Test:\")\nevaluate(model, test_loader, device)\nprint(\"Train:\")\nevaluate(model, train_loader, device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The evaluation API requires that you set up a server which will respond to inference requests. We have already defined the server; you just need write the predict function. When we evaluate your submission on the hidden test set the client defined in `rsna_gateway` will run in a different container with direct access to the hidden test set and hand off the data series by series.\n\nYour code will always have access to the published copies of the files.","metadata":{}},{"cell_type":"markdown","source":"When your notebook is run on the hidden test set, `inference_server.serve` must be called within 15 minutes of the notebook starting or the gateway will throw an error. If you need more than 15 minutes to load your model you can do so during the very first `predict` call.","metadata":{}}]}