{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np \nimport pandas as pd \n\nimport cv2\nimport pydicom\nfrom PIL import Image\nfrom IPython.display import Image as IPyImage, display\n\nimport os\nimport re\nimport glob\nimport random\nfrom tqdm import tqdm\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nsns.set(style=\"whitegrid\")\n\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom torch.optim import AdamW\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\nfrom sklearn.model_selection import train_test_split\n\nfrom torchvision import transforms\nimport timm\n\nimport yaml\n\nimport albumentations as A\n\nfrom sklearn.model_selection import KFold","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-11-03T10:06:42.586776Z","iopub.execute_input":"2024-11-03T10:06:42.587167Z","iopub.status.idle":"2024-11-03T10:06:42.597423Z","shell.execute_reply.started":"2024-11-03T10:06:42.587128Z","shell.execute_reply":"2024-11-03T10:06:42.596309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install ultralytics","metadata":{"execution":{"iopub.status.busy":"2024-11-03T05:54:10.401328Z","iopub.execute_input":"2024-11-03T05:54:10.401887Z","iopub.status.idle":"2024-11-03T05:54:24.288239Z","shell.execute_reply.started":"2024-11-03T05:54:10.401851Z","shell.execute_reply":"2024-11-03T05:54:24.286983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from ultralytics import YOLO\nfrom sklearn.model_selection import train_test_split\n","metadata":{"execution":{"iopub.status.busy":"2024-11-03T05:54:24.289648Z","iopub.execute_input":"2024-11-03T05:54:24.289950Z","iopub.status.idle":"2024-11-03T05:54:24.424144Z","shell.execute_reply.started":"2024-11-03T05:54:24.289919Z","shell.execute_reply":"2024-11-03T05:54:24.423254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\nlabel_coords_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv')\nseries_desc_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')","metadata":{"execution":{"iopub.status.busy":"2024-11-03T05:54:24.426036Z","iopub.execute_input":"2024-11-03T05:54:24.426874Z","iopub.status.idle":"2024-11-03T05:54:24.594853Z","shell.execute_reply.started":"2024-11-03T05:54:24.426829Z","shell.execute_reply":"2024-11-03T05:54:24.594047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Getting a list of all the study IDs and paths to their images\nimages_dir_path = r'/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images'\nstudy_id_list = os.listdir(images_dir_path)\nstudy_id_paths = [(x, f\"{images_dir_path}/{x}\") for x in study_id_list]\n\n# Initialize the metadata dictionary\nmeta_df = {}\n\n# Process each study and its series\nfor study_id, study_folder_path in study_id_paths:\n    series_ids = []\n    series_descriptions = []\n    \n    # Get all the series IDs (folders) within the study folder\n    try:\n        series_folders = os.listdir(study_folder_path)\n    except FileNotFoundError as e:\n        print(f\"Error: Folder not found for study {study_id}. Skipping this study.\")\n        continue  # Skip this study if the folder doesn't exist\n\n    # Process each series in the study folder\n    for series_id in series_folders:\n        try:\n            # Fetch the series description from the dataframe\n            series_description = series_desc_df[series_desc_df['series_id'] == int(series_id)]['series_description'].iloc[0]\n        except (IndexError, ValueError):\n            # Handle cases where series_id is not found in the dataframe or can't be converted to int\n            series_description = 'Unknown'\n\n        # Append series ID and description to the lists\n        series_ids.append(series_id)\n        series_descriptions.append(series_description)\n    \n    # Add metadata for the current study_id\n    meta_df[int(study_id)] = {\n        'folder_path': study_folder_path,\n        'series_ids': series_ids,\n        'series_descriptions': series_descriptions\n    }","metadata":{"execution":{"iopub.status.busy":"2024-11-03T06:48:30.763815Z","iopub.execute_input":"2024-11-03T06:48:30.764582Z","iopub.status.idle":"2024-11-03T06:48:35.046938Z","shell.execute_reply.started":"2024-11-03T06:48:30.764539Z","shell.execute_reply":"2024-11-03T06:48:35.046043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def add_desc(df, meta_df):\n    df['series_desc'] = None\n\n    # Iterate over rows in the dataframe\n    for idx, coor_row in df.iterrows():\n        try:\n            # Find the meta_df for the study_id\n            meta_info = meta_df[int(coor_row['study_id'])]\n\n            # Find the index of the series_id in the meta_info\n            series_index = meta_info['series_ids'].index(str(coor_row['series_id']))\n\n            # Get the corresponding series description\n            series_desc = meta_info['series_descriptions'][series_index]\n\n            # Update the series_desc column\n            df.at[idx, 'series_desc'] = series_desc\n\n        except KeyError:\n            print(f\"Error processing study_id: {coor_row['study_id']} - Study ID not found in meta_df\")\n            df.at[idx, 'series_desc'] = 'Unknown'\n        except ValueError:\n            print(f\"Error processing study_id: {coor_row['study_id']} - Series ID not found in meta_df\")\n            df.at[idx, 'series_desc'] = 'Unknown'\n        except Exception as e:\n            print(f\"Error processing study_id: {coor_row['study_id']} - {e}\")\n            df.at[idx, 'series_desc'] = 'Unknown'\n    \n    return df\n\n# Apply the function\ncoords_with_desc = label_coords_df.copy()\ncoords_with_desc = add_desc(coords_with_desc, meta_df)\ncoords_with_desc.head(20)","metadata":{"execution":{"iopub.status.busy":"2024-11-03T06:48:35.048616Z","iopub.execute_input":"2024-11-03T06:48:35.048927Z","iopub.status.idle":"2024-11-03T06:48:38.964484Z","shell.execute_reply.started":"2024-11-03T06:48:35.048894Z","shell.execute_reply":"2024-11-03T06:48:38.963610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sagt1_df = coords_with_desc[coords_with_desc['series_desc'] == 'Sagittal T1'].copy()\nsagt2_df = coords_with_desc[coords_with_desc['series_desc'] == 'Sagittal T2/STIR'].copy()\naxialt2_df = coords_with_desc[coords_with_desc['series_desc'] == 'Axial T2'].copy()","metadata":{"execution":{"iopub.status.busy":"2024-11-03T06:48:46.277949Z","iopub.execute_input":"2024-11-03T06:48:46.278408Z","iopub.status.idle":"2024-11-03T06:48:46.331349Z","shell.execute_reply.started":"2024-11-03T06:48:46.278359Z","shell.execute_reply":"2024-11-03T06:48:46.329948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sagt1_df.head(20)","metadata":{"execution":{"iopub.status.busy":"2024-11-03T06:55:38.105704Z","iopub.execute_input":"2024-11-03T06:55:38.106096Z","iopub.status.idle":"2024-11-03T06:55:38.124173Z","shell.execute_reply.started":"2024-11-03T06:55:38.106058Z","shell.execute_reply":"2024-11-03T06:55:38.122834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for (study_id, series_id), instance_data in sagt1_df.groupby(['study_id', 'series_id',]):\n    for level in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']:\n            level_data = instance_data[instance_data['level'] == level]\n            print(level_data)\n            for _, row in level_data.iterrows():\n                print(row['x'], row['y'])\n                break\n            break\n    break","metadata":{"execution":{"iopub.status.busy":"2024-11-03T07:17:19.011643Z","iopub.execute_input":"2024-11-03T07:17:19.012053Z","iopub.status.idle":"2024-11-03T07:17:19.035218Z","shell.execute_reply.started":"2024-11-03T07:17:19.012011Z","shell.execute_reply":"2024-11-03T07:17:19.034299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SagT1_YOLO:\n    def __init__(self, df, images_dir, output_dir, img_size=384):\n        \"\"\"\n        Initialize Sagittal T1 YOLO detector\n        Args:\n            df: DataFrame with annotations\n            images_dir: Path to DICOM images\n            output_dir: Path to save processed dataset\n            img_size: Target image size for YOLO\n        \"\"\"\n        self.df = df\n        self.images_dir = images_dir\n        self.output_dir = output_dir\n        self.img_size = img_size\n        \n        # Create initial directory structure\n        self.create_dataset_structure()\n        \n        # Process dataset\n        self.instance_coords_df = self.process_instance_coordinates()\n        self.processed_df = self.process_spine_dataset()\n        \n        # Create YAML and train model\n        self.yaml_path = self.create_dataset_yaml()\n        self.model, self.results = self.train_yolo()\n\n    def create_dataset_structure(self):\n        \"\"\"Create YOLO dataset directory structure\"\"\"\n        for split in ['train', 'val']:\n            for subdir in ['images', 'labels']:\n                path = os.path.join(self.output_dir, split, subdir)\n                os.makedirs(path, exist_ok=True)\n\n    def process_instance_coordinates(self):\n        \"\"\"\n        Process coordinates for Sagittal T1 images with coordinate sharing across instances\n        Returns DataFrame with coordinates for all instances, sharing information across the series\n        \"\"\"\n        result_records = []\n\n        # Group by series to process related instances\n        for (study_id, series_id), series_data in self.df.groupby(['study_id', 'series_id']):\n            # First, collect all coordinates for each level in the series\n            series_level_coords = {}\n            for level in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']:\n                level_key = level.lower().replace('/', '_')\n                level_data = series_data[series_data['level'] == level]\n\n                if not level_data.empty:\n                    right_coords = []\n                    left_coords = []\n                    instances_with_data = set()  # Track which instances have data for this level\n\n                    for _, row in level_data.iterrows():\n                        instances_with_data.add(row['instance_number'])\n                        if 'Right' in row['condition']:\n                            right_coords.append((row['x'], row['y']))\n                        elif 'Left' in row['condition']:\n                            left_coords.append((row['x'], row['y']))\n\n                    # Calculate coordinates for this level\n                    level_coords = {\n                        'instances': instances_with_data,\n                        'coords': {\n                            'right': (np.mean([x for x, _ in right_coords]) if right_coords else None,\n                                    np.mean([y for _, y in right_coords]) if right_coords else None),\n                            'left': (np.mean([x for x, _ in left_coords]) if left_coords else None,\n                                   np.mean([y for _, y in left_coords]) if left_coords else None)\n                        }\n                    }\n\n                    # Calculate center coordinates if possible\n                    if right_coords and left_coords:\n                        level_coords['coords']['center'] = (\n                            (level_coords['coords']['right'][0] + level_coords['coords']['left'][0]) / 2,\n                            (level_coords['coords']['right'][1] + level_coords['coords']['left'][1]) / 2\n                        )\n                    elif right_coords:\n                        level_coords['coords']['center'] = level_coords['coords']['right']\n                    elif left_coords:\n                        level_coords['coords']['center'] = level_coords['coords']['left']\n                    else:\n                        level_coords['coords']['center'] = (None, None)\n\n                    series_level_coords[level_key] = level_coords\n\n            # Get all unique instance numbers in the series\n            all_instances = series_data['instance_number'].unique()\n\n            # Create records for each instance, sharing coordinates across the series\n            for instance_number in all_instances:\n                record = {\n                    'study_id': study_id,\n                    'series_id': series_id,\n                    'instance_number': instance_number\n                }\n\n                # Add source tracking\n                record['coordinate_sources'] = {}\n\n                # Add coordinates for all levels to this instance\n                for level in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']:\n                    level_key = level.lower().replace('/', '_')\n                    if level_key in series_level_coords:\n                        level_data = series_level_coords[level_key]\n\n                        # Record which instances provided data for this level\n                        record['coordinate_sources'][level_key] = list(level_data['instances'])\n\n                        # Add right coordinates\n                        if level_data['coords']['right'][0] is not None:\n                            record[f'{level_key}_right_x'] = level_data['coords']['right'][0]\n                            record[f'{level_key}_right_y'] = level_data['coords']['right'][1]\n\n                        # Add left coordinates\n                        if level_data['coords']['left'][0] is not None:\n                            record[f'{level_key}_left_x'] = level_data['coords']['left'][0]\n                            record[f'{level_key}_left_y'] = level_data['coords']['left'][1]\n\n                        # Add center coordinates\n                        if level_data['coords']['center'][0] is not None:\n                            record[f'{level_key}_center_x'] = level_data['coords']['center'][0]\n                            record[f'{level_key}_center_y'] = level_data['coords']['center'][1]\n\n                result_records.append(record)\n\n        # Convert to DataFrame\n        result_df = pd.DataFrame(result_records)\n\n        # Add metadata about coordinate availability\n        result_df['available_levels'] = result_df.apply(\n            lambda row: [\n                level for level in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\n                if not pd.isna(row.get(f\"{level.lower().replace('/', '_')}_center_x\"))\n            ],\n            axis=1\n        )\n\n        result_df['total_levels'] = result_df['available_levels'].apply(len)\n\n        # Print statistics\n        print(\"\\nDataset Statistics:\")\n        print(f\"Total series processed: {len(result_df['series_id'].unique())}\")\n        print(f\"Total instances processed: {len(result_df)}\")\n        print(\"\\nLevel availability:\")\n        for level in ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']:\n            count = result_df[f'{level}_center_x'].notna().sum()\n            print(f\"{level}: {count} instances ({count/len(result_df)*100:.1f}%)\")\n\n        return result_df\n\n    def create_yolo_annotation(self, row, image_width, image_height):\n        \"\"\"Create YOLO format annotations for Sagittal T1 images\"\"\"\n        annotations = []\n        box_width = 0.05\n        box_height = 0.05\n        \n        level_map = {'l1_l2': 0, 'l2_l3': 1, 'l3_l4': 2, 'l4_l5': 3, 'l5_s1': 4}\n        \n        for level, idx in level_map.items():\n            x_coord = row.get(f'{level}_center_x')\n            y_coord = row.get(f'{level}_center_y')\n            \n            if pd.notna(x_coord) and pd.notna(y_coord):\n                x_norm = x_coord / image_width\n                y_norm = y_coord / image_height\n                annotations.append(f\"{idx} {x_norm:.6f} {y_norm:.6f} {box_width:.6f} {box_height:.6f}\")\n        \n        return annotations\n\n    def process_spine_dataset(self):\n        \"\"\"Process and save dataset in YOLO format\"\"\"\n        # Split studies\n        studies = self.instance_coords_df['study_id'].unique()\n        train_studies, val_studies = train_test_split(studies, train_size=0.8, random_state=42)\n        \n        processed_counts = {'train': 0, 'val': 0}\n        failed_cases = []\n        \n        for _, row in tqdm(self.instance_coords_df.iterrows(), desc=\"Processing Sagittal T1 images\"):\n            try:\n                # Convert IDs to integers for path construction\n                study_id = str(int(row['study_id']))\n                series_id = str(int(row['series_id']))\n                instance_number = str(int(row['instance_number']))\n                \n                # Construct image path\n                img_path = os.path.join(self.images_dir, study_id, series_id, instance_number)\n                \n                # Try with and without .dcm extension\n                if os.path.exists(img_path + '.dcm'):\n                    img_path = img_path + '.dcm'\n                elif not os.path.exists(img_path):\n                    raise FileNotFoundError(f\"Image not found: {img_path}\")\n                \n                # Read and process image\n                ds = pydicom.dcmread(img_path)\n                image = ds.pixel_array\n                h, w = image.shape\n                \n                # Create annotations\n                annotations = self.create_yolo_annotation(row, w, h)\n                if not annotations:\n                    continue\n                \n                # Prepare image\n                image_normalized = cv2.normalize(image, None, 0, 255, cv2.NORM_MINMAX, cv2.CV_8U)\n                image_resized = cv2.resize(image_normalized, (self.img_size, self.img_size))\n                \n                # Save files\n                is_train = row['study_id'] in train_studies\n                split = 'train' if is_train else 'val'\n                \n                img_filename = f\"{study_id}_{series_id}_{instance_number}.png\"\n                label_filename = f\"{study_id}_{series_id}_{instance_number}.txt\"\n                \n                cv2.imwrite(os.path.join(self.output_dir, split, 'images', img_filename), \n                           image_resized)\n                with open(os.path.join(self.output_dir, split, 'labels', label_filename), 'w') as f:\n                    f.write('\\n'.join(annotations))\n                \n                # Add split information to row\n                row['split'] = split\n                processed_counts[split] += 1\n                \n            except Exception as e:\n                failed_cases.append((study_id, series_id, instance_number, str(e)))\n        \n        print(f\"\\nProcessing Summary:\")\n        print(f\"Training images: {processed_counts['train']}\")\n        print(f\"Validation images: {processed_counts['val']}\")\n        \n        if failed_cases:\n            print(\"\\nFailed cases:\")\n            for case in failed_cases:\n                print(f\"Study {case[0]}, Series {case[1]}, Instance {case[2]}: {case[3]}\")\n        \n        return self.instance_coords_df\n\n    def create_dataset_yaml(self):\n        \"\"\"Create YOLO dataset configuration file\"\"\"\n        yaml_content = {\n            'path': os.path.abspath(self.output_dir),\n            'train': 'train/images',\n            'val': 'val/images',\n            'nc': 5,  # number of classes\n            'names': {\n                0: 'L1/L2',\n                1: 'L2/L3',\n                2: 'L3/L4',\n                3: 'L4/L5',\n                4: 'L5/S1'\n            }\n        }\n\n        yaml_path = os.path.join(self.output_dir, 'dataset.yaml')\n        with open(yaml_path, 'w') as f:\n            yaml.dump(yaml_content, f, sort_keys=False)\n\n        return yaml_path\n\n    def train_yolo(self):\n        \"\"\"Train YOLO model for Sagittal T1 images\"\"\"\n        try:\n            model = YOLO('yolov8x.pt')\n            \n            config = {\n                'data': self.yaml_path,\n                'imgsz': self.img_size,\n                'batch': 16,\n                'epochs': 20,\n                'patience': 5,\n                'device': '0',\n                'workers': 8,\n                'project': 'spine_detection',\n                'name': 'sagittal_t1_yolo',\n                'exist_ok': True,\n                'pretrained': True,\n                'optimizer': 'AdamW',\n                'verbose': True,\n                'seed': 42,\n                'deterministic': True,\n                'dropout': 0.2,\n                'lr0': 0.001,\n                'lrf': 0.01,\n                'momentum': 0.937,\n                'weight_decay': 0.0005,\n                'warmup_epochs': 10,\n                'warmup_momentum': 0.8,\n                'box': 7.5,\n                'cls': 0.5,\n                'dfl': 1.5,\n                'close_mosaic': 10,\n                'amp': True,\n                # Augmentation settings\n                #'degrees': 5,\n                #'translate': 0.1,\n                #'scale': 0.5,\n                #'shear': 2.0,\n                #'flipud': 0.5,\n                #'mosaic': 1.0,\n                #'mixup': 0.3,\n                #'copy_paste': 0.3\n            }\n            \n            results = model.train(**config)\n            return model, results\n            \n        except Exception as e:\n            print(f\"Error training model: {str(e)}\")\n            return None, None\n\n    def get_training_stats(self):\n        \"\"\"Get statistics about the processed dataset\"\"\"\n        stats = {\n            'total_series': len(self.instance_coords_df['series_id'].unique()),\n            'total_instances': len(self.instance_coords_df),\n            'train_images': len(self.processed_df[self.processed_df['split'] == 'train']),\n            'val_images': len(self.processed_df[self.processed_df['split'] == 'val']),\n            'level_coverage': {}\n        }\n        \n        for level in ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']:\n            count = self.instance_coords_df[f'{level}_center_x'].notna().sum()\n            stats['level_coverage'][level] = {\n                'count': count,\n                'percentage': count/len(self.instance_coords_df)*100\n            }\n            \n        return stats","metadata":{"execution":{"iopub.status.busy":"2024-11-03T08:41:17.603829Z","iopub.execute_input":"2024-11-03T08:41:17.604254Z","iopub.status.idle":"2024-11-03T08:41:17.655658Z","shell.execute_reply.started":"2024-11-03T08:41:17.604204Z","shell.execute_reply":"2024-11-03T08:41:17.654655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images_dir  = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images'\noutput_dir = '/kaggle/working/spine_dataset_sagt1'\n\nsag_t1_model = SagT1_YOLO(\n    df=sagt1_df,\n    images_dir=images_dir,\n    output_dir=output_dir,\n    img_size=384\n)","metadata":{"execution":{"iopub.status.busy":"2024-11-03T08:41:20.873426Z","iopub.execute_input":"2024-11-03T08:41:20.873821Z","iopub.status.idle":"2024-11-03T10:05:25.592697Z","shell.execute_reply.started":"2024-11-03T08:41:20.873785Z","shell.execute_reply":"2024-11-03T10:05:25.591454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_trained_model(model, save_path='/kaggle/working/sagt1_yolo_model.pt'):\n    \"\"\"\n    Save the trained YOLO model\n    \"\"\"\n    if model is not None:\n        # The model is saved during training in the 'best.pt' file\n        # We can copy it to our desired location\n        train_dir = os.path.join('spine_detection', 'sagittal_t1_yolo')\n        best_model_path = os.path.join(train_dir, 'weights', 'best.pt')\n        \n        if os.path.exists(best_model_path):\n            import shutil\n            shutil.copy(best_model_path, save_path)\n            print(f\"Model saved to {save_path}\")\n        else:\n            print(\"Best model weights not found\")\n    else:\n        print(\"No model to save\")\n        \nsave_trained_model(sag_t1_model.model)","metadata":{"execution":{"iopub.status.busy":"2024-11-03T10:12:31.097734Z","iopub.execute_input":"2024-11-03T10:12:31.098557Z","iopub.status.idle":"2024-11-03T10:12:31.216545Z","shell.execute_reply.started":"2024-11-03T10:12:31.098515Z","shell.execute_reply":"2024-11-03T10:12:31.215535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_spine_predictions(image_path, model_path='best_spine_model.pt', \n                         conf_threshold=0.25, iou_threshold=0.45, img_size=384):\n    \"\"\"\n    Plot YOLO predictions for spine levels on a DICOM image\n    \n    Args:\n        image_path: Path to DICOM image\n        model_path: Path to saved YOLO model\n        conf_threshold: Confidence threshold for predictions\n        iou_threshold: IOU threshold for NMS\n        img_size: Image size for model input\n    \"\"\"\n    # Load model\n    model = YOLO(model_path)\n    \n    # Read DICOM\n    ds = pydicom.dcmread(image_path)\n    image = ds.pixel_array\n    \n    # Normalize and resize\n    image_normalized = cv2.normalize(image, None, 0, 255, cv2.NORM_MINMAX, cv2.CV_8U)\n    image_resized = cv2.resize(image_normalized, (img_size, img_size))\n    \n    # Convert grayscale to RGB\n    image_rgb = np.stack([image_resized] * 3, axis=-1)\n    \n    # Create figure\n    plt.figure(figsize=(15, 7))\n    \n    # Plot original image\n    plt.subplot(1, 2, 1)\n    plt.imshow(image_resized, cmap='gray')\n    plt.title('Original Image')\n    plt.axis('off')\n    \n    # Plot image with predictions\n    plt.subplot(1, 2, 2)\n    plt.imshow(image_resized, cmap='gray')\n    plt.title('Predictions')\n    \n    # Get predictions\n    results = model.predict(\n        source=image_rgb,\n        conf=conf_threshold,\n        iou=iou_threshold\n    )\n    \n    # Define colors for each level\n    colors = ['red', 'green', 'blue', 'yellow', 'purple']\n    level_names = ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\n    \n    if results[0].boxes is not None:\n        boxes = results[0].boxes.cpu().numpy()\n        \n        # Sort boxes by y-coordinate to display levels in order\n        box_data = []\n        for box in boxes:\n            cls_id = int(box.cls[0])\n            conf = box.conf[0]\n            x1, y1, x2, y2 = box.xyxy[0]\n            box_data.append((y1, cls_id, conf, x1, y1, x2, y2))\n        \n        box_data.sort()  # Sort by y1 coordinate\n        \n        # Plot each detection\n        for i, (_, cls_id, conf, x1, y1, x2, y2) in enumerate(box_data):\n            color = colors[cls_id]\n            level_name = level_names[cls_id]\n            \n            # Draw bounding box\n            plt.gca().add_patch(plt.Rectangle(\n                (x1, y1), x2-x1, y2-y1,\n                fill=False, color=color, linewidth=2\n            ))\n            \n            # Add label\n            plt.text(\n                x2 + 5, (y1 + y2) / 2, \n                f'{level_name}: {conf:.2f}',\n                color=color, fontsize=8, verticalalignment='center',\n                bbox=dict(facecolor='white', alpha=0.7, edgecolor='none')\n            )\n            \n            # Print detection info\n            print(f\"Found {level_name} with confidence {conf:.2f}\")\n    else:\n        print(\"No detections found\")\n    \n    plt.axis('off')\n    plt.tight_layout()\n    plt.show()\n\nplot_spine_predictions(\n    image_path=\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/4646740/3486248476/15.dcm\",\n    model_path=\"/kaggle/working/sagt1_yolo_model.pt\",\n    conf_threshold=0.25,\n    iou_threshold=0.45\n)","metadata":{"execution":{"iopub.status.busy":"2024-11-03T10:14:54.549975Z","iopub.execute_input":"2024-11-03T10:14:54.550385Z","iopub.status.idle":"2024-11-03T10:14:56.337163Z","shell.execute_reply.started":"2024-11-03T10:14:54.550349Z","shell.execute_reply":"2024-11-03T10:14:56.336187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_multiple_images(image_paths, model_path='/kaggle/working/sagt1_yolo_model.pt', \n                          conf_threshold=0.25, iou_threshold=0.45):\n    \"\"\"\n    Process multiple images and display their predictions\n    \n    Args:\n        image_paths: List of paths to DICOM images\n        model_path: Path to saved YOLO model\n        conf_threshold: Confidence threshold for predictions\n        iou_threshold: IOU threshold for NMS\n    \"\"\"\n    n_images = len(image_paths)\n    cols = min(3, n_images)  # Max 3 images per row\n    rows = (n_images - 1) // cols + 1\n    \n    plt.figure(figsize=(6*cols, 6*rows))\n    \n    for i, img_path in enumerate(image_paths, 1):\n        plt.subplot(rows, cols, i)\n        plot_spine_predictions(img_path, model_path, conf_threshold, iou_threshold)\n        plt.title(f'Image {i}')\n    \n    plt.tight_layout()\n    plt.show()\n\n# Usage example:\nimage_paths = glob.glob('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/4646740/3486248476/*.dcm')\nprocess_multiple_images(image_paths)","metadata":{"execution":{"iopub.status.busy":"2024-11-03T10:17:06.620405Z","iopub.execute_input":"2024-11-03T10:17:06.621279Z","iopub.status.idle":"2024-11-03T10:17:41.805286Z","shell.execute_reply.started":"2024-11-03T10:17:06.621238Z","shell.execute_reply":"2024-11-03T10:17:41.804392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}