{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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"}],"dockerImageVersionId":30822,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom collections import Counter\nimport json\nimport numpy as np\nimport torch\nimport pydicom\nimport cv2\nfrom tqdm.auto import tqdm\nimport torchvision.transforms as T\nimport pandas as pd\nfrom collections import defaultdict\nfrom torch.utils.data import Dataset\nimport torch.nn as nn\nfrom torchvision.models import resnet50, efficientnet_b0\nimport torch.optim as optim\n\n\n# Define paths\nbase_path = Path(\"../input/rsna-2024-lumbar-spine-degenerative-classification\")\noutput_path = Path(\"../working\")\ntrain_csv_path = base_path / \"train.csv\"\nlabel_coordinates_csv_path = base_path / \"train_label_coordinates.csv\"\nseries_descriptions_csv_path = base_path / \"train_series_descriptions.csv\"\ntrain_images_path = base_path / \"train_images\"\n\n# Load CSV files\ntrain_df = pd.read_csv(train_csv_path)\nlabel_coords_df = pd.read_csv(label_coordinates_csv_path)\nseries_desc_df = pd.read_csv(series_descriptions_csv_path)\n\n# Display a few rows from each CSV\nprint(\"Train CSV Sample:\")\ndisplay(train_df.head())\n\nprint(\"Label Coordinates CSV Sample:\")\ndisplay(label_coords_df.head())\n\nprint(\"Series Descriptions CSV Sample:\")\ndisplay(series_desc_df.head())\n\nprint(\"Coordinates nan counts\")\ndisplay(label_coords_df.isna().sum())\n# Check class distribution\nseverity_cols = [\n    col for col in train_df.columns if \"stenosis\" in col or \"narrowing\" in col\n]\nseverity_counts = train_df[severity_cols].apply(pd.Series.value_counts).sum(axis=1)\nprint(\"\\nClass Distribution Across Severity Levels:\")\nprint(severity_counts)\n\n# Visualize class distribution\nseverity_counts.plot(kind=\"bar\", figsize=(10, 6), title=\"Severity Class Distribution\")\nplt.show()\n\n# Explore a single study's image files\nsample_study_id = train_df.iloc[0][\"study_id\"]\nsample_study_path = train_images_path / str(sample_study_id)\nsample_series = os.listdir(sample_study_path)\nprint(f\"\\nSample Study ID: {sample_study_id}\")\nprint(f\"Available Series in the Study: {sample_series}\")\n\n# Visualize a single image\ndef show_dicom_image(series_path):\n    dicom_files = list(Path(series_path).glob(\"*.dcm\"))\n    if not dicom_files:\n        print(f\"No DICOM files found in {series_path}\")\n        return\n    sample_dicom = pydicom.dcmread(dicom_files[0])\n    image = sample_dicom.pixel_array\n    plt.figure(figsize=(6, 6))\n    plt.imshow(image, cmap=\"gray\")\n    plt.title(f\"DICOM Image from {series_path}\")\n    plt.axis(\"off\")\n    plt.show()\n\nfor series_id in sample_series:\n    print(f\"\\nDisplaying series: {series_id}\")\n    show_dicom_image(sample_study_path / series_id)\n\ndel sample_study_id, sample_study_path, sample_series, severity_counts, train_df, label_coords_df, series_desc_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-12T13:42:38.703385Z","iopub.execute_input":"2025-01-12T13:42:38.703694Z","iopub.status.idle":"2025-01-12T13:42:39.865451Z","shell.execute_reply.started":"2025-01-12T13:42:38.703662Z","shell.execute_reply":"2025-01-12T13:42:39.864523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\nfigure, axis = plt.subplots(1,3, figsize=(20,5)) \nfor idx, d in enumerate(['foraminal', 'subarticular', 'canal']):\n    diagnosis = list(filter(lambda x: x.find(d) > -1, train.columns))\n    dff = train[diagnosis]\n    \n    value_counts = dff.apply(pd.value_counts).fillna(0).T\n    value_counts.plot(kind='bar', stacked=True, ax=axis[idx])\n    axis[idx].set_title(f'{d} distribution')\ndel train","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-12T13:43:00.116232Z","iopub.execute_input":"2025-01-12T13:43:00.116514Z","iopub.status.idle":"2025-01-12T13:43:01.087317Z","shell.execute_reply.started":"2025-01-12T13:43:00.116493Z","shell.execute_reply":"2025-01-12T13:43:01.086492Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SpineDataProcessor:\n\n    def __init__(self, base_path, output_path, roi_size=224):\n\n        self.base_path = Path(base_path)\n\n        self.output_path = Path(output_path)\n\n        self.roi_size = roi_size\n\n        self.processed_data = {\"samples\": {}, \"processed_rois\": {}}\n\n        # Load metadata\n\n        self.train_df = pd.read_csv(self.base_path / \"train.csv\")\n\n        self.coords_df = pd.read_csv(self.base_path / \"train_label_coordinates.csv\")\n\n        self.series_df = pd.read_csv(self.base_path / \"train_series_descriptions.csv\")\n\n    def extract_roi(self, image, points):\n        \"\"\"\n        Extract ROI that contains all points while maintaining aspect ratio\n        points: list of (x,y) coordinates for conditions in the same image\n        \"\"\"\n        h, w = image.shape\n\n        # Get bounding box that contains all points\n        x_coords = [p[0] for p in points]\n        y_coords = [p[1] for p in points]\n\n        # Add padding to ensure we capture enough context\n        padding = 32\n        min_x = max(0, int(min(x_coords) - padding))\n        max_x = min(w, int(max(x_coords) + padding))\n        min_y = max(0, int(min(y_coords) - padding))\n        max_y = min(h, int(max(y_coords) + padding))\n\n        # Extract region containing all points\n        roi = image[min_y:max_y, min_x:max_x]\n\n        # Calculate padding needed to maintain aspect ratio\n        roi_h, roi_w = roi.shape\n        if roi_h > roi_w:\n            # Add padding to width\n            target_w = int(roi_h * (224 / 224))  # maintain square aspect ratio\n            pad_w = target_w - roi_w\n            pad_left = pad_w // 2\n            pad_right = pad_w - pad_left\n            roi = np.pad(roi, ((0, 0), (pad_left, pad_right)), mode=\"constant\")\n        else:\n            # Add padding to height\n            target_h = int(roi_w * (224 / 224))\n            pad_h = target_h - roi_h\n            pad_top = pad_h // 2\n            pad_bottom = pad_h - pad_top\n            roi = np.pad(roi, ((pad_top, pad_bottom), (0, 0)), mode=\"constant\")\n\n        # Resize to 224x224 while maintaining aspect ratio\n        roi_resized = cv2.resize(roi, (224, 224), interpolation=cv2.INTER_AREA)\n\n        # Calculate scaled coordinates for all points\n        scaled_points = []\n        for x, y in points:\n            # Adjust for padding and scaling\n            if roi_h > roi_w:\n                scaled_x = ((x - min_x + pad_left) * 224) / roi.shape[1]\n                scaled_y = ((y - min_y) * 224) / roi.shape[0]\n            else:\n                scaled_x = ((x - min_x) * 224) / roi.shape[1]\n                scaled_y = ((y - min_y + pad_top) * 224) / roi.shape[0]\n            scaled_points.append((scaled_x, scaled_y))\n\n        return roi_resized, scaled_points\n\n    def augment_image(self, image, num_augmentations=4):\n        \"\"\"Apply basic augmentations for moderate/severe cases\"\"\"\n        augmented = []\n        for _ in range(num_augmentations):\n            # Convert to tensor for torchvision transforms\n            img_tensor = torch.from_numpy(image).unsqueeze(0)\n\n            # Apply random transformations\n            transforms = T.Compose(\n                [\n                    T.RandomHorizontalFlip(p=0.5),\n                    T.GaussianBlur(kernel_size=3),\n                    T.RandomRotation(15),\n                    T.RandomAffine(degrees=0, translate=(0.1, 0.1)),\n                ]\n            )\n\n            aug_tensor = transforms(img_tensor)\n            augmented.append(aug_tensor.squeeze(0).numpy())\n\n        return augmented\n\n    def load_dicom(\n        self, study_id: str, series_id: str, instance_number: int\n    ) -> np.ndarray:\n        \"\"\"Load and preprocess DICOM image\"\"\"\n\n        dicom_path = (\n            self.base_path\n            / \"train_images\"\n            / str(study_id)\n            / str(series_id)\n            / f\"{instance_number}.dcm\"\n        )\n\n        try:\n\n            dcm = pydicom.dcmread(str(dicom_path))\n\n            image = dcm.pixel_array\n\n            # Convert to float and normalize\n\n            image = image.astype(float)\n\n            image = ((image - image.min()) / (image.max() - image.min()) * 255).astype(\n                np.uint8\n            )\n\n            # Convert to grayscale if needed\n\n            if len(image.shape) > 2:\n\n                image = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)\n\n            return image\n\n        except Exception as e:\n\n            print(f\"Error loading DICOM {dicom_path}: {e}\")\n\n            return None\n\n    def save_processed_data(self):\n        \"\"\"Save processed data to disk with type conversion\"\"\"\n\n        def convert_numpy_types(obj):\n            if isinstance(obj, np.integer):\n                return int(obj)\n            elif isinstance(obj, np.floating):\n                return float(obj)\n            elif isinstance(obj, np.ndarray):\n                return obj.tolist()\n            elif isinstance(obj, dict):\n                return {str(k): convert_numpy_types(v) for k, v in obj.items()}\n            elif isinstance(obj, list):\n                return [convert_numpy_types(i) for i in obj]\n            return obj\n\n        # Convert and save metadata\n        converted_samples = convert_numpy_types(self.processed_data[\"samples\"])\n        self.output_path.mkdir(parents=True, exist_ok=True)\n\n        with open(self.output_path / \"metadata.json\", \"w\") as f:\n            json.dump(converted_samples, f)\n\n        # Save ROIs\n        # Convert instance numbers to strings and ensure all numpy arrays are properly handled\n        processed_rois = {}\n        for study_id, study_data in self.processed_data[\"processed_rois\"].items():\n            processed_rois[str(study_id)] = {}\n            for instance_num, instance_data in study_data.items():\n                processed_rois[str(study_id)][str(instance_num)] = {\n                    \"rois\": [\n                        {\n                            \"series_id\": roi[\"series_id\"],\n                            \"condition\": roi[\"condition\"],\n                            \"level\": roi[\"level\"],\n                            \"image\": roi[\"image\"],  # Keep as numpy array\n                            \"original_coords\": convert_numpy_types(\n                                roi[\"original_coords\"]\n                            ),\n                            \"scaled_coords\": convert_numpy_types(roi[\"scaled_coords\"]),\n                        }\n                        for roi in instance_data[\"rois\"]\n                    ]\n                }\n\n        # Save ROIs using numpy's save function\n        np.save(self.output_path / \"processed_rois.npy\", processed_rois)\n\n    def process_study(self, study_id: int):\n        \"\"\"Process a single study with all its series\"\"\"\n        try:\n            study_id_str = str(study_id)\n            # Pre-allocate dictionaries with expected structure\n            processed_samples = {\"series_info\": [], \"conditions\": {}}\n            processed_rois = {}\n\n            # Filter dataframes once using boolean indexing instead of multiple times\n            study_series = self.series_df[self.series_df[\"study_id\"] == study_id]\n            study_coords = self.coords_df[self.coords_df[\"study_id\"] == study_id]\n\n            # Batch process series info\n            processed_samples[\"series_info\"].extend(\n                [\n                    {\n                        \"series_id\": series[\"series_id\"],\n                        \"series_description\": series[\"series_description\"],\n                    }\n                    for _, series in study_series.iterrows()\n                ]\n            )\n\n            # Create temporary data structure for instance grouping\n            instance_data = {}\n\n            # Process coordinates in batch\n            for _, coord_row in study_coords.iterrows():\n                instance_num = coord_row[\"instance_number\"]\n                series_id = coord_row[\"series_id\"]\n                condition = coord_row[\"condition\"]\n                level = coord_row[\"level\"].lower().replace(\"/\", \"_\")\n                x, y = coord_row[\"x\"], coord_row[\"y\"]\n\n                # Get label once per condition/level pair\n                label = self.get_label(study_id, condition, level)\n\n                # Initialize nested dictionaries if needed using dict.setdefault\n                level_dict = processed_samples[\"conditions\"].setdefault(level, {})\n                condition_list = level_dict.setdefault(condition, [])\n\n                # Append coordinate information\n                coord_info = {\n                    \"series_id\": series_id,\n                    \"instance_number\": instance_num,\n                    \"x\": x,\n                    \"y\": y,\n                    \"label\": label,\n                }\n                condition_list.append(coord_info)\n\n                # Group by instance number for ROI processing\n                instance_key = (instance_num, series_id)\n                if instance_key not in instance_data:\n                    instance_data[instance_key] = {\n                        \"coords\": [],\n                        \"conditions\": set(),\n                        \"levels\": set(),\n                        \"label\": label,\n                    }\n                instance_data[instance_key][\"coords\"].append((x, y))\n                instance_data[instance_key][\"conditions\"].add(condition)\n                instance_data[instance_key][\"levels\"].add(level)\n\n            # Process ROIs in batch\n            for (instance_num, series_id), data in instance_data.items():\n                # Load DICOM image once per instance\n                image = self.load_dicom(study_id, series_id, instance_num)\n                if image is None:\n                    continue\n\n                # Extract ROI once per instance\n                roi, scaled_coords = self.extract_roi(image, data[\"coords\"])\n\n                # Initialize ROI storage using dict.setdefault\n                instance_rois = processed_rois.setdefault(instance_num, {\"rois\": []})\n\n                # Create base ROI data\n                for condition in data[\"conditions\"]:\n                    for level in data[\"levels\"]:\n                        roi_data = {\n                            \"series_id\": series_id,\n                            \"condition\": condition,\n                            \"level\": level,\n                            \"image\": roi,\n                            \"original_coords\": data[\"coords\"],\n                            \"scaled_coords\": scaled_coords,\n                        }\n                        instance_rois[\"rois\"].append(roi_data)\n\n                        # Apply augmentation if needed\n                        if data[\"label\"] in [\"Moderate\", \"Severe\"]:\n                            augmented_rois = self.augment_image(roi)\n                            for aug_roi in augmented_rois:\n                                aug_data = roi_data.copy()\n                                aug_data[\"image\"] = aug_roi\n                                instance_rois[\"rois\"].append(aug_data)\n\n            # Update final processed data\n            self.processed_data[\"samples\"][study_id_str] = processed_samples\n            self.processed_data[\"processed_rois\"][study_id_str] = processed_rois\n\n            return True\n\n        except Exception as e:\n            print(f\"Error processing study {study_id}: {e}\")\n            return False\n\n    def get_label(self, study_id: str, condition: str, level: str) -> str:\n        \"\"\"Get label (Normal/Mild, Moderate, Severe) for a specific condition and level\"\"\"\n        condition_col = f\"{condition.replace(' ', '_').lower()}_{level}\"\n        study_row = self.train_df[self.train_df[\"study_id\"] == study_id]\n        if not study_row.empty and condition_col in study_row.columns:\n\n            try:\n                return study_row[condition_col].iloc[0]\n            except:\n                pass\n\n        return \"Normal/Mild\"  # Default case","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T20:50:44.945173Z","iopub.execute_input":"2025-01-11T20:50:44.945400Z","iopub.status.idle":"2025-01-11T20:50:44.969484Z","shell.execute_reply.started":"2025-01-11T20:50:44.945381Z","shell.execute_reply":"2025-01-11T20:50:44.968624Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"processor = SpineDataProcessor(\n    base_path=base_path,\n    output_path=output_path\n)\n\n# Get unique study IDs\nstudy_ids = processor.train_df['study_id'].unique()\n# Process each study with progress bar\nfor study_id in tqdm(study_ids, desc=\"Processing studies\"):\n    processor.process_study(study_id)\n\nprocessor.save_processed_data()\ndel processor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T20:50:44.972738Z","iopub.execute_input":"2025-01-11T20:50:44.972950Z","iopub.status.idle":"2025-01-11T21:00:37.312345Z","shell.execute_reply.started":"2025-01-11T20:50:44.972931Z","shell.execute_reply":"2025-01-11T21:00:37.311381Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DataVisualizer:\n    def __init__(self, base_path, processed_path):\n        self.base_path = Path(base_path)\n        self.processed_path = Path(processed_path)\n        \n        print(\"Loading metadata and processed ROIs...\")\n        # Load processed data\n        with open(self.processed_path / \"metadata.json\", \"r\") as f:\n            self.metadata = json.load(f)\n        self.processed_rois = np.load(self.processed_path / \"processed_rois.npy\", allow_pickle=True).item()\n        \n        print(f\"Found {len(self.metadata)} studies in metadata\")\n        print(f\"Found {len(self.processed_rois)} studies in processed ROIs\")\n\n    def visualize_study(self, study_id: str):\n        \"\"\"Simple visualization of original and processed images with coordinates.\"\"\"\n        metadata = self.metadata[study_id]\n        processed_rois = self.processed_rois[study_id]\n        \n        for instance_num, roi_data in processed_rois.items():\n            for roi_info in roi_data['rois']:\n                # Get series ID from ROI info\n                roi_series_id = roi_info['series_id']\n                \n                # Load original image with matching series ID\n                original_img = self.load_original_dicom(study_id, roi_series_id, instance_num)\n                if original_img is None:\n                    continue\n                \n                # Create figure with subplots\n                fig, (ax1, ax2) = plt.subplots(1, 2)\n                \n                # Plot original\n                ax1.imshow(original_img, cmap='gray')\n                ax1.set_title('Original')\n                if roi_info['original_coords']:\n                    coords = np.array(roi_info['original_coords'])\n                    ax1.scatter(coords[:, 0], coords[:, 1], c='red', marker='x', s=100)\n                ax1.axis('off')\n                \n                # Plot processed\n                ax2.imshow(roi_info['image'], cmap='gray')\n                ax2.set_title('Processed ROI')\n                if roi_info['scaled_coords']:\n                    coords = np.array(roi_info['scaled_coords'])\n                    ax2.scatter(coords[:, 0], coords[:, 1], c='red', marker='x', s=100)\n                ax2.axis('off')\n                \n                plt.show()\n            \n    def load_original_dicom(self, study_id, series_id, instance_number):\n        \"\"\"Load original DICOM image with error handling\"\"\"\n        try:\n            dicom_path = self.base_path / 'train_images' / str(study_id) / str(series_id) / f\"{instance_number}.dcm\"\n            if not dicom_path.exists():\n                return None\n                \n            print(f\"Loading DICOM from: {dicom_path}\")\n            dcm = pydicom.dcmread(dicom_path)\n            image = dcm.pixel_array\n            \n            # Normalize to 0-255\n            image = ((image - image.min()) / (image.max() - image.min()) * 255).astype(np.uint8)\n            if len(image.shape) > 2:\n                image = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)\n            return image\n            \n        except Exception as e:\n            print(f\"Error loading DICOM: {e}\")\n            return None\n\n    def get_condition_counts(self, num_samples=3):\n        \"\"\"Count occurrences of each condition-level pair and their severity\"\"\"\n        counts = defaultdict(lambda: defaultdict(lambda: defaultdict(int)))\n        \n        for study_data in self.metadata.values():\n            for level, conditions in study_data['conditions'].items():\n                for condition, instances in conditions.items():\n                    for instance in instances:\n                        counts[condition][level][instance['label']] += 1\n        \n        # Print counts in a structured way\n        for condition in counts:\n            print(f\"\\n{condition}:\")\n            for level in counts[condition]:\n                print(f\"  {level}:\")\n                for severity, count in counts[condition][level].items():\n                    print(f\"    {severity}: {count}\")\n\n        \"\"\"Visualize original and processed images for given severity\"\"\"\n        plt.figure(figsize=(15, 5*num_samples))\n        sample_count = 0\n        \n        for study_id, study_data in self.processed_rois.items():\n            for instance_num, instance_data in study_data.items():\n                for roi_data in instance_data['rois']:\n                    if roi_data.get('label') == severity and sample_count < num_samples:\n                        # Get original image\n                        original_img = self.load_original_dicom(\n                            study_id, \n                            roi_data['series_id'],\n                            roi_data['instance_number']\n                        )\n                        \n                        # Plot original image\n                        plt.subplot(num_samples, 2, sample_count*2 + 1)\n                        plt.imshow(original_img, cmap='gray')\n                        plt.plot(roi_data['original_coords']['x'], \n                               roi_data['original_coords']['y'], \n                               'r.', markersize=10)\n                        circle = plt.Circle((roi_data['original_coords']['x'], \n                                          roi_data['original_coords']['y']), \n                                         radius=112,  # half of 224 for ROI visualization\n                                         color='r', \n                                         fill=False)\n                        plt.gca().add_patch(circle)\n                        plt.title(f\"Original - {roi_data['condition']} {roi_data['level']}\")\n                        \n                        # Plot processed ROI\n                        plt.subplot(num_samples, 2, sample_count*2 + 2)\n                        plt.imshow(roi_data['image'], cmap='gray')\n                        plt.plot(roi_data['scaled_coords']['x'], \n                               roi_data['scaled_coords']['y'], \n                               'r.', markersize=10)\n                        plt.title(f\"Processed ROI - {severity}\")\n                        \n                        sample_count += 1\n                        if sample_count >= num_samples:\n                            break\n            if sample_count >= num_samples:\n                break\n                \n        plt.tight_layout()\n        plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:00:37.313480Z","iopub.execute_input":"2025-01-11T21:00:37.313780Z","iopub.status.idle":"2025-01-11T21:00:37.329616Z","shell.execute_reply.started":"2025-01-11T21:00:37.313747Z","shell.execute_reply":"2025-01-11T21:00:37.328596Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualizer = DataVisualizer(\n    base_path=base_path,\n    processed_path=output_path\n)\nprint(\"Condition-Level Pair Counts:\")\nvisualizer.get_condition_counts()\nvisualizer.visualize_study(\"4003253\")\ndel visualizer\n# print(visualizer.processed_rois['4003253'].keys())\n# print(visualizer.processed_rois['4003253']['8'].keys())\n# print(visualizer.processed_rois['4003253']['8']['rois'][0].keys())\n# print(visualizer.processed_rois['4003253']['8']['rois'][0].values())\n# print(visualizer.processed_rois['4003253'].values())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:00:37.332013Z","iopub.execute_input":"2025-01-11T21:00:37.332321Z","iopub.status.idle":"2025-01-11T21:00:50.306414Z","shell.execute_reply.started":"2025-01-11T21:00:37.332300Z","shell.execute_reply":"2025-01-11T21:00:50.305488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SpineDataset(Dataset):\n    def __init__(self, processed_rois, metadata, study_ids):\n        self.processed_rois = processed_rois\n        self.metadata = metadata\n        self.study_ids = study_ids\n        self.series_type_mapping = {\n            \"Sagittal T1\": 0,\n            \"Axial T2\": 1,\n            \"Sagittal T2/STIR\": 2,\n        }\n\n        # Define all possible conditions and levels\n        self.conditions = [\n            \"Spinal Canal Stenosis\",\n            \"Left Neural Foraminal Narrowing\",\n            \"Right Neural Foraminal Narrowing\",\n            \"Left Subarticular Stenosis\",\n            \"Right Subarticular Stenosis\",\n        ]\n        self.levels = [\"l1_l2\", \"l2_l3\", \"l3_l4\", \"l4_l5\", \"l5_s1\"]\n\n        # Filter out studies with no ROIs and preprocess\n        self.valid_study_ids = []\n        self.cached_data = {}\n        self._preprocess_data()\n\n    def __len__(self):\n        return len(self.valid_study_ids)\n\n    def get_series_type_encoding(self, series_id, study_id):\n        series_desc = next(\n            s[\"series_description\"]\n            for s in self.metadata[study_id][\"series_info\"]\n            if s[\"series_id\"] == series_id\n        )\n        return self.series_type_mapping[series_desc]\n\n    def _preprocess_data(self):\n        \"\"\"Preprocess and cache all data during initialization\"\"\"\n        print(\"Preprocessing and caching dataset...\")\n\n        for study_id in tqdm(self.study_ids):\n            try:\n                study_rois = self.processed_rois[study_id]\n                study_metadata = self.metadata[study_id]\n\n                # Process images and series types\n                all_images = []\n                all_series_types = []\n                all_conditions = []\n\n                # Check if study has any ROIs\n                has_rois = False\n                for instance_num, instance_data in study_rois.items():\n                    if \"rois\" in instance_data and instance_data[\"rois\"]:\n                        has_rois = True\n                        for roi in instance_data[\"rois\"]:\n                            # Normalize and convert to tensor (keeping in CPU)\n                            image = torch.FloatTensor(roi[\"image\"]) / 255.0\n                            series_type = self.get_series_type_encoding(\n                                roi[\"series_id\"], study_id\n                            )\n\n                            all_images.append(image)\n                            all_series_types.append(series_type)\n                            all_conditions.append((roi[\"condition\"], roi[\"level\"]))\n\n                # Skip studies with no ROIs\n                if not has_rois or not all_images:\n                    print(f\"Warning: Study {study_id} has no ROIs, skipping...\")\n                    continue\n\n                # Stack images and convert series types to tensor\n                images = torch.stack(all_images)\n                series_types = torch.tensor(all_series_types)\n\n                # Create and fill labels tensor\n                labels = torch.zeros(25, 3)\n\n                for i, condition in enumerate(self.conditions):\n                    for j, level in enumerate(self.levels):\n                        idx = i * 5 + j\n                        if level in study_metadata[\"conditions\"]:\n                            if condition in study_metadata[\"conditions\"][level]:\n                                label = study_metadata[\"conditions\"][level][condition][\n                                    0\n                                ][\"label\"]\n                                if label == \"Normal/Mild\":\n                                    labels[idx] = torch.tensor([1, 0, 0])\n                                elif label == \"Moderate\":\n                                    labels[idx] = torch.tensor([0, 1, 0])\n                                elif label == \"Severe\":\n                                    labels[idx] = torch.tensor([0, 0, 1])\n\n                # Cache processed data\n                self.cached_data[study_id] = {\n                    \"images\": images,\n                    \"series_types\": series_types,\n                    \"labels\": labels,\n                }\n                self.valid_study_ids.append(study_id)\n\n            except Exception as e:\n                print(f\"Error processing study {study_id}: {str(e)}\")\n                continue\n\n        print(\n            f\"Dataset preprocessing completed! Valid studies: {len(self.valid_study_ids)}\"\n        )\n\n    def __getitem__(self, idx):\n        \"\"\"Get preprocessed data from cache\"\"\"\n        study_id = self.valid_study_ids[idx]\n        cached_item = self.cached_data[study_id]\n\n        return {\n            \"images\": cached_item[\"images\"],\n            \"series_types\": cached_item[\"series_types\"],\n            \"labels\": cached_item[\"labels\"],\n            \"study_id\": study_id,\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:00:50.307962Z","iopub.execute_input":"2025-01-11T21:00:50.308229Z","iopub.status.idle":"2025-01-11T21:00:50.320360Z","shell.execute_reply.started":"2025-01-11T21:00:50.308207Z","shell.execute_reply":"2025-01-11T21:00:50.319566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiHeadAttention(nn.Module):\n    def __init__(self, input_dim):\n        super().__init__()\n        self.attention = nn.Sequential(\n            nn.Linear(input_dim, input_dim),\n            nn.Tanh(),\n            nn.Linear(input_dim, 1)\n        )\n    \n    def forward(self, x):\n        # x: [batch_size, input_dim]\n        weights = self.attention(x)  # [batch_size, 1]\n        weights = torch.softmax(weights, dim=1)\n        return x * weights  # [batch_size, input_dim]\n\nclass SpineModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n        # Initialize EfficientNet-B0 without pretrained weights\n        self.efficient_net = efficientnet_b0(weights=None)\n        \n        # Modify first conv layer for single channel\n        self.efficient_net.features[0][0] = nn.Conv2d(\n            1, 32, kernel_size=3, stride=2, padding=1, bias=False\n        )\n        \n        # Get feature dimensions (1280 for EfficientNet-B0)\n        self.feature_dim = 1280\n        \n        # Series type embedding\n        self.series_embedding = nn.Embedding(3, 32)\n        \n        # Classification head with attention\n        self.attention1 = MultiHeadAttention(self.feature_dim + 32)\n        self.fc1 = nn.Linear(self.feature_dim + 32, 256)\n        self.attention2 = MultiHeadAttention(256)\n        self.dropout1 = nn.Dropout(0.3)\n        self.fc2 = nn.Linear(256, 25 * 3)\n        \n        # Initialize weights\n        self._initialize_weights()\n\n    def _initialize_weights(self):\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                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.Linear):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n\n    def forward(self, images, series_types):\n        batch_size = images.size(0)\n        num_rois = images.size(1)\n        \n        # Process each image through EfficientNet\n        images = images.view(-1, 1, images.size(-2), images.size(-1))\n        features = self.efficient_net.features(images)\n        features = self.efficient_net.avgpool(features)\n        features = torch.flatten(features, 1)  # [batch_size * num_rois, feature_dim]\n        \n        # Get series type embeddings\n        series_embeddings = self.series_embedding(series_types.view(-1))  # [batch_size * num_rois, 32]\n        \n        # Concatenate features and embeddings\n        combined = torch.cat([features, series_embeddings], dim=1)  # [batch_size * num_rois, feature_dim + 32]\n        \n        # Apply attention and classification\n        x = self.attention1(combined)\n        x = self.fc1(x)\n        x = torch.relu(x)  # Use ReLU as in original\n        x = self.attention2(x)\n        x = self.dropout1(x)\n        predictions = self.fc2(x)\n        \n        # Reshape to [batch_size, num_rois, 25 * 3]\n        predictions = predictions.view(batch_size, num_rois, -1)\n        \n        # Create a mask for valid ROIs (non-zero images)\n        roi_mask = (images.view(batch_size, num_rois, -1).sum(dim=-1) != 0).float()\n        roi_mask = roi_mask.unsqueeze(-1)  # Add dimension for broadcasting\n        \n        # Apply mask and average predictions across valid ROIs only\n        masked_predictions = predictions * roi_mask\n        predictions = masked_predictions.sum(dim=1) / (roi_mask.sum(dim=1) + 1e-6)\n        \n        # Reshape to [batch_size, 25, 3] and apply softmax\n        predictions = predictions.view(batch_size, 25, 3)\n        predictions = torch.softmax(predictions, dim=-1)\n        \n        return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:00:50.321275Z","iopub.execute_input":"2025-01-11T21:00:50.321565Z","iopub.status.idle":"2025-01-11T21:00:50.336380Z","shell.execute_reply.started":"2025-01-11T21:00:50.321543Z","shell.execute_reply":"2025-01-11T21:00:50.335493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SpineTrainer:\n    def __init__(self, processed_rois, metadata, device=\"cuda\"):\n        self.processed_rois = processed_rois\n        self.metadata = metadata\n        self.device = device\n\n        # Create model\n        self.model = SpineModel().to(device)\n\n        # Calculate class weights\n        self.class_weights = self.calculate_class_weights()\n        print(f\"Class weights device: {self.class_weights.device}\")\n        # Loss function and optimizer\n        self.criterion = nn.CrossEntropyLoss(weight=self.class_weights)\n        self.optimizer = optim.AdamW(self.model.parameters(), lr=1e-4)\n        self.scaler = torch.amp.GradScaler()\n\n    def calculate_class_weights(self):\n        # Count occurrences of each class\n        class_counts = defaultdict(int)\n        total_samples = 0\n\n        for study_data in self.metadata.values():\n            for level_data in study_data[\"conditions\"].values():\n                for condition_data in level_data.values():\n                    label = condition_data[0][\"label\"]\n                    class_counts[label] += 1\n                    total_samples += 1\n\n        # Calculate weights\n        weights = torch.zeros(3)\n        weights[0] = total_samples / (3 * class_counts[\"Normal/Mild\"])\n        weights[1] = total_samples / (3 * class_counts[\"Moderate\"])\n        weights[2] = total_samples / (3 * class_counts[\"Severe\"])\n\n        return weights.to(self.device)\n\n    def train_epoch(self, train_loader):\n        print(\"Train epoch started\")\n        self.model.train()\n        print(\"Model is training...\")\n        total_loss = 0\n        print(f\"\\nStarting to iterate through {len(train_loader)} batches...\")\n        for batch in tqdm(\n            train_loader, total=len(train_loader), desc=\"iterate through train batches\"\n        ):\n            try:\n                images = batch[\"images\"].to(self.device)\n                series_types = batch[\"series_types\"].to(self.device)\n                labels = batch[\"labels\"].to(self.device)\n\n                self.optimizer.zero_grad()\n\n                with torch.amp.autocast(device.type):\n                    outputs = self.model(images, series_types)\n                    # Calculate loss for each condition\n                    loss = 0\n                    for i in range(25):\n                        loss += self.criterion(outputs[:, i], labels[:, i])\n                    loss = loss / 25\n                self.scaler.scale(loss).backward()\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n\n                total_loss += loss.item()\n            except Exception as err:\n                print(err)\n\n        print(f\"\\nEpoch completed. Average loss: {total_loss / len(train_loader)}\")\n        return total_loss / len(train_loader)\n\n    def validate(self, val_loader):\n        print(\"\\nStarting validation...\")\n        self.model.eval()\n        total_loss = 0\n\n        with torch.no_grad():\n            for batch in tqdm(\n                val_loader,\n                total=len(val_loader),\n                desc=\"iterate through validation batches\",\n            ):\n                images = batch[\"images\"].to(self.device)\n                series_types = batch[\"series_types\"].to(self.device)\n                labels = batch[\"labels\"].to(self.device)\n\n                outputs = self.model(images, series_types)\n\n                loss = 0\n                for i in range(25):\n                    loss += self.criterion(outputs[:, i], labels[:, i])\n                loss = loss / 25\n\n                total_loss += loss.item()\n\n        return total_loss / len(val_loader)\n\n    def train(self, train_loader, val_loader, num_epochs=10):\n        best_val_loss = float(\"inf\")\n\n        for epoch in range(num_epochs):\n            train_loss = self.train_epoch(train_loader)\n            val_loss = self.validate(val_loader)\n\n            print(f\"Epoch {epoch+1}/{num_epochs}\")\n            print(f\"Train Loss: {train_loss:.4f}\")\n            print(f\"Val Loss: {val_loss:.4f}\")\n\n            if val_loss < best_val_loss:\n                best_val_loss = val_loss\n                torch.save(self.model.state_dict(), \"best_model.pth\")\n\n    def predict(self, test_loader):\n        self.model.eval()\n        predictions = {}\n\n        with torch.no_grad():\n            for batch in tqdm(test_loader, total=len(test_loader), desc=\"Prediction\"):\n                images = batch[\"images\"].to(self.device)\n                series_types = batch[\"series_types\"].to(self.device)\n                study_ids = batch[\"study_id\"]\n\n                outputs = self.model(images, series_types)\n\n                # Store predictions\n                for i, study_id in enumerate(study_ids):\n                    predictions[study_id] = outputs[i].cpu().numpy()\n\n        return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:00:50.337220Z","iopub.execute_input":"2025-01-11T21:00:50.337506Z","iopub.status.idle":"2025-01-11T21:00:50.352387Z","shell.execute_reply.started":"2025-01-11T21:00:50.337472Z","shell.execute_reply":"2025-01-11T21:00:50.351589Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_data(processed_rois_path, metadata_path):\n    \"\"\"Load the processed ROIs and metadata\"\"\"\n    # Load processed ROIs\n    processed_rois = np.load(processed_rois_path, allow_pickle=True).item()\n\n    # Load metadata\n    with open(metadata_path, \"r\") as f:\n        metadata = json.load(f)\n\n    return processed_rois, metadata\n\ndef custom_collate(batch):\n    \"\"\"Custom collate function to handle variable number of ROIs per study\"\"\"\n    # Get max number of ROIs in this batch\n    max_rois = max([b[\"images\"].size(0) for b in batch])\n\n    # Get other dimensions from first item\n    first = batch[0]\n    img_h, img_w = first[\"images\"].size(-2), first[\"images\"].size(-1)\n\n    # Initialize tensors for the batch\n    batch_size = len(batch)\n    batched_images = torch.zeros(batch_size, max_rois, img_h, img_w)\n    batched_series_types = torch.zeros(batch_size, max_rois, dtype=torch.long)\n    batched_labels = torch.stack([b[\"labels\"] for b in batch])\n    study_ids = [b[\"study_id\"] for b in batch]\n\n    # Fill in the batched tensors\n    for i, item in enumerate(batch):\n        num_rois = item[\"images\"].size(0)\n        batched_images[i, :num_rois] = item[\"images\"]\n        batched_series_types[i, :num_rois] = item[\"series_types\"]\n\n    return {\n        \"images\": batched_images,\n        \"series_types\": batched_series_types,\n        \"labels\": batched_labels,\n        \"study_id\": study_ids,\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:00:50.353220Z","iopub.execute_input":"2025-01-11T21:00:50.353457Z","iopub.status.idle":"2025-01-11T21:00:50.368933Z","shell.execute_reply.started":"2025-01-11T21:00:50.353438Z","shell.execute_reply":"2025-01-11T21:00:50.368014Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestDataset(torch.utils.data.Dataset):\n    def __init__(self, processed_rois, study_series, target_size=(256, 256)):\n        self.processed_rois = processed_rois\n        self.study_series = {s[\"series_id\"]: s for s in study_series}\n        self.target_size = target_size\n        self.series_type_mapping = {\n            \"Sagittal T1\": 0,\n            \"Axial T2\": 1,\n            \"Sagittal T2/STIR\": 2,\n        }\n\n        # Prepare instance list\n        self.instances = []\n        for instance_num, instance_data in processed_rois.items():\n            self.instances.append((instance_num, instance_data))\n\n    def get_series_type(self, series_id):\n        series_info = self.study_series.get(series_id)\n        if series_info:\n            return self.series_type_mapping[series_info[\"series_description\"]]\n        return 0  # Default value if not found\n\n    def resize_image(self, image):\n        \"\"\"Resize image to target size\"\"\"\n        h, w = image.shape\n        if (h, w) != self.target_size:\n            # Convert to PIL Image for resizing\n            from PIL import Image\n            import numpy as np\n\n            img_pil = Image.fromarray(image)\n            img_pil = img_pil.resize(self.target_size, Image.Resampling.BILINEAR)\n            return np.array(img_pil)\n        return image\n\n    def __getitem__(self, idx):\n        instance_num, instance_data = self.instances[idx]\n\n        # Process all ROIs in this instance\n        all_images = []\n        all_series_types = []\n\n        for roi in instance_data[\"rois\"]:\n            # Resize image before converting to tensor\n            resized_image = self.resize_image(roi[\"image\"])\n            image = torch.FloatTensor(resized_image) / 255.0\n            series_type = self.get_series_type(roi[\"series_id\"])\n\n            all_images.append(image)\n            all_series_types.append(series_type)\n\n        try:\n            # Stack tensors\n            images = torch.stack(all_images)\n            series_types = torch.tensor(all_series_types)\n\n            return {\n                \"images\": images,\n                \"series_types\": series_types,\n                \"instance_num\": instance_num,\n            }\n        except Exception as e:\n            print(f\"Error stacking tensors for instance {instance_num}:\")\n            print(f\"Image shapes: {[img.shape for img in all_images]}\")\n            raise e\n\n    def __len__(self):\n        return len(self.instances)\n\n\ndef load_test_metadata(csv_path):\n    \"\"\"Load and structure test series descriptions\"\"\"\n    df = pd.read_csv(csv_path)\n\n    # Group by study_id\n    study_metadata = defaultdict(list)\n    for _, row in df.iterrows():\n        study_metadata[str(row[\"study_id\"])].append(\n            {\n                \"series_id\": str(row[\"series_id\"]),\n                \"series_description\": row[\"series_description\"],\n            }\n        )\n\n    return study_metadata\n\ndef process_test_image(dcm_path):\n    \"\"\"Process a single DICOM image\"\"\"\n    try:\n        # Read DICOM\n        dcm = pydicom.dcmread(dcm_path)\n        image = dcm.pixel_array.astype(float)\n\n        # Normalize image\n        if image.max() != image.min():\n            image = ((image - image.min()) / (image.max() - image.min()) * 255).astype(\n                np.uint8\n            )\n        else:\n            image = np.zeros_like(image, dtype=np.uint8)\n\n        # print(f\"Processed image shape: {image.shape}\")  # Debug info\n        return image\n\n    except Exception as e:\n        print(f\"Error processing DICOM {dcm_path}: {str(e)}\")\n        raise e\n    \ndef process_test_study(test_path, study_id, study_series):\n    \"\"\"Process all images for a study\"\"\"\n    processed_rois = {}\n\n    # Process each series\n    for series_info in study_series:\n        series_id = series_info[\"series_id\"]\n        series_path = test_path / str(study_id) / str(series_id)\n\n        if not series_path.exists():\n            print(f\"Warning: Series path {series_path} does not exist\")\n            continue\n\n        # Process each DICOM in series\n        for dcm_path in series_path.glob(\"*.dcm\"):\n            instance_num = int(dcm_path.stem)\n            try:\n                processed_image = process_test_image(dcm_path)\n\n                if instance_num not in processed_rois:\n                    processed_rois[instance_num] = {\"rois\": []}\n\n                processed_rois[instance_num][\"rois\"].append(\n                    {\"series_id\": series_id, \"image\": processed_image}\n                )\n\n            except Exception as e:\n                print(f\"Error processing {dcm_path}: {str(e)}\")\n                continue\n\n    return processed_rois","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:00:50.369948Z","iopub.execute_input":"2025-01-11T21:00:50.370247Z","iopub.status.idle":"2025-01-11T21:00:50.384440Z","shell.execute_reply.started":"2025-01-11T21:00:50.370220Z","shell.execute_reply":"2025-01-11T21:00:50.383535Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_study(model, test_loader, device):\n    \"\"\"Generate predictions for a study\"\"\"\n    model.eval()  # Set model to evaluation mode\n    predictions = []\n    \n    with torch.no_grad():\n        for batch in test_loader:\n            images = batch['images'].to(device)\n            series_types = batch['series_types'].to(device)\n            \n            try:\n                outputs = model(images, series_types)\n                # Ensure outputs are the right shape (batch_size, 25, 3)\n                if len(outputs.shape) != 3 or outputs.shape[1:] != (25, 3):\n                    print(f\"WARNING: Unexpected output shape: {outputs.shape}\")\n                    continue\n                    \n                predictions.append(outputs.cpu().numpy())\n                \n            except Exception as e:\n                print(f\"Error during prediction: {str(e)}\")\n                print(f\"Images shape: {images.shape}\")\n                print(f\"Series types shape: {series_types.shape}\")\n                continue\n    \n    if not predictions:\n        raise RuntimeError(\"No valid predictions were generated for this study\")\n    \n    # Stack all predictions\n    predictions = np.concatenate(predictions, axis=0)\n    \n    # Average predictions across all instances\n    final_prediction = np.mean(predictions, axis=0)\n    \n    # Ensure final prediction has shape (25, 3)\n    if final_prediction.shape != (25, 3):\n        print(f\"WARNING: Final prediction has unexpected shape: {final_prediction.shape}\")\n        if len(final_prediction.shape) == 3 and final_prediction.shape[0] == 1:\n            final_prediction = final_prediction.squeeze(0)\n    \n    # Ensure probabilities sum to 1\n    final_prediction = final_prediction / final_prediction.sum(axis=1, keepdims=True)\n    \n    return final_prediction\n\ndef process_and_predict(test_path, model, series_csv_path, device):\n    \"\"\"Main function to process test data and generate predictions\"\"\"\n    study_metadata = load_test_metadata(series_csv_path)\n\n    test_path = Path(test_path)\n    predictions = {}\n\n    # Process each study\n    for study_id in study_metadata.keys():\n        print(f\"Processing study {study_id}\")\n\n        study_series = study_metadata[study_id]\n\n        # Process study images\n        processed_rois = process_test_study(test_path, study_id, study_series)\n\n        if not processed_rois:\n            print(f\"Warning: No ROIs processed for study {study_id}\")\n            continue\n\n        # Create dataset and dataloader\n        test_dataset = TestDataset(processed_rois, study_series)\n        test_loader = torch.utils.data.DataLoader(\n            test_dataset, batch_size=1, shuffle=False, num_workers=2\n        )\n\n        # Generate predictions\n        study_predictions = predict_study(model, test_loader, device)\n        predictions[study_id] = study_predictions\n\n    return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:00:50.385328Z","iopub.execute_input":"2025-01-11T21:00:50.385541Z","iopub.status.idle":"2025-01-11T21:00:50.400020Z","shell.execute_reply.started":"2025-01-11T21:00:50.385523Z","shell.execute_reply":"2025-01-11T21:00:50.399277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_submission(predictions, output_path):\n    \"\"\"Create submission file from model predictions\"\"\"\n    rows = []\n    conditions = [\n        \"spinal_canal_stenosis\",\n        \"left_neural_foraminal_narrowing\",\n        \"right_neural_foraminal_narrowing\",\n        \"left_subarticular_stenosis\",\n        \"right_subarticular_stenosis\",\n    ]\n    levels = [\"l1_l2\", \"l2_l3\", \"l3_l4\", \"l4_l5\", \"l5_s1\"]\n\n    for study_id, study_preds in predictions.items():\n\n        # Reshape predictions if necessary\n        if len(study_preds.shape) == 3:  # If shape is (1, 25, 3)\n            study_preds = study_preds.squeeze(0)  # Convert to (25, 3)\n\n        if len(study_preds.shape) != 2 or study_preds.shape != (25, 3):\n            print(\n                f\"WARNING: Unexpected prediction shape for study {study_id}: {study_preds.shape}\"\n            )\n            continue\n\n        for condition_idx, condition in enumerate(conditions):\n            for level_idx, level in enumerate(levels):\n                pred_idx = condition_idx * len(levels) + level_idx\n                probs = study_preds[pred_idx]\n\n                row_id = f\"{study_id}_{condition}_{level}\"\n                rows.append(\n                    {\n                        \"row_id\": row_id,\n                        \"normal_mild\": float(probs[0]),\n                        \"moderate\": float(probs[1]),\n                        \"severe\": float(probs[2]),\n                    }\n                )\n\n    if not rows:\n        raise ValueError(\"No predictions were processed successfully!\")\n\n    submission_df = pd.DataFrame(rows)\n\n    # Verify probabilities sum to approximately 1\n    prob_sum = submission_df[[\"normal_mild\", \"moderate\", \"severe\"]].sum(axis=1)\n    if not np.allclose(prob_sum, 1.0, atol=1e-5):\n        print(\"WARNING: Some probability rows don't sum to 1!\")\n        print(\"Probability sums range:\", prob_sum.min(), \"to\", prob_sum.max())\n\n    submission_df.to_csv(output_path, index=False)\n    print(f\"Saved submission to {output_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:00:50.400899Z","iopub.execute_input":"2025-01-11T21:00:50.401234Z","iopub.status.idle":"2025-01-11T21:00:50.415742Z","shell.execute_reply.started":"2025-01-11T21:00:50.401201Z","shell.execute_reply":"2025-01-11T21:00:50.414989Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader\n# Set random seeds for reproducibility\ntorch.manual_seed(42)\nnp.random.seed(42)\ntest_base_path = (\n    base_path / \"test_images\"\n)\n# # Device configuration\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# # Hyperparameters\nBATCH_SIZE = 1\nNUM_EPOCHS = 30\nVAL_SIZE = 0.2\nNUM_WORKERS = 0\n\n# # Load data\nprocessed_rois, metadata = load_data(\n    output_path /\"processed_rois.npy\", output_path / \"metadata.json\"\n)\n\n# Get all study IDs\nstudy_ids = list(metadata.keys())\n\n# Split into train and validation sets\ntrain_ids, val_ids = train_test_split(\n    study_ids, test_size=VAL_SIZE, random_state=42\n)\n\nprint(f\"Number of training studies: {len(train_ids)}\")\nprint(f\"Number of validation studies: {len(val_ids)}\")\n\n# Create datasets\ntrain_dataset = SpineDataset(processed_rois, metadata, train_ids)\nval_dataset = SpineDataset(processed_rois, metadata, val_ids)\n\n# Create dataloaders\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n    pin_memory=True,\n    collate_fn=custom_collate,\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=True,\n    collate_fn=custom_collate,\n)\n\nprint(\"Length of the train data :\", len(train_loader))\nprint(\"Length of the val data :\", len(val_loader))\n\n# # Initialize trainer\ntrainer = SpineTrainer(\n    processed_rois=processed_rois, metadata=metadata, device=device\n)\n\n# # Training loop with checkpointing\nbest_val_loss = float(\"inf\")\nsave_dir = output_path\nsave_dir.mkdir(exist_ok=True)\n\nfor epoch in range(NUM_EPOCHS):\n    print(f\"\\nEpoch {epoch+1}/{NUM_EPOCHS}\")\n    print(\"-\" * 20)\n\n    # Train and validate\n    train_loss = trainer.train_epoch(train_loader)\n    val_loss = trainer.validate(val_loader)\n\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss: {val_loss:.4f}\")\n\n    # Save checkpoint if validation loss improved\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        checkpoint_path = (\n            save_dir / f\"model_epoch_{epoch+1}_valloss_{val_loss:.4f}.pth\"\n        )\n\n        # Save checkpoint\n        torch.save(\n            {\n                \"epoch\": epoch,\n                \"model_state_dict\": trainer.model.state_dict(),\n                \"optimizer_state_dict\": trainer.optimizer.state_dict(),\n                \"val_loss\": val_loss,\n                \"train_loss\": train_loss,\n            },\n            checkpoint_path,\n        )\n\n        print(f\"Saved checkpoint to {checkpoint_path}\")\n\nprint(\"Training completed!\")\n\n# # Load best model for predictions\nbest_model_path = sorted(\n    save_dir.glob(\"*.pth\"),\n    key=lambda x: float(str(x).split(\"valloss_\")[1].split(\".pth\")[0]),\n)[0]\nprint(f\"Loading best model from {best_model_path}\")\n\ncheckpoint = torch.load(best_model_path)\ntrainer.model.load_state_dict(checkpoint[\"model_state_dict\"])\n\n# Generate predictions for validation set\n# val_predictions = trainer.predict(val_loader)\n\ntest_predictions = process_and_predict(\n    test_base_path,\n    trainer.model,\n    base_path /\"test_series_descriptions.csv\",\n    device=device,\n)\n\n# Create submission file for validation set\n# create_submission(val_predictions, \"validation_predictions.csv\")\ncreate_submission(test_predictions, output_path / \"submission.csv\")\n\nprint(\"Created validation predictions file!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:00:50.416644Z","iopub.execute_input":"2025-01-11T21:00:50.416919Z","iopub.status.idle":"2025-01-11T21:37:49.237436Z","shell.execute_reply.started":"2025-01-11T21:00:50.416893Z","shell.execute_reply":"2025-01-11T21:37:49.236301Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(\"../working/submission.csv\")\n\ndisplay(df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T21:55:17.628724Z","iopub.execute_input":"2025-01-11T21:55:17.629041Z","iopub.status.idle":"2025-01-11T21:55:17.644663Z","shell.execute_reply.started":"2025-01-11T21:55:17.629017Z","shell.execute_reply":"2025-01-11T21:55:17.643454Z"}},"outputs":[],"execution_count":null}]}