{"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":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9462464,"sourceType":"datasetVersion","datasetId":5753101},{"sourceId":9462730,"sourceType":"datasetVersion","datasetId":5753305},{"sourceId":9562240,"sourceType":"datasetVersion","datasetId":5827317},{"sourceId":90871,"sourceType":"modelInstanceVersion","modelInstanceId":76180,"modelId":100857}],"dockerImageVersionId":30761,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"import torch\nimport sys\n\nsys.path.append('/kaggle/input/packages/kaggle/working/python-packages')\nsys.path.append('/kaggle/input/mysegment-anything/kaggle/working/segment-anything-2')\nfrom sam2.build_sam import build_sam2\nfrom sam2.sam2_image_predictor import SAM2ImagePredictor\nfrom sam2.automatic_mask_generator import SAM2AutomaticMaskGenerator\ncheckpoint = \"/kaggle/input/segment-anything-2/pytorch/sam2-hiera-large/1/sam2_hiera_large.pt\"\nmodel_cfg = \"sam2_hiera_l.yaml\"\n#predictor = SAM2ImagePredictor(build_sam2(model_cfg, checkpoint))\n#predictor = SAM2AutomaticMaskGenerator(sam2)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T08:52:43.383384Z","iopub.execute_input":"2024-10-08T08:52:43.383760Z","iopub.status.idle":"2024-10-08T08:52:49.506281Z","shell.execute_reply.started":"2024-10-08T08:52:43.383723Z","shell.execute_reply":"2024-10-08T08:52:49.505481Z"}}},{"cell_type":"code","source":"import os\nimport pandas as pd\n\nroot_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\"\n\n# Load the data\ntrain_main = pd.read_csv(os.path.join(root_dir, \"train.csv\"))\ntrain_labels = pd.read_csv(os.path.join(root_dir, \"train_label_coordinates.csv\"))\ntrain_des = pd.read_csv(os.path.join(root_dir, \"train_series_descriptions.csv\"))\ntest_des = pd.read_csv(os.path.join(root_dir, \"test_series_descriptions.csv\"))\nsample_sub = pd.read_csv(os.path.join(root_dir, \"sample_submission.csv\"))\n\n# Create coordinates map and pairs list\ncoordinates_map = {}\nstudy_series_pairs = []\n\n# Populate coordinates\nfor index, row in train_labels.iterrows():\n    key = (row['study_id'], row['series_id'])  # Ensure the key is a tuple\n    if key not in coordinates_map:\n        coordinates_map[key] = {'coords': [], 'description': '', 'severity': ''}  # Initialize all fields\n    coordinates_map[key]['coords'].append({\n        'x': row['x'], \n        'y': row['y'], \n        'label': row['condition'], \n        'level': row['level'], \n        'instance': row['instance_number']\n    })\n\n# Populate descriptions\nfor index, row in train_des.iterrows():\n    key = (row['study_id'], row['series_id'])\n    if key not in coordinates_map:\n        coordinates_map[key] = {'coords': [], 'description': '', 'severity': ''}  # Initialize if not present\n    coordinates_map[key]['description'] = row['series_description']\n\n# Populate severity information dynamically and append to matching labels in coordinates\nfor index, row in train_main.iterrows():\n    study_id = row['study_id']\n    for series_id in coordinates_map.keys():\n        if series_id[0] == study_id:  # Match based on study_id\n            # Loop through the columns of train_main to match label and level\n            for col_name in train_main.columns:\n                if col_name != 'study_id':  # Skip the study_id column\n                    # Match label and level combination with the column name\n                    for coord in coordinates_map[series_id]['coords']:\n                        combined_label_level = f\"{coord['label']} {coord['level']}\".lower().replace(\"/\", \"_\").replace(\" \", \"_\")\n                        if combined_label_level in col_name.lower():\n                            # Append severity value from train_main to the 'severity' key in the coordinate\n                            coord['severity'] = row[col_name]\n# Function to calculate average bounding box for each series description\n# Function to calculate average min and max bounding box for each series description\n# Function to calculate average min and max bounding box for each series description\ndef calculate_avg_min_max_bounding_box(data):\n    bbox_averages = {}\n\n    for description, group in data.groupby('Description'):\n        total_min_x, total_min_y = 0, 0\n        total_max_x, total_max_y = 0, 0\n        total_boxes = 0\n\n        # Initialize min and max variables for each description\n        all_min_x, all_min_y = [], []\n        all_max_x, all_max_y = [], []\n\n        for index, row in group.iterrows():\n            coordinates_list = row['Coordinates']\n\n            # Skip this row if there are no coordinates\n            if not coordinates_list:\n                continue\n            \n            # For each series, find the min and max x, y values\n            min_x = min(coord['x'] for coord in coordinates_list)\n            min_y = min(coord['y'] for coord in coordinates_list)\n            max_x = max(coord['x'] for coord in coordinates_list)\n            max_y = max(coord['y'] for coord in coordinates_list)\n\n            # Collect the min and max values across all instances\n            all_min_x.append(min_x)\n            all_min_y.append(min_y)\n            all_max_x.append(max_x)\n            all_max_y.append(max_y)\n\n            total_boxes += 1\n\n        if total_boxes > 0:\n            # Calculate the average min and max values\n            avg_min_x = sum(all_min_x) / total_boxes\n            avg_min_y = sum(all_min_y) / total_boxes\n            avg_max_x = sum(all_max_x) / total_boxes\n            avg_max_y = sum(all_max_y) / total_boxes\n\n            # Store the results in the bbox_averages dictionary\n            bbox_averages[description] = {\n                'avg_min_x': avg_min_x, 'avg_min_y': avg_min_y,\n                'avg_max_x': avg_max_x, 'avg_max_y': avg_max_y,\n                'total_boxes': total_boxes\n            }\n\n    return bbox_averages\n\n# Function to aggregate data\ndef aggregate_data(base_dir, coordinates_map):\n    data = []\n    for (study_id, series_id), details in coordinates_map.items():\n        path = os.path.join(base_dir, str(study_id), str(series_id))\n        if not os.path.exists(path):\n            print(f\"Path not found: {path}\")\n            continue\n        \n        filepaths = [os.path.join(path, filename) for filename in sorted(os.listdir(path)) if filename.endswith('.dcm')]\n\n        # Aggregate filepaths and coordinate information\n        record = {\n            \"StudyID\": study_id,\n            \"SeriesID\": series_id,\n            \"FilePaths\": filepaths,\n            \"Coordinates\": details['coords'],\n            \"Description\": details['description']\n        }\n\n        # Add all other fields from train_main to the record\n        for key, value in details.items():\n            if key not in ['coords', 'description']:  # Skip 'coords' and 'description' fields\n                record[key] = value\n        \n        data.append(record)\n    \n    return pd.DataFrame(data)\n\n# Process data and create DataFrame\nbase_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/\"\ntrain_matrix_new = aggregate_data(base_dir, coordinates_map)\n\n# Further processing to include severity in class_map\nlevel = []\nclass_map = {}\n\n# Iterate over the DataFrame and its Coordinates list to extract label, level, and severity\nfor index, row in train_matrix_new.iterrows():\n    coordinates_list = row['Coordinates']\n    for coord in coordinates_list:\n        rlevel = coord['level']\n        rlabel = coord['label']\n        severity = coord.get('severity', 'None')  # Get severity or set to 'None' if missing\n        #level.append([rlabel, rlevel, severity])  # Include severity in the tuple\n        \n        # Only include non-NaN severities in the class_map\n        if pd.notna(severity) and severity != 'None':  # Filter out NaN and 'None'\n            level.append([rlabel, rlevel, severity])\n\n# Create a set of unique tuples for label, level, and severity\nunique_tuples = set(tuple(item) for item in level)\n\n# Convert tuples back to list of lists\nunique_level = [list(item) for item in unique_tuples]\n\n# Create the class_map including severity\nclass_map = {tuple(pair): index for index, pair in enumerate(unique_level)}\n\n# Output the class_map to check the result\nprint(\"Class Map with Severity:\")\nprint(len(class_map))\nprint(class_map)\n\n\n# Show counts\nnum_unique_study_ids = train_matrix_new['StudyID'].nunique()\nnum_unique_series_ids = train_matrix_new['SeriesID'].nunique()\nprint(\"Number of unique StudyIDs:\", num_unique_study_ids)\nprint(\"Number of unique SeriesIDs:\", num_unique_series_ids)\n\n# Filter the data\nSTIR_Data = train_matrix_new[train_matrix_new['Description'].str.contains('STIR', na=False)]\nT1_Data = train_matrix_new[train_matrix_new['Description'] == 'Sagittal T1']\nT2_Data = train_matrix_new[train_matrix_new['Description'] == 'Axial T2']\n\nprint(len(STIR_Data))\nprint(len(T1_Data))\nprint(len(T2_Data))\n#print(STIR_Data.head())\nprint(STIR_Data.iloc[1]['Coordinates'])\n\n# Calculate average min and max bounding boxes for each series description\nglobal_bounding_box_averages = calculate_avg_min_max_bounding_box(train_matrix_new)\n\n# Example to print average min/max bounding box info\nSTIR_avg_bbox = global_bounding_box_averages.get('Sagittal T2/STIR', {})\nT1_avg_bbox = global_bounding_box_averages.get('Sagittal T1', {})\nT2_avg_bbox = global_bounding_box_averages.get('Axial T2', {})\n\nprint(\"STIR Avg Min/Max Bounding Box:\", STIR_avg_bbox)\nprint(\"T1 Avg Min/Max Bounding Box:\", T1_avg_bbox)\nprint(\"T2 Avg Min/Max Bounding Box:\", T2_avg_bbox)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T10:55:02.022796Z","iopub.execute_input":"2024-10-08T10:55:02.023358Z","iopub.status.idle":"2024-10-08T10:56:05.079581Z","shell.execute_reply.started":"2024-10-08T10:55:02.023292Z","shell.execute_reply":"2024-10-08T10:56:05.078483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"DICOMDataset Class","metadata":{}},{"cell_type":"markdown","source":"import os\nimport pydicom\nimport numpy as np\nimport cv2\nfrom PIL import Image\nimport torch\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\nfrom scipy.ndimage import gaussian_filter  # For Gaussian smoothing\nfrom skimage import exposure\nfrom skimage.transform import resize\nfrom sam2.automatic_mask_generator import SAM2AutomaticMaskGenerator\nfrom sam2.build_sam import build_sam2\n\nclass DICOMDataset(Dataset):\n    def __init__(self, dataframe, class_map, target_shape=(20, 256, 256), max_samples=None, save_path='./kaggle/working/'):\n        self.dataframe = dataframe\n        self.target_shape = target_shape\n        self.class_map = class_map if class_map is not None else {}\n        self.series_dict = self.group_series_by_id()\n        self.series_ids = list(self.series_dict.keys())\n        self.max_samples = max_samples if max_samples is not None else len(self.series_ids)\n        self.save_path = save_path  # Directory to save test volumes and coordinates\n        self.device = torch.device(\"cuda\")  # Define the device here\n        \n        # Initialize the SAM2 model\n        checkpoint = \"/kaggle/input/segment-anything-2/pytorch/sam2-hiera-large/1/sam2_hiera_large.pt\"\n        model_cfg = \"sam2_hiera_l.yaml\"\n        self.sam2 = build_sam2(model_cfg, checkpoint, device=self.device, apply_postprocessing=True)\n\n    def group_series_by_id(self):\n        series_dict = {}\n        for _, row in self.dataframe.iterrows():\n            series_uid = row['SeriesID']\n            if series_uid not in series_dict:\n                series_dict[series_uid] = []\n            series_dict[series_uid].append(row)\n        return series_dict\n\n    def extract_file_number(self, file_path):\n        filename = os.path.basename(file_path)\n        file_number = filename.split('.')[0]\n        return int(file_number)\n\n    def calculate_target_spacing(self, dicom_slice_thickness, dicom_pixel_spacing, method='original_resolution'):\n        pixel_spacing_x, pixel_spacing_y = dicom_pixel_spacing\n        slice_thickness = dicom_slice_thickness\n\n        if method == 'isotropic_max':\n            target_spacing = max(slice_thickness, pixel_spacing_x, pixel_spacing_y)\n            return (target_spacing, target_spacing, target_spacing)\n        \n        elif method == 'isotropic_avg':\n            target_spacing = (slice_thickness + pixel_spacing_x + pixel_spacing_y) / 3.0\n            return (target_spacing, target_spacing, target_spacing)\n        \n        elif method == 'original_resolution':\n            return (slice_thickness, pixel_spacing_x, pixel_spacing_y)\n        \n        else:\n            raise ValueError(\"Invalid method. Choose from 'isotropic_max', 'isotropic_avg', or 'original_resolution'.\")\n\n    def apply_noise_reduction(self, image_array, h=25, templateWindowSize=9, searchWindowSize=31):\n            \"\"\"\n            Apply Non-Local Means Denoising to reduce noise in the image.\n            Args:\n                image_array (np.ndarray): The input 2D slice as a NumPy array.\n                h (int): Parameter for controlling the filter strength.\n                templateWindowSize (int): Size of the template patch used to compute weights.\n                searchWindowSize (int): Size of the window used to search for similar patches.\n            Returns:\n                np.ndarray: The denoised image array.\n            \"\"\"\n            # Ensure the image is in the correct format (uint8 for OpenCV NLM function)\n            #image_array_uint8 = image_array.astype(np.uint8)\n\n            # Apply Non-Local Means Denoising\n            denoised_image = cv2.fastNlMeansDenoising(image_array, None, h, templateWindowSize, searchWindowSize)\n\n            return denoised_image\n\n    def load_dicom_with_pil(self, dicom_file, resize_shape=(256, 256)):\n        try:\n            #print(\"Processing DICOM file:\", dicom_file)\n            dicom_data = pydicom.dcmread(dicom_file)\n            image_array = dicom_data.pixel_array.astype(np.float32)\n            self.o_shape = image_array.shape\n\n            # Normalize the image to [0, 255]\n            image_array = (image_array - image_array.min()) / (image_array.max() - image_array.min()) * 255\n            image_array = image_array.astype(np.uint8)\n\n            # Apply OpenCV enhancements\n            image_array = self.apply_noise_reduction(image_array)\n            image_array = self.apply_opencv_enhancements(Image.fromarray(image_array))\n            #image_array = self.apply_noise_reduction(image_array)\n\n            # Convert grayscale to RGB for compatibility with SAM2\n            image_rgb = np.stack([image_array] * 3, axis=-1)\n\n            # Extract SliceThickness and PixelSpacing for later use\n            slice_spacing = float(dicom_data.SliceThickness) if hasattr(dicom_data, 'SliceThickness') else 1.0\n            pixel_spacing = dicom_data.PixelSpacing if hasattr(dicom_data, 'PixelSpacing') else [1.0, 1.0]\n            self.o_space = pixel_spacing\n            #print(\"Processed image array shape:\", image_rgb.shape)\n\n            # Return the preprocessed image and metadata without any mask application\n            return image_rgb, slice_spacing, pixel_spacing\n\n        except Exception as e:\n            print(f\"Error loading DICOM file {dicom_file}: {e}\")\n            return None, None, None\n\n\n\n    def apply_opencv_enhancements(self, pil_image):\n        \"\"\"\n        Apply automatic contrast and brightness adjustments using OpenCV.\n        Args:\n            pil_image (PIL.Image): The input image.\n        Returns:\n            PIL.Image: The enhanced image.\n        \"\"\"\n        # Convert PIL image to OpenCV format (numpy array)\n        image_cv = np.array(pil_image)\n\n        # Apply CLAHE (Contrast Limited Adaptive Histogram Equalization) for contrast enhancement\n        clahe = cv2.createCLAHE(clipLimit=0.15, tileGridSize=(64, 64))\n        image_cv = clahe.apply(image_cv)\n\n        # Calculate histogram of the image\n        hist, bins = np.histogram(image_cv.flatten(), 256, [0, 256])\n\n        # Determine the intensity values covering the central 90% of the histogram (5th to 95th percentile)\n        cdf = hist.cumsum()  # Cumulative distribution function\n        cdf_normalized = cdf / cdf[-1]  # Normalize to [0, 1]\n\n        # Find the 5th and 95th percentiles\n        lower_bound = np.searchsorted(cdf_normalized, 0.05)\n        upper_bound = np.searchsorted(cdf_normalized, 0.95)\n\n        # Clip the image intensities to this range\n        clipped_image = np.clip(image_cv, lower_bound, upper_bound)\n\n        # Normalize to [0, 255] based on the clipped range\n        clipped_image = (clipped_image - lower_bound) / (upper_bound - lower_bound) * 255.0\n\n        # Convert back to uint8\n        enhanced_image = np.clip(clipped_image, 0, 255).astype(np.uint8)\n\n        # Convert back to PIL image\n        #enhanced_image = Image.fromarray(enhanced_image)\n\n        return enhanced_image\n    \n    \n    def fill_structures_in_volume(self, volume, kernel_size=1):\n        \"\"\"\n        Fill structures in the 3D volume using morphological operations.\n        Args:\n            volume (np.ndarray): The input 3D volume.\n            kernel_size (int): Size of the structuring element for morphological operations.\n        Returns:\n            np.ndarray: The volume with filled structures.\n        \"\"\"\n        filled_volume = np.zeros_like(volume)\n        for i in range(volume.shape[0]):  # Process slice by slice\n            # Convert slice to uint8 for morphological operations\n            slice_uint8 = cv2.normalize(volume[i], None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n\n            # Create structuring element\n            kernel = np.ones((kernel_size, kernel_size), np.uint8)\n            # Apply morphological closing to fill structures\n            filled_slice = cv2.morphologyEx(slice_uint8, cv2.MORPH_CLOSE, kernel)\n            # Store the filled slice back\n            filled_volume[i] = filled_slice\n\n        return filled_volume\n\n\n    def resize_mask(self, mask, target_shape):\n        \"\"\"\n        Resize the mask to match the target shape.\n        Args:\n            mask (np.ndarray): The mask to resize.\n            target_shape (tuple): The target shape (H, W).\n        Returns:\n            np.ndarray: The resized mask.\n        \"\"\"\n        return resize(mask, target_shape, mode='constant', anti_aliasing=False, preserve_range=True)\n\n    def stack_and_resample_dicom_slices(self, file_paths_list, coordinates, pixel_spacing, slice_thickness):\n        # Create a list to hold tuples of (instance_number, file_path)\n        dicom_files_with_instance = []\n        \n        z_depth = self.get_z_depth(file_paths_list)\n        #z_depth = 1\n        #self.o_shape = (z_depth, self.o_shape[0], self.o_shape[1])\n\n        for file_path in file_paths_list:\n            try:\n                # Read the DICOM file\n                dicom_data = pydicom.dcmread(file_path)\n                # Extract the InstanceNumber\n                instance_number = int(dicom_data.InstanceNumber)\n                # Append the tuple (instance_number, file_path) to the list\n                dicom_files_with_instance.append((instance_number, file_path))\n            except Exception as e:\n                print(f\"Error reading DICOM file {file_path}: {e}\")\n                continue\n\n        # Sort the list by instance_number\n        dicom_files_with_instance.sort(key=lambda x: x[0])\n\n        # Extract the sorted file paths\n        sorted_file_paths = [file_path for _, file_path in dicom_files_with_instance]\n\n        # Initialize lists to collect slices and other data\n        slices = []\n        mask_scores = []  # To keep track of mask scores for each slice\n\n        for file_path in sorted_file_paths:\n            # Continue processing as usual\n            #print(f\"Processing DICOM file: {file_path}\")\n            image_array, slice_spacing, pixel_spacing = self.load_dicom_with_pil(file_path)\n            self.o_shape = (z_depth, self.o_shape[0], self.o_shape[1])\n            if image_array is not None:\n                # Ensure the image is in RGB format (3 channels)\n                if len(image_array.shape) == 2:  # If grayscale, convert to RGB\n                    image_array = np.stack([image_array] * 3, axis=-1)\n\n                slices.append(image_array)\n\n                # Set up SAM2 mask predictor and generate masks for each slice\n                predictor = SAM2ImagePredictor(\n                    self.sam2,\n                    device=torch.device(\"cuda\"),\n                    apply_postprocessing=True,\n                    pred_iou_thresh=0.60,\n                    stability_score_thresh=0.70,\n                )\n\n                predictor.set_image(image_array)\n                x_coords = [coord['x'] for coord in coordinates]\n                y_coords = [coord['y'] for coord in coordinates]\n                xmin, xmax = int(min(x_coords)), int(max(x_coords))\n                ymin, ymax = int(min(y_coords)), int(max(y_coords))\n\n                width = xmax - xmin\n                height = ymax - ymin\n\n                # Expand bounding box by 10%\n                xmin = max(0, int(xmin - 0.2 * width))\n                xmax = min(image_array.shape[1] - 0.2, int(xmax + 1 * width))\n                ymin = max(0, int(ymin - 0.2 * height))\n                ymax = min(image_array.shape[0] - 0.2, int(ymax + 1 * height))\n\n                input_box = [xmin, ymin, xmax, ymax]\n\n                masks, scores, logits = predictor.predict(box=input_box, multimask_output=False)\n\n                # Select the highest scoring mask for this slice\n                best_mask_index = np.argmax(scores)\n                #best_mask_index = np.argsort(scores)[len(scores) // 2]  # Select the index of the median score\n                best_mask = masks[best_mask_index]\n                best_score = scores[best_mask_index]\n\n                mask_scores.append((best_mask, best_score))\n\n           # Find the best mask across all slices\n            # Find the best mask across all slices\n        if mask_scores:\n            best_mask_overall = max(mask_scores, key=lambda x: x[1])[0]  # Mask with the highest score\n\n            for i in range(len(slices)):\n                best_mask_for_slice = mask_scores[i][0]  # Get the best mask for the current slice\n                \n                try:\n                    if best_mask_for_slice.shape != best_mask_overall.shape:\n                        best_mask_for_slice = resize(best_mask_for_slice, best_mask_overall.shape, preserve_range=True, anti_aliasing=True)\n\n                    current_mask = np.clip(best_mask_for_slice, best_mask_overall * 0.4, best_mask_overall * 1.6)\n                except Exception as e:\n                    print(f\"Error processing mask at index {idx}: {e}\")\n                    return None, None\n                # Now clip the mask safely\n                current_mask = np.clip(best_mask_for_slice, best_mask_overall * 0.4, best_mask_overall * 1.6)\n\n                current_mask = np.where(current_mask > 0.5, 1, 0)  # Ensure binary mask (0 or 1)\n\n                # Ensure the slice is in grayscale by averaging over RGB channels\n                slices[i] = np.mean(slices[i], axis=-1).astype(np.uint8)  # Convert (H, W, 3) to (H, W)\n\n                # Extract the feature using the best mask for the current slice\n                feature_region = slices[i] * current_mask\n                feature_region = self.apply_gaussian_smoothing(feature_region, sigma=0.5)\n                #feature_region = self.fill_structures_in_volume(feature_region)\n\n                # Create a dimmed version of the original image for the background\n                dimmed_image = (slices[i] * 0.3).astype(np.uint8)  # Adjust the factor to control the dimming\n\n                # Combine the feature and the dimmed image\n                combined_image = np.where(current_mask > 0, feature_region * 2.5, dimmed_image).astype(np.uint8)\n\n                # Clip to ensure values are in the [0, 255] range\n                combined_image = np.clip(combined_image, 0, 255)\n\n                # Update the current slice with the combined image\n                slices[i] = combined_image\n\n        # Stack the slices into a volume and continue with processing\n        if slices:\n            volume = np.stack(slices, axis=0)  # Now (D, H, W)\n\n            # Calculate target_spacing\n            target_spacing = self.calculate_target_spacing(slice_thickness, pixel_spacing, method='original_resolution')\n\n            # Define original spacing from DICOM metadata\n            original_spacing = (slice_thickness, pixel_spacing[0], pixel_spacing[1])\n\n            # Resample volume using target_spacing and finally to target_shape\n            volume = self.resample_volume_to_target_shape(volume)\n\n            # Further processing\n            volume = self.apply_gaussian_smoothing(volume, sigma=0.5)\n            volume = self.apply_sharpening(volume)\n            volume = self.fill_structures_in_volume(volume)\n\n            return volume, pixel_spacing, slice_thickness\n        else:\n            return None, None, None\n\n\n\n\n\n\n\n\n    def apply_gaussian_smoothing(self, volume, sigma=1):\n        smoothed_volume = gaussian_filter(volume, sigma=sigma)\n        return smoothed_volume\n\n    def apply_sharpening(self, volume):\n        sharpened_volume = np.zeros_like(volume)\n        for i in range(volume.shape[0]):  # Apply sharpening slice by slice\n            blurred = cv2.GaussianBlur(volume[i], (0, 0), sigmaX=0.8)\n            sharpened = cv2.addWeighted(volume[i], 2.0, blurred, -1.0, 0)\n            sharpened_volume[i] = np.clip(sharpened, 0, 255)  # Keep values in [0, 255]\n        return sharpened_volume\n\n    def resample_volume_to_target_shape(self, volume):\n        # Check the shape and remove any singleton dimensions\n        if volume.ndim == 4 and volume.shape[-1] == 1:\n            volume = np.squeeze(volume, axis=-1)  # Remove the singleton dimension\n\n        # Ensure volume shape is (D, H, W)\n        if volume.ndim != 3:\n            raise ValueError(f\"Expected 3D volume data but got {volume.ndim}D data.\")\n\n        # Resample the volume to this intermediate shape\n        volume_tensor = torch.from_numpy(volume).unsqueeze(0).unsqueeze(0).float()  # Convert to float32\n        resampled_volume_tensor = F.interpolate(volume_tensor, size=tuple(self.target_shape), mode='trilinear', align_corners=True)\n\n\n        final_resampled_volume = resampled_volume_tensor.squeeze().numpy()  # Remove the batch and channel dimensions\n\n        return final_resampled_volume\n\n\n\n\n    def extract_and_adjust_coordinates(self, row, original_shape, original_spacing, resampled_spacing, coordinates):\n        #coordinates = row['Coordinates']\n        #print(coordinates)\n        adjusted_coords = []\n        target_shape = np.array(self.target_shape, dtype=np.float32)\n        original_shape = np.array([self.o_shape[0], self.o_shape[1], self.o_shape[2]], dtype=np.float32)\n\n        # Calculate scaling factors based on spacing changes\n        scale = (target_shape / original_shape)\n    \n        #print(\"Scaling factors:\", scale[1])\n\n        for coord in coordinates:\n            x, y = coord['x'], coord['y']\n            # Use original z if available, otherwise set to mid-point\n            label = coord['label']\n            level = coord.get('level', None)\n            severity = coord['severity']\n            instance = coord['instance']\n            #print(label)\n            #print(level)\n            #print(instance)\n            #print(severity)\n\n            label_level_pair = (label, level, severity)\n            \n            #print(label_level_pair)\n            class_id = self.class_map.get(label_level_pair, -1)\n\n            # Adjust the coordinates according to the scaling factors\n            adjusted_x = x * scale[1]\n            adjusted_y = y * scale[2]\n\n            # Ensure the adjusted coordinates stay within the bounds of the target shape\n            #adjusted_x = np.clip(adjusted_x, 255, self.target_shape[2] - 1)\n            #adjusted_y = np.clip(adjusted_y, 255, self.target_shape[1] - 1)\n            adjusted_z = ((int(instance))*(self.target_shape[0]/self.o_shape[0]))\n            #print('z', adjusted_z)\n\n            adjusted_coords.append(((adjusted_x, adjusted_y), class_id, adjusted_z))\n\n        # Ensure at least 6 adjusted coordinates\n        while len(adjusted_coords) < 6:\n            adjusted_coords.append(((0.0, 0.0), -1, 0))\n\n        adjusted_coords_tensor = torch.tensor([list(c[0]) + [c[2], c[1]] for c in adjusted_coords], dtype=torch.float32)\n\n        return adjusted_coords_tensor\n\n\n    def save_test_volume_and_coords(self, volume_tensor, coords_tensor, volume_filename='test_volume.pt', coords_filename='test_coords.pt'):\n        torch.save(volume_tensor, os.path.join(self.save_path, volume_filename))\n        torch.save(coords_tensor, os.path.join(self.save_path, coords_filename))\n\n    def __len__(self):\n        return min(len(self.series_ids), self.max_samples)\n\n    def __getitem__(self, idx):\n        try:\n            series_id = self.series_ids[idx]\n            rows = self.series_dict[series_id]\n            file_paths_list = rows[0]['FilePaths']\n            coordinates = rows[0]['Coordinates']\n\n            # Get the first DICOM file to read pixel spacing and slice thickness\n            if len(file_paths_list) > 0:\n                dicom_data = pydicom.dcmread(file_paths_list[0])\n                current_pixel_spacing = dicom_data.PixelSpacing if hasattr(dicom_data, 'PixelSpacing') else [1.0, 1.0]\n                current_slice_thickness = float(dicom_data.SliceThickness) if hasattr(dicom_data, 'SliceThickness') else 1.0\n\n                # Stack and resample DICOM slices\n                volume_data, pixel_spacing, slice_thickness = self.stack_and_resample_dicom_slices(\n                    file_paths_list, coordinates, current_pixel_spacing, current_slice_thickness\n                )\n\n                if volume_data is None:\n                    return torch.tensor([]), torch.tensor([]), torch.tensor([])  # Return empty tensors\n\n                # Fix axis order and normalize volume data\n                volume_data = self.check_and_fix_axis_order(volume_data)\n                volume_tensor = torch.from_numpy(volume_data).unsqueeze(0).float()\n                volume_tensor = (volume_tensor - volume_tensor.min()) / (volume_tensor.max() - volume_tensor.min())\n\n                # Extract original shape and spacing\n                target_shape = np.array(self.target_shape, dtype=np.float32)\n                original_shape = np.array([self.o_shape[0], self.o_shape[1], self.o_shape[2]], dtype=np.float32)\n                original_spacing = np.array([slice_thickness, pixel_spacing[0], pixel_spacing[1]], dtype=np.float32)\n                resampled_spacing = (original_shape * original_spacing) / target_shape\n\n                # Adjust coordinates\n                adjusted_coords_tensor = self.extract_and_adjust_coordinates(\n                    rows[0], self.o_shape, self.o_space, resampled_spacing, coordinates\n                )\n\n                resampled_spacing = torch.tensor(resampled_spacing, dtype=torch.float32)\n\n                return volume_tensor, adjusted_coords_tensor, resampled_spacing\n            else:\n                return torch.tensor([]), torch.tensor([]), torch.tensor([])\n\n        except Exception as e:\n            print(f\"Error processing data at index {idx}: {e}\")\n            return torch.tensor([]), torch.tensor([]), torch.tensor([])\n\n    def get_z_depth(self, file_paths_list):\n        # The number of slices corresponds to the number of DICOM files\n        z_depth = len(file_paths_list)\n        return z_depth\n        \n    def apply_contouring(self, volume):\n        \"\"\"\n        Apply contouring to the 3D volume using OpenCV's findContours function.\n        Args:\n            volume (np.ndarray): The input 3D volume.\n        Returns:\n            np.ndarray: The volume with contours drawn on each slice.\n        \"\"\"\n        contoured_volume = np.zeros_like(volume)\n\n        for i in range(volume.shape[0]):  # Process slice by slice\n            slice_data = volume[i]\n\n            # Normalize the slice to [0, 255] for contour detection\n            slice_uint8 = cv2.normalize(slice_data, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n\n            # Apply thresholding to create a binary mask\n            _, binary_mask = cv2.threshold(slice_uint8, 50, 255, cv2.THRESH_BINARY)\n\n            # Find contours\n            contours, _ = cv2.findContours(binary_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n\n            # Create a blank image for drawing contours\n            contour_image = np.zeros_like(slice_uint8)\n\n            # Draw contours on the blank image\n            cv2.drawContours(contour_image, contours, -1, (255), 1)  # Draw contours in white (255)\n\n            # Add the contoured image back into the volume\n            contoured_volume[i] = contour_image\n\n        return contoured_volume\n\n\n    def check_and_fix_axis_order(self, volume_data):\n        if volume_data.shape[0] < volume_data.shape[1] and volume_data.shape[0] < volume_data.shape[2]:\n            return volume_data\n        elif volume_data.shape[-1] < volume_data.shape[0]:\n            return np.transpose(volume_data, (2, 0, 1))\n        else:\n            return np.moveaxis(volume_data, [0, 1, 2], [2, 1, 0])","metadata":{"execution":{"iopub.status.busy":"2024-10-06T09:21:51.331803Z","iopub.execute_input":"2024-10-06T09:21:51.332438Z","iopub.status.idle":"2024-10-06T09:21:51.745667Z","shell.execute_reply.started":"2024-10-06T09:21:51.332401Z","shell.execute_reply":"2024-10-06T09:21:51.744863Z"}}},{"cell_type":"markdown","source":"Dataset Loaders Train/Val","metadata":{}},{"cell_type":"markdown","source":"Single dataloader with saved volume for validation","metadata":{}},{"cell_type":"markdown","source":"from torch.utils.data import DataLoader, random_split\n\n# Assuming `STIR_Data or T1_Data or T2_Data` is your dataframe\ndataset = DICOMDataset(T2_Data, class_map, target_shape=(50, 320, 320))\n\n# Define split ratio\ntrain_ratio = 0.8\nval_ratio = 1 - train_ratio\n\n# Calculate lengths for training and validation sets\ntotal_len = len(dataset)\ntrain_len = int(train_ratio * total_len)\nval_len = total_len - train_len\n\n# Split the dataset into training and validation sets\ntrain_dataset, val_dataset = random_split(dataset, [train_len, val_len])\n\n# Create DataLoaders for training and validation\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True, drop_last=True)\nval_loader = DataLoader(val_dataset, batch_size=1, shuffle=False, drop_last=True)\n\nprint('Train_loader Size', len(train_loader))\nprint('Val_loader Size', len(val_loader))\n\ndef save_volume_and_coordinates(volume_tensor, coords_tensor, resampled_spacing, save_dir='./kaggle/working/', volume_filename='saved_volume_sam.pt', coords_filename='saved_coords.pt', spacing_filename='saved_spacing.pt'):\n    os.makedirs(save_dir, exist_ok=True)  # Create the directory if it doesn't exist\n    torch.save(volume_tensor, os.path.join(save_dir, volume_filename))\n    torch.save(coords_tensor, os.path.join(save_dir, coords_filename))\n    torch.save(resampled_spacing, os.path.join(save_dir, spacing_filename))\n    print(f'Saved volume to {os.path.join(save_dir, volume_filename)}')\n    print(f'Saved coordinates to {os.path.join(save_dir, coords_filename)}')\n    print(f'Saved spacing to {os.path.join(save_dir, spacing_filename)}')\n\n# Now iterate over the train_loader and save the volume, targets, and spacing\nfor i, (volumes, targets, resampled_spacing) in enumerate(train_loader):\n    print(f'Batch {i + 1}:')\n    print('  Volumes shape:', volumes.shape)  # Shape of volumes\n    print('  Volumes max:', torch.max(volumes))\n    print('  Targets:', targets)  # Displaying targets (list of coordinates)\n    print('  Resampled spacing:', resampled_spacing)  # Displaying the resampled spacing\n\n    # Save the first volume, coordinates, and spacing to a file\n    if i == 0:  # Change the condition as needed to select which batch to save\n        save_volume_and_coordinates(volumes, targets, resampled_spacing, save_dir='/kaggle/working/saved_data', volume_filename=f'saved_volume_sam.pt', coords_filename=f'saved_coords_{i+1}.pt', spacing_filename=f'saved_spacing_{i+1}.pt')\n\n    # Break after the first batch to avoid excessive output (optional)\n    if i == 1:\n        break\n\n\n# Example usage of the loaders\n#for volume, targets in train_loader:\n #   print(\"Training batch:\")\n  #  print('Volume', volume.shape)  # Should print the shape of the batched volumes\n   # print('Targets', targets)  # List of adjusted coordinates and labels\n\n#for volume, targets in val_loader:\n  #  print(\"Validation batch:\")\n   # print(volume.shape)  # Should print the shape of the batched volumes\n    #print(targets)  # List of adjusted coordinates and labels\n","metadata":{"execution":{"iopub.status.busy":"2024-10-06T09:21:51.747850Z","iopub.execute_input":"2024-10-06T09:21:51.748323Z","iopub.status.idle":"2024-10-06T09:22:38.839792Z","shell.execute_reply.started":"2024-10-06T09:21:51.748288Z","shell.execute_reply":"2024-10-06T09:22:38.838725Z"}}},{"cell_type":"markdown","source":"Single dataset dataloader","metadata":{}},{"cell_type":"markdown","source":"from torch.utils.data import DataLoader, random_split, Subset\nimport random\n\n# Assuming `STIR_Data or T1_Data or T2_Data` is your dataframe\ndataset = DICOMDataset(T2_Data, class_map, target_shape=(20, 320, 320))\n\n# Define sample ratio (20% of total dataset)\nsample_ratio = 0.1\n\n# Calculate sample size (20% of total dataset)\nsample_size = int(len(dataset) * sample_ratio)\n\n# Generate random indices for the sampled subset\nindices = random.sample(range(len(dataset)), sample_size)\n\n# Create a subset using the sampled indices\nsampled_dataset = Subset(dataset, indices)\n\n# Define train/validation split ratio\ntrain_ratio = 0.8\n\n# Calculate lengths for training and validation sets from the sample\ntrain_len = int(sample_size * train_ratio)\nval_len = sample_size - train_len\n\n# Split the sampled dataset into training and validation sets\ntrain_dataset, val_dataset = random_split(sampled_dataset, [train_len, val_len])\n\n# Create DataLoaders for training and validation\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True, drop_last=True)\nval_loader = DataLoader(val_dataset, batch_size=1, shuffle=False, drop_last=True)\n\n# Print sizes\nprint(f\"Total dataset size: {len(dataset)}\")\nprint(f\"Sample size (20% of total): {sample_size}\")\nprint(f\"Training dataset size: {train_len}\")\nprint(f\"Validation dataset size: {val_len}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-06T09:22:38.841302Z","iopub.execute_input":"2024-10-06T09:22:38.841981Z","iopub.status.idle":"2024-10-06T09:22:41.867967Z","shell.execute_reply.started":"2024-10-06T09:22:38.841916Z","shell.execute_reply":"2024-10-06T09:22:41.866975Z"}}},{"cell_type":"markdown","source":"3DUNET ","metadata":{}},{"cell_type":"markdown","source":"Model for inference remember to add Sigmoid return","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass DoubleConv(nn.Module):\n    \"\"\"(convolution => [BN] => ReLU) * 2\"\"\"\n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        self.double_conv = nn.Sequential(\n            nn.Conv3d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n\nclass Down(nn.Module):\n    \"\"\"Downscaling with maxpool then double conv\"\"\"\n    def __init__(self, in_channels, out_channels):\n        super(Down, self).__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool3d(2),\n            DoubleConv(in_channels, out_channels)\n        )\n\n    def forward(self, x):\n        return self.maxpool_conv(x)\n\nclass Up(nn.Module):\n    \"\"\"Upscaling then double conv\"\"\"\n    def __init__(self, in_channels, out_channels):\n        super(Up, self).__init__()\n        self.up = nn.ConvTranspose3d(in_channels, out_channels, kernel_size=2, stride=2)\n        self.conv = DoubleConv(in_channels, out_channels)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        # Input is CHWD (Channel, Height, Width, Depth), we need to pad to match dimensions\n        diffD = x2.size(2) - x1.size(2)\n        diffH = x2.size(3) - x1.size(3)\n        diffW = x2.size(4) - x1.size(4)\n        x1 = F.pad(x1, [diffW // 2, diffW - diffW // 2,\n                        diffH // 2, diffH - diffH // 2,\n                        diffD // 2, diffD - diffD // 2])\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\nclass OutConv(nn.Module):\n    \"\"\"Output convolution to reduce the number of channels to the desired output\"\"\"\n    def __init__(self, in_channels, out_channels):\n        super(OutConv, self).__init__()\n        self.conv = nn.Conv3d(in_channels, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        return self.conv(x)\n\nclass UNet3DClassifier(nn.Module):\n    def __init__(self, n_channels, n_classes):\n        super(UNet3DClassifier, self).__init__()\n        self.in_conv = DoubleConv(n_channels, 64)\n        self.down1 = Down(64, 128)\n        self.down2 = Down(128, 256)\n        self.down3 = Down(256, 512)\n        self.down4 = Down(512, 1024)\n        self.up1 = Up(1024, 512)\n        self.up2 = Up(512, 256)\n        self.up3 = Up(256, 128)\n        self.up4 = Up(128, 64)\n        self.out_conv = OutConv(64, n_classes)\n\n    def forward(self, x):\n        x1 = self.in_conv(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        x = self.up1(x5, x4)\n        x = self.up2(x, x3)\n        x = self.up3(x, x2)\n        x = self.up4(x, x1)\n        x = self.out_conv(x)  # Output is raw logits\n        return x \n        #return torch.sigmoid(x)  # Apply sigmoid for multi-label classification\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T10:56:05.081299Z","iopub.execute_input":"2024-10-08T10:56:05.081683Z","iopub.status.idle":"2024-10-08T10:56:07.826902Z","shell.execute_reply.started":"2024-10-08T10:56:05.081646Z","shell.execute_reply":"2024-10-08T10:56:07.826027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T10:56:07.828504Z","iopub.execute_input":"2024-10-08T10:56:07.829103Z","iopub.status.idle":"2024-10-08T10:56:07.833883Z","shell.execute_reply.started":"2024-10-08T10:56:07.829045Z","shell.execute_reply":"2024-10-08T10:56:07.832924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"import random\nfrom torch.utils.data import DataLoader, random_split, Subset\n\n# Define the sample ratio (for example, 10% of the total dataset)\nsample_ratio = 0.1  # 10% of the dataset\n\n# Function to sample a subset of the dataset\ndef sample_dataset(dataset, sample_ratio):\n    sample_size = int(len(dataset) * sample_ratio)\n    indices = random.sample(range(len(dataset)), sample_size)\n    return Subset(dataset, indices)\n\n# Create individual datasets for STIR, T1, and T2\nstir_dataset = DICOMDataset(STIR_Data, class_map, target_shape=(20, 320, 320))\nt1_dataset = DICOMDataset(T1_Data, class_map, target_shape=(20, 320, 320))\nt2_dataset = DICOMDataset(T2_Data, class_map, target_shape=(20, 320, 320))\n\n# Sample a subset from each dataset (10% in this case)\nstir_subset = sample_dataset(stir_dataset, sample_ratio)\nt1_subset = sample_dataset(t1_dataset, sample_ratio)\nt2_subset = sample_dataset(t2_dataset, sample_ratio)\n\n# Define train/val split ratio\ntrain_ratio = 0.8\n\n# Split each sampled subset into training and validation sets\ntrain_len_stir = int(len(stir_subset) * train_ratio)\nval_len_stir = len(stir_subset) - train_len_stir\ntrain_stir, val_stir = random_split(stir_subset, [train_len_stir, val_len_stir])\n\ntrain_len_t1 = int(len(t1_subset) * train_ratio)\nval_len_t1 = len(t1_subset) - train_len_t1\ntrain_t1, val_t1 = random_split(t1_subset, [train_len_t1, val_len_t1])\n\ntrain_len_t2 = int(len(t2_subset) * train_ratio)\nval_len_t2 = len(t2_subset) - train_len_t2\ntrain_t2, val_t2 = random_split(t2_subset, [train_len_t2, val_len_t2])\n\n# Create DataLoaders for each dataset\ntrain_loader_stir = DataLoader(train_stir, batch_size=1, shuffle=True, drop_last=True)\ntrain_loader_t1 = DataLoader(train_t1, batch_size=1, shuffle=True, drop_last=True)\ntrain_loader_t2 = DataLoader(train_t2, batch_size=1, shuffle=True, drop_last=True)\n\n# Create validation DataLoaders\nval_loader_stir = DataLoader(val_stir, batch_size=1, shuffle=False, drop_last=True)\nval_loader_t1 = DataLoader(val_t1, batch_size=1, shuffle=False, drop_last=True)\nval_loader_t2 = DataLoader(val_t2, batch_size=1, shuffle=False, drop_last=True)\n\n# Print dataset sizes to verify\nprint(f\"STIR subset size: {len(stir_subset)}\")\nprint(f\"T1 subset size: {len(t1_subset)}\")\nprint(f\"T2 subset size: {len(t2_subset)}\")\n\n","metadata":{"execution":{"iopub.status.busy":"2024-10-06T09:22:41.934379Z","iopub.execute_input":"2024-10-06T09:22:41.934725Z","iopub.status.idle":"2024-10-06T09:22:49.538016Z","shell.execute_reply.started":"2024-10-06T09:22:41.934685Z","shell.execute_reply":"2024-10-06T09:22:49.536972Z"}}},{"cell_type":"markdown","source":"Concat Dataloaders sphere 20 and learning rate e3 10% of dataset","metadata":{}},{"cell_type":"markdown","source":"import torch\nimport torch.optim as optim\nimport torch.nn as nn\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.utils.data import ConcatDataset, DataLoader, random_split\nimport os\n\n# Set environment variable to help with memory fragmentation\nos.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'expandable_segments:True'\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nn_channels = 1  # Input channels (e.g., grayscale volumes)\nn_classes = len(class_map)  # Number of classes in your classification task\nmodel = UNet3DClassifier(n_channels, n_classes).to(device)\n\n# Use BCEWithLogitsLoss for multi-label classification\ncriterion = nn.BCEWithLogitsLoss()  # Binary cross-entropy with logits for multi-label classification\noptimizer = optim.AdamW(model.parameters(), lr=1e-3)  # Using AdamW optimizer\nscaler = GradScaler()  # For mixed precision training\n\nnum_epochs = 3  # Define the number of epochs\naccumulation_steps = 8  # Number of batches to accumulate before an optimization step\nsphere_radius = 25  # Radius of the spherical mask in voxels\ncheckpoint_interval = 500  # Interval for saving checkpoints (in batches)\n\n# Function for spherical mask creation\ndef create_spherical_mask(shape, center, radius):\n    z, y, x = torch.meshgrid(\n        torch.arange(shape[0], device=device),\n        torch.arange(shape[1], device=device),\n        torch.arange(shape[2], device=device),\n        indexing='ij'  # Use 'ij' indexing to align with (D, H, W) dimensions\n    )\n    z = z - center[0]\n    y = y - center[1]\n    x = x - center[2]\n    mask = (z ** 2 + y ** 2 + x ** 2) <= radius ** 2\n    return mask.float()\n\ndef validate(model, val_loader, max_batches=4):\n    model.eval()  # Set model to evaluation mode\n    val_loss = 0.0\n    num_batches = 0\n    with torch.no_grad():  # No need to compute gradients during validation\n        for i, (volumes, targets, _) in enumerate(val_loader):\n            if i >= max_batches:  # Stop validation after processing 'max_batches' batches\n                break\n            if volumes.nelement() == 0 or targets.nelement() == 0:\n                print(f\"Skipping batch {i} due to invalid data (empty tensors).\")\n                continue  # Skip this batch\n\n            volumes = volumes.to(device)\n            targets = targets.to(device)\n\n            # Forward pass\n            with autocast():  # Enable mixed precision for validation\n                outputs = model(volumes)  # Raw logits, no sigmoid\n\n                batch_size, n_rois, _ = targets.shape\n                roi_predictions = []\n                valid_class_ids = []\n\n                for b in range(batch_size):\n                    valid_mask = targets[b, :, 3] != -1  # Boolean mask for valid ROIs\n                    coords = targets[b, valid_mask, :3].long()  # Extract valid ROI coordinates\n                    class_ids = targets[b, valid_mask, 3].long()  # Extract valid class IDs\n\n                    if coords.shape[0] == 0:\n                        continue  # Skip if no valid ROIs\n\n                    # Create spherical masks and extract ROI predictions\n                    for idx, coord in enumerate(coords):\n                        x, y, z = coord\n                        mask = create_spherical_mask(outputs.shape[2:], (z, y, x), sphere_radius)\n\n                        # Apply mask to the output volume\n                        roi_prediction = outputs[b, :, mask.bool()].mean(dim=1)  # Average the predictions within the mask\n                        roi_predictions.append(roi_prediction)\n\n                        # Convert class_ids to one-hot encoded format\n                        one_hot_class_ids = F.one_hot(class_ids[idx], num_classes=n_classes).float()  # `n_classes` = 75\n                        valid_class_ids.append(one_hot_class_ids)\n\n                if roi_predictions:\n                    roi_predictions = torch.stack(roi_predictions, dim=0)  # Shape: [n_valid_rois, n_classes]\n                    class_ids = torch.stack(valid_class_ids, dim=0)  # Shape: [n_valid_rois, n_classes]\n\n                    if roi_predictions.size(0) > 0:\n                        # Use BCEWithLogitsLoss for validation loss\n                        loss = criterion(roi_predictions, class_ids)\n                        val_loss += loss.item()\n                        num_batches += 1\n\n    # Average validation loss\n    return val_loss / num_batches if num_batches > 0 else 0.0\n\n# Function to calculate the total number of batches from the combined DataLoaders\ndef get_total_batches(loaders):\n    total_batches = sum(len(loader) for loader in loaders)\n    return total_batches\n\n# Function to alternate between the DataLoaders for different datasets\ndef alternate_dataloaders(loaders):\n    while True:\n        for loader in loaders:\n            for data in loader:\n                # Skip batches where any part of the data is None\n                if any(d.nelement() == 0 for d in data):\n                    print(\"Skipping batch due to empty tensors in data\")\n                    continue\n\n                yield data\n\n# Combine train and validation datasets using ConcatDataset\ntrain_loaders = [train_loader_stir, train_loader_t1, train_loader_t2]  # Add your loaders here\nval_loaders = [val_loader_stir, val_loader_t1, val_loader_t2]\n\n# Create iterators for alternating DataLoaders\ntrain_loader = alternate_dataloaders(train_loaders)\nval_loader = alternate_dataloaders(val_loaders)\n\n# Calculate total batches for the training loop\ntotal_batches = get_total_batches(train_loaders)\n\n# Training loop\nscheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.1)\n\n# Updated Training loop\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    num_batches = 0\n    optimizer.zero_grad()\n\n    # Loop through all batches with enumerate\n    for i, data in enumerate(train_loader):\n        if i >= total_batches:\n            break  # Stop after all batches in the dataset have been processed\n\n        volumes, targets, _ = data  # Get batch from alternated DataLoader\n        # Check for invalid data (None values)\n        if volumes.nelement() == 0 or targets.nelement() == 0:\n            print(f\"Skipping batch {i} due to invalid data (empty tensors).\")\n            continue  # Skip this batch\n\n        volumes = volumes.to(device)\n        targets = targets.to(device)\n\n        with autocast():  # Enable mixed precision\n            # Forward pass\n            outputs = model(volumes)\n\n            batch_size, n_rois, _ = targets.shape\n            roi_predictions = []\n            valid_class_ids = []\n\n            for b in range(batch_size):\n                valid_mask = targets[b, :, 3] != -1\n                coords = targets[b, valid_mask, :3].long()\n                class_ids = targets[b, valid_mask, 3].long()\n\n                if coords.shape[0] == 0:\n                    continue\n\n                # Create spherical masks and extract ROI predictions\n                for idx, coord in enumerate(coords):\n                    x, y, z = coord\n                    mask = create_spherical_mask(outputs.shape[2:], (z, y, x), sphere_radius)\n\n                    # Apply mask to the output volume\n                    roi_prediction = outputs[b, :, mask.bool()].mean(dim=1)  # Average the predictions within the mask\n                    roi_predictions.append(roi_prediction)\n\n                    # Convert class_ids to one-hot encoded format\n                    one_hot_class_ids = F.one_hot(class_ids[idx], num_classes=n_classes).float()  # `n_classes` = 75\n                    valid_class_ids.append(one_hot_class_ids)\n\n            if roi_predictions:\n                roi_predictions = torch.stack(roi_predictions, dim=0)\n                class_ids = torch.stack(valid_class_ids, dim=0)\n\n                if roi_predictions.size(0) > 0:\n                    # Use BCEWithLogitsLoss for training\n                    loss = criterion(roi_predictions, class_ids) / accumulation_steps\n                    scaler.scale(loss).backward()\n                    running_loss += loss.item()\n                    num_batches += 1\n\n                    # Perform optimization step every `accumulation_steps` batches\n                    if (i + 1) % accumulation_steps == 0:\n                        scaler.step(optimizer)\n                        scaler.update()\n                        optimizer.zero_grad()\n\n                        # Print the accumulated loss\n                        print(f'Epoch [{epoch + 1}/{num_epochs}], Batch [{i + 1}], Loss: {running_loss:.4f}')\n                        running_loss = 0.0\n\n                        # Save checkpoint every `checkpoint_interval` batches\n                        # Save checkpoint when you're just past the checkpoint interval\n                        if ((i + 1) * accumulation_steps) >= checkpoint_interval:\n                            checkpoint_path = f'/kaggle/working/model_checkpoint_epoch_{epoch + 1}_batch_{i + 1}.pth'\n                            torch.save(model.state_dict(), checkpoint_path)\n                            print(f\"Checkpoint saved at {checkpoint_path}\")\n                            checkpoint_interval += checkpoint_interval  # Adjust to the next checkpoint interval\n                            # Validate after saving checkpoint\n                            val_loss = validate(model, val_loader, max_batches=4)\n                            print(f\"Validation Loss after batch {i + 1}: {val_loss:.4f}\")\n\n        # Free up memory after each batch\n        del volumes, targets, outputs, roi_predictions, class_ids\n        torch.cuda.empty_cache()\n\n    # Final optimization step for remaining gradients\n    if (i + 1) % accumulation_steps != 0:\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad()\n\n    # Validate at the end of each epoch\n    val_loss = validate(model, val_loader, _)\n    print(f'Validation Loss after epoch {epoch + 1}: {val_loss:.4f}')\n\n    # Adjust learning rate\n    scheduler.step()\n\n    # Save model checkpoint at the end of each epoch\n    torch.save(model.state_dict(), f'/kaggle/working/model_epoch_final_{epoch + 1}.pth')\n    print(f\"Model saved after epoch {epoch + 1}\")\n\nprint('Training finished!')\n","metadata":{"execution":{"iopub.status.busy":"2024-10-06T09:22:49.645681Z","iopub.execute_input":"2024-10-06T09:22:49.646364Z","iopub.status.idle":"2024-10-06T15:18:59.304931Z","shell.execute_reply.started":"2024-10-06T09:22:49.646319Z","shell.execute_reply":"2024-10-06T15:18:59.303349Z"}}},{"cell_type":"markdown","source":"Inference predictions test data","metadata":{}},{"cell_type":"code","source":"import os\nimport pydicom\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\nfrom scipy.ndimage import gaussian_filter\nfrom skimage.transform import resize\nfrom PIL import Image\nimport cv2\n\nclass DICOMTestDataset(Dataset):\n    def __init__(self, root_dir, target_shape=(20, 320, 320)):\n        self.root_dir = root_dir\n        self.target_shape = target_shape\n        self.series_groups = self.group_series_by_id()\n        self.current_study_id = None\n\n    def group_series_by_id(self):\n        study_series_dict = {}\n        \n        # Iterate through the study_id directories\n        for study_id in sorted(os.listdir(self.root_dir)):\n            study_path = os.path.join(self.root_dir, study_id)\n            if os.path.isdir(study_path):\n                series_ids = []\n                \n                # Iterate through the series_id directories within each study_id\n                for series_id in sorted(os.listdir(study_path)):\n                    series_path = os.path.join(study_path, series_id)\n                    if os.path.isdir(series_path):\n                        series_ids.append(series_id)\n                \n                if series_ids:\n                    study_series_dict[study_id] = series_ids\n        \n        return study_series_dict\n\n    def load_dicom_with_pil(self, dicom_file):\n        try:\n            dicom_data = pydicom.dcmread(dicom_file)\n            image_array = dicom_data.pixel_array.astype(np.float32)\n            # Normalize the image to [0, 255]\n            image_array = (image_array - image_array.min()) / (image_array.max() - image_array.min()) * 255\n            image_array = image_array.astype(np.uint8)\n            return image_array\n        except Exception as e:\n            print(f\"Error loading DICOM file {dicom_file}: {e}\")\n            return None\n\n    def stack_and_resample_dicom_slices(self, file_paths_list, z_spacing=3.5):\n        slices = []\n        \n        for file_path in file_paths_list:\n            image_array = self.load_dicom_with_pil(file_path)\n            if image_array is None:\n                continue\n\n            # Resize the image_array to the target shape (height and width)\n            image_array = resize(image_array, self.target_shape[1:], mode='constant', preserve_range=True)\n            slices.append(image_array)\n\n        if not slices:\n            return None\n\n        # Stack slices to form a 3D volume (D, H, W)\n        volume = np.stack(slices, axis=0)\n\n        # Adjust the depth (z-dimension) for the slice spacing\n        target_shape_with_z_spacing = (\n            #int(volume.shape[0] * z_spacing),  # Adjust depth based on z_spacing\n            self.target_shape[0],\n            self.target_shape[1],              # Height remains the same\n            self.target_shape[2]               # Width remains the same\n        )\n\n        # Resample the volume to the target shape\n        volume_tensor = torch.from_numpy(volume).unsqueeze(0).unsqueeze(0).float()\n\n        resampled_volume_tensor = torch.nn.functional.interpolate(\n            volume_tensor, size=target_shape_with_z_spacing, mode='trilinear', align_corners=True\n        )\n\n        final_resampled_volume = resampled_volume_tensor.squeeze().numpy()\n        final_resampled_volume = (final_resampled_volume - final_resampled_volume.min()) / (final_resampled_volume.max() - final_resampled_volume.min())\n\n        return final_resampled_volume\n\n    def __len__(self):\n        return len(self.series_groups)\n\n    def __getitem__(self, idx):\n        self.current_study_id = list(self.series_groups.keys())[idx]\n        series_ids = self.series_groups[self.current_study_id]\n        print(f\"Study {self.current_study_id} has series: {series_ids}\")\n\n        volumes = []\n\n        # Iterate over all series_ids for the given study_id\n        for series_id in series_ids:\n            file_paths_list = []\n            series_dir = os.path.join(self.root_dir, str(self.current_study_id), str(series_id))\n            file_paths_list += [os.path.join(series_dir, f) for f in sorted(os.listdir(series_dir)) if f.endswith('.dcm')]\n\n            if len(file_paths_list) == 0:\n                print(f\"No DICOM files found in {series_dir}. Skipping.\")\n                continue\n\n            volume = self.stack_and_resample_dicom_slices(file_paths_list, z_spacing=2)\n\n            if volume is None:\n                print(f\"Failed to load volume for series {series_id}. Skipping.\")\n                continue\n\n            # Apply post-processing\n            volume = self.apply_postprocessing(volume)\n\n            # Convert the volume to a torch tensor\n            volume_tensor = torch.from_numpy(volume).unsqueeze(0).float()\n\n            # Save the volume to a file in /kaggle/working/\n            #save_path = f'/kaggle/working/{self.current_study_id}_{series_id}_volume.pt'\n            #torch.save(volume_tensor, save_path)\n            #print(f\"Volume for study {self.current_study_id}, series {series_id} saved to {save_path}\")\n\n            volumes.append(volume_tensor)\n\n        return volumes\n\n    def get_current_study_id(self):\n        return self.current_study_id\n    \n    def apply_postprocessing(self, volume):\n        volume = gaussian_filter(volume, sigma=0.5)\n        return volume\n\n\n# Example Usage:\ncsv_file = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv\"\nroot_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/\"\n\n# Create test dataset\ntest_dataset = DICOMTestDataset(root_dir=root_dir, target_shape=(20, 320, 320))\n\n# Create DataLoader for the test dataset\ntest_loader = torch.utils.data.DataLoader(test_dataset, batch_size=1, shuffle=False)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T10:56:07.836322Z","iopub.execute_input":"2024-10-08T10:56:07.836914Z","iopub.status.idle":"2024-10-08T10:56:08.438689Z","shell.execute_reply.started":"2024-10-08T10:56:07.836865Z","shell.execute_reply":"2024-10-08T10:56:08.437798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(STIR_avg_bbox['avg_min_x'])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T10:56:08.439834Z","iopub.execute_input":"2024-10-08T10:56:08.440276Z","iopub.status.idle":"2024-10-08T10:56:08.445270Z","shell.execute_reply.started":"2024-10-08T10:56:08.440241Z","shell.execute_reply":"2024-10-08T10:56:08.444337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\n# Recreate the model architecture\nn_channels = 1  # Adjust according to your model\nn_classes = 75  # Adjust based on the number of classes in your model\nmodel = UNet3DClassifier(n_channels, n_classes)\n\n# Load the saved weights\n#model.load_state_dict(torch.load('/kaggle/working/model_epoch_final_2.pth'))\nmodel.load_state_dict(torch.load('/kaggle/input/model-checkpoint/model_epoch_final_2.pth', map_location=torch.device('cpu')))\ndef init_weights(m):\n    if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear):\n        nn.init.xavier_uniform_(m.weight)\n\nmodel.apply(init_weights)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# Ensure the model is on the same device as the input data\nmodel = model.to(device)\n\nmodel.eval()\n\nall_predictions = []\n\nwith torch.no_grad():\n    for batch in test_loader:\n        batch_probabilities = []\n        \n        # Each batch contains multiple volumes (3 in your case)\n        for volume in batch:\n            # Pass each volume through the model individually\n            outputs = model(volume.to(device))  # Assuming model and volume are on the same device\n\n            # Perform sigmoid to get probabilities\n            probabilities = torch.sigmoid(outputs)\n\n            # Average the predictions across the spatial volume (depth, height, width)\n            spatial_avg_probabilities = probabilities.mean(dim=(2, 3, 4))  # Averaging over depth, height, width\n\n            # Append the averaged probabilities for each volume to a list\n            batch_probabilities.append(spatial_avg_probabilities)\n\n        # Instead of averaging the predictions, apply softmax across the stacked batch probabilities\n        stacked_probabilities = torch.stack(batch_probabilities)  # Shape [3, 75] if 3 volumes\n\n        # Apply softmax across the batch dimension (dim=0) for each class\n        #aggregated_probabilities = F.softmax(stacked_probabilities, dim=0)\n\n        # Optionally, you can reduce this further by summing the probabilities or keeping them as is\n        all_predictions = torch.mean(stacked_probabilities, dim=0)\n\n# The `all_predictions` list now contains the aggregated predictions for each batch of volumes.\nprint(all_predictions)\nprint(stacked_probabilities)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T10:56:08.446976Z","iopub.execute_input":"2024-10-08T10:56:08.447275Z","iopub.status.idle":"2024-10-08T10:57:57.796278Z","shell.execute_reply.started":"2024-10-08T10:56:08.447243Z","shell.execute_reply":"2024-10-08T10:57:57.795209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_pred = []\n\n# Assuming you already have study_id extracted or available in your data structure\n#study_id = \"44036939\"  # Replace this with dynamic study_id extraction if needed\ncurrent_study_id = test_dataset.get_current_study_id()\nprint(current_study_id)\nstudy_id = current_study_id\n# Step 1: Reverse the class_map to map index to class label\nreverse_class_map = {v: k for k, v in class_map.items()}\n\n# Step 2: Loop over all batches in all_predictions\nfor batch_index, class_probabilities in enumerate(all_predictions):\n    print(f\"Processing batch {batch_index + 1}\")\n\n    # Step 3: Loop over class probabilities within the batch\n    for i, prob in enumerate(class_probabilities.squeeze(0)):  # Squeeze to handle batch dimension\n        class_label = reverse_class_map.get(i, f'Class {i}')  # Get class label from reversed map\n\n        # Add the study_id at the beginning of the label to match the row_id format\n        row_id_label = f\"{study_id}_{class_label[0].lower().replace(' ', '_')}_{class_label[1].lower().replace('/', '_')}\"\n        \n        class_pred.append(f'{row_id_label}: {prob.item():.4f}')\n        print(f'{row_id_label}: {prob.item():.4f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T10:57:57.797848Z","iopub.execute_input":"2024-10-08T10:57:57.798300Z","iopub.status.idle":"2024-10-08T10:57:57.810015Z","shell.execute_reply.started":"2024-10-08T10:57:57.798245Z","shell.execute_reply":"2024-10-08T10:57:57.808755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport torch\n\n# Load sample_submission.csv\nsubmission_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv')\n\n# Assuming `class_pred` is a list of strings in the format \"44036939_right_neural_foraminal_narrowing_l5_s1: 2.0877\"\npred_dict = {}\n\nfor pred in class_pred:\n    try:\n        # Split label and probability\n        label, prob = pred.split(': ')\n        \n        # Split the label into parts\n        label_parts = label.split('_')\n        \n        # Parse label parts\n        study_id = label_parts[0]\n        condition = '_'.join(label_parts[1:-2])  # Combine parts up to the last two for the condition\n        region = label_parts[-2]  # Second last part is the region (e.g., l5)\n        level = label_parts[-1]  # Last part is the level (e.g., s1)\n        \n        # Create row_id format like the one in sample_submission.csv\n        row_id = f\"{study_id}_{condition.lower()}_{region.lower()}_{level.lower()}\"\n\n        # Initialize the dictionary entry for the row_id if it doesn't exist\n        if row_id not in pred_dict:\n            pred_dict[row_id] = {\"normal_mild\": 0, \"moderate\": 0, \"severe\": 0}\n        \n        # Example assumes you are classifying for \"severe\"\n        pred_dict[row_id]['severe'] = float(prob)\n    \n    except ValueError as e:\n        print(f\"Error processing prediction '{pred}': {e}\")\n\n# Apply softmax to normalize each row\ndef apply_softmax(probabilities_dict):\n    \"\"\"Applies softmax to a dictionary of probabilities and returns a dictionary with normalized values.\"\"\"\n    probs = torch.tensor([probabilities_dict['normal_mild'], probabilities_dict['moderate'], probabilities_dict['severe']])\n    softmax_probs = torch.softmax(probs, dim=0)  # Apply softmax\n    return {\n        'normal_mild': softmax_probs[0].item(),\n        'moderate': softmax_probs[1].item(),\n        'severe': softmax_probs[2].item()\n    }\n\n# Loop through your sample_submission and fill in the predicted values from pred_dict\nfor i, row in submission_df.iterrows():\n    row_id = row['row_id']\n    \n    # Check if the row_id exists in pred_dict\n    if row_id in pred_dict:\n        # Normalize the predictions using softmax\n        normalized_probs = apply_softmax(pred_dict[row_id])\n\n        # Populate the submission DataFrame with normalized values\n        submission_df.at[i, 'normal_mild'] = normalized_probs['normal_mild']\n        submission_df.at[i, 'moderate'] = normalized_probs['moderate']\n        submission_df.at[i, 'severe'] = normalized_probs['severe']\n    else:\n        print(f\"Row ID not found in predictions: {row_id}\")\n        # If no prediction exists, set default values (optional)\n        submission_df.at[i, 'normal_mild'] = 0.0\n        submission_df.at[i, 'moderate'] = 0.0\n        submission_df.at[i, 'severe'] = 0.0\n\n# Save the final submission CSV\nsubmission_df.to_csv('/kaggle/working/submission.csv', index=False)\n\nprint(\"Submission file created successfully with softmax normalization!\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T10:57:57.811235Z","iopub.execute_input":"2024-10-08T10:57:57.811917Z","iopub.status.idle":"2024-10-08T10:57:57.846100Z","shell.execute_reply.started":"2024-10-08T10:57:57.811881Z","shell.execute_reply":"2024-10-08T10:57:57.845130Z"},"trusted":true},"execution_count":null,"outputs":[]}]}