{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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"},{"sourceId":9378220,"sourceType":"datasetVersion","datasetId":5688949},{"sourceId":112462,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":94286,"modelId":118511},{"sourceId":112463,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":94287,"modelId":118512}],"dockerImageVersionId":30761,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install torch timm pillow numpy pandas opencv-python albumentations tqdm pydicom scikit-learn","metadata":{"execution":{"iopub.status.busy":"2024-09-13T11:37:33.460575Z","iopub.execute_input":"2024-09-13T11:37:33.460918Z","iopub.status.idle":"2024-09-13T11:38:07.675765Z","shell.execute_reply.started":"2024-09-13T11:37:33.460879Z","shell.execute_reply":"2024-09-13T11:38:07.674673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport timm\nfrom torchvision import transforms\nfrom PIL import Image\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport albumentations as A\nfrom tqdm import tqdm\nimport pydicom\nimport glob\nfrom torch.utils.data import Dataset\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom sklearn.model_selection import train_test_split\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n","metadata":{"execution":{"iopub.status.busy":"2024-09-13T11:38:07.677932Z","iopub.execute_input":"2024-09-13T11:38:07.678748Z","iopub.status.idle":"2024-09-13T11:38:35.679716Z","shell.execute_reply.started":"2024-09-13T11:38:07.678696Z","shell.execute_reply":"2024-09-13T11:38:35.678898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-09-13T11:38:35.680820Z","iopub.execute_input":"2024-09-13T11:38:35.681302Z","iopub.status.idle":"2024-09-13T11:38:35.709226Z","shell.execute_reply.started":"2024-09-13T11:38:35.681267Z","shell.execute_reply":"2024-09-13T11:38:35.708038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the model\nclass LumbarLineModel(nn.Module):\n    def __init__(self, backbone='resnet18'):\n        super().__init__()\n        self.backbone = timm.create_model(backbone, pretrained=False, num_classes=20)\n        \n    def forward(self, x):\n        return torch.sigmoid(self.backbone(x))\n\n# Helper functions\ndef angle_of_line(x1, y1, x2, y2):\n    return np.degrees(np.arctan2(-(y2-y1), x2-x1))\n\ndef crop_between_keypoints(img, keypoint1, keypoint2):\n    h, w = img.shape[:2]\n    x1, y1 = int(keypoint1[0]), int(keypoint1[1])\n    x2, y2 = int(keypoint2[0]), int(keypoint2[1])\n    \n    left = int(min(x1, x2) + (w * 0.1))\n    right = int(max(x1, x2) + (w * 0.1))\n    top = int(min(y1, y2) - (h * 0.05))\n    bottom = int(max(y1, y2) + (h * 0.05))\n            \n    return img[top:bottom, left:right]\n\ndef convert_to_8bit(x):\n    lower, upper = np.percentile(x, (1, 99))\n    x = np.clip(x, lower, upper)\n    x = x - np.min(x)\n    x = x / np.max(x) \n    return (x * 255).astype(\"uint8\")\n\ndef load_dicom_stack(dicom_folder, plane):\n    dicom_files = glob.glob(os.path.join(dicom_folder, \"*.dcm\"))\n    dicoms = [pydicom.dcmread(f) for f in dicom_files]\n    plane = {\"sagittal\": 0, \"coronal\": 1, \"axial\": 2}[plane.lower()]\n    positions = np.asarray([float(d.ImagePositionPatient[plane]) for d in dicoms])\n    idx = np.argsort(positions)\n    array = np.stack([d.pixel_array.astype(\"float32\") for d in dicoms])\n    array = array[idx]\n    return {\"array\": convert_to_8bit(array), \"positions\": positions[idx], \"pixel_spacing\": np.asarray(dicoms[0].PixelSpacing).astype(\"float\")}\n\n# Main function to extract crops using the trained model\ndef extract_crops_with_model(model, image_dir, output_dir, device):\n\n    dfd = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv\")\n    dfd_filtered = dfd[dfd.series_description == \"Sagittal T2/STIR\"]\n    series_ids = dfd_filtered['series_id'].tolist()\n    # Define transforms\n    transform = transforms.Compose([\n        transforms.Resize((512, 512)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ])\n\n    # Resize transform for the crops\n    resize_transform = A.Compose([\n        A.LongestMaxSize(max_size=224, interpolation=cv2.INTER_CUBIC, always_apply=True),\n        A.PadIfNeeded(min_height=224, min_width=224, border_mode=cv2.BORDER_CONSTANT, value=(0, 0, 0), always_apply=True),\n    ])\n\n    # Get all study directories\n    study_dirs = [d for d in os.listdir(image_dir) if os.path.isdir(os.path.join(image_dir, d))]\n\n    for study_id in tqdm(study_dirs, desc=\"Processing studies\"):\n        study_path = os.path.join(image_dir, study_id)\n        series_dirs = [d for d in os.listdir(study_path) if ((os.path.isdir(os.path.join(study_path, d))) and (int(d) in series_ids))]\n\n        for series_id in series_dirs:\n            series_path = os.path.join(study_path, series_id)\n            \n            # Load DICOM stack\n            try:\n                sag_t2 = load_dicom_stack(series_path, plane=\"sagittal\")\n            except Exception as e:\n                print(f\"Error loading DICOM stack for study {study_id}, series {series_id}: {str(e)}\")\n                continue\n            \n            # Select middle slice\n            middle_slice = sag_t2[\"array\"][len(sag_t2[\"array\"])//2]\n            \n            # Resize to 512x512\n            middle_slice_resized = cv2.resize(middle_slice, (512, 512))\n            \n            # Convert to RGB (DICOM images are typically grayscale)\n            img_rgb = cv2.cvtColor(middle_slice_resized, cv2.COLOR_GRAY2RGB)\n            \n            # Convert to PIL Image and apply transform\n            img_pil = Image.fromarray(img_rgb)\n            img_tensor = transform(img_pil).unsqueeze(0).to(device)\n            \n            # Get model prediction\n            model.eval()\n            with torch.no_grad():\n                pred = model(img_tensor).squeeze().cpu().numpy()\n            \n            # Convert prediction to image coordinates\n            h, w = 512, 512  # We resized the image to 512x512\n            pred[0::2] *= w\n            pred[1::2] *= h\n            \n            # Extract crops for each level\n            levels = ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\n            \n            for i, level in enumerate(levels):\n                idx = i * 4\n                left = pred[idx:idx+2]\n                right = pred[idx+2:idx+4]\n                \n                # Rotate\n                rotate_angle = angle_of_line(left[0], left[1], right[0], right[1])\n                augmented = A.Compose([\n                    A.Rotate(limit=(-rotate_angle, -rotate_angle), p=1.0),\n                ], keypoint_params=A.KeypointParams(format='xy', remove_invisible=False))(image=img_rgb, keypoints=[left, right])\n\n                img_rotated = augmented[\"image\"]\n                left_rotated, right_rotated = augmented[\"keypoints\"]\n                \n                # Crop and resize\n                img_cropped = crop_between_keypoints(img_rotated, left_rotated, right_rotated)\n                img_resized = resize_transform(image=img_cropped)[\"image\"]\n                \n                # Save the crop\n                output_path = os.path.join(output_dir, study_id, series_id, f\"{level.replace('/', '_')}.png\")\n                os.makedirs(os.path.dirname(output_path), exist_ok=True)\n                cv2.imwrite(output_path, img_resized)","metadata":{"execution":{"iopub.status.busy":"2024-09-13T11:38:35.711787Z","iopub.execute_input":"2024-09-13T11:38:35.712123Z","iopub.status.idle":"2024-09-13T11:38:35.741962Z","shell.execute_reply.started":"2024-09-13T11:38:35.712090Z","shell.execute_reply":"2024-09-13T11:38:35.741043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = \"/kaggle/input/best_model/pytorch/default/1/best_model.pth\"\nimage_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/\"\noutput_dir = \"/kaggle/working/model_extracted_crops_test/\"\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Load the trained model\nmodel = LumbarLineModel().to(device)\nmodel.load_state_dict(torch.load(model_path, map_location=device))\n\n# Extract crops using the model\nextract_crops_with_model(model, image_dir, output_dir, device)\n\nprint(\"Crop extraction completed.\")","metadata":{"execution":{"iopub.status.busy":"2024-09-13T11:38:35.743137Z","iopub.execute_input":"2024-09-13T11:38:35.743445Z","iopub.status.idle":"2024-09-13T11:38:38.844277Z","shell.execute_reply.started":"2024-09-13T11:38:35.743412Z","shell.execute_reply":"2024-09-13T11:38:38.843508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LumbarSpineDataset(Dataset):\n    def __init__(self, csv_file, img_dir, transform=None):\n        self.data = pd.read_csv(csv_file)\n        self.img_dir = img_dir\n        self.transform = transform\n        self.levels = ['L1_L2', 'L2_L3', 'L3_L4', 'L4_L5', 'L5_S1']\n        self.valid_indices = self._get_valid_indices()\n        \n    def _get_valid_indices(self):\n        valid_indices = []\n        for idx, row in self.data.iterrows():\n            study_id = str(row['study_id'])\n            study_folder = os.path.join(self.img_dir, study_id)\n            if not os.path.exists(study_folder):\n                continue\n            subfolders = [f for f in os.listdir(study_folder) if os.path.isdir(os.path.join(study_folder, f)) and row['series_description']==\"Sagittal T2/STIR\"]\n            if not subfolders:\n                continue\n            subfolder = subfolders[0]\n            if all(os.path.exists(os.path.join(study_folder, subfolder, f\"{level}.png\")) for level in self.levels):\n                valid_indices.append(idx)\n        return valid_indices\n        \n    def __len__(self):\n        return len(self.valid_indices)\n    \n    def __getitem__(self, idx):\n        real_idx = self.valid_indices[idx]\n        study_id = str(self.data.iloc[real_idx]['study_id'])\n        images = []\n        \n        study_folder = os.path.join(self.img_dir, study_id)\n        subfolders = [f for f in os.listdir(study_folder) if os.path.isdir(os.path.join(study_folder, f))]\n        subfolder = subfolders[0]  # Assume there's only one subfolder\n        \n        for level in self.levels:\n            img_path = os.path.join(study_folder, subfolder, f\"{level}.png\")\n            image = Image.open(img_path).convert('RGB')\n            if self.transform:\n                image = self.transform(image)\n            images.append(image)\n        \n        return images, study_id\n","metadata":{"execution":{"iopub.status.busy":"2024-09-13T11:38:38.845370Z","iopub.execute_input":"2024-09-13T11:38:38.845659Z","iopub.status.idle":"2024-09-13T11:38:38.858566Z","shell.execute_reply.started":"2024-09-13T11:38:38.845627Z","shell.execute_reply":"2024-09-13T11:38:38.857522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LumbarSpineModel(nn.Module):\n    def __init__(self, num_classes=75, weights_path='/kaggle/input/maxvit-224/maxvit_tiny_tf_224_weights.pth'):\n        super().__init__()\n        self.backbone = timm.create_model('maxvit_tiny_tf_224.in1k', pretrained=False, num_classes=0)\n        \n        # Load the weights\n        state_dict = torch.load(weights_path)\n        # Remove the 'head' keys from the state dict\n        state_dict = {k: v for k, v in state_dict.items() if not k.startswith('head.')}\n        self.backbone.load_state_dict(state_dict, strict=False)\n        \n        # Get the number of output features from the backbone\n        with torch.no_grad():\n            dummy_input = torch.randn(1, 3, 224, 224)\n            features = self.backbone(dummy_input)\n            num_features = features.shape[1]\n        \n        self.dropout = nn.Dropout(0.5)\n        self.fc = nn.Linear(num_features * 5, num_classes)\n        \n    def forward(self, x):\n        if isinstance(x, list):\n            features = []\n            for img in x:\n                feat = self.backbone(img)\n                features.append(feat)\n            combined_features = torch.cat(features, dim=1)\n        else:\n            batch_size, num_images, C, H, W = x.shape\n            x = x.view(batch_size * num_images, C, H, W)\n            features = self.backbone(x)\n            combined_features = features.view(batch_size, -1)\n        \n        x = self.dropout(combined_features)\n        return self.fc(x)\n\n# Create the model\nmodel = LumbarSpineModel().to(device)","metadata":{"execution":{"iopub.status.busy":"2024-09-13T11:38:38.859996Z","iopub.execute_input":"2024-09-13T11:38:38.860679Z","iopub.status.idle":"2024-09-13T11:38:41.681633Z","shell.execute_reply.started":"2024-09-13T11:38:38.860634Z","shell.execute_reply":"2024-09-13T11:38:41.680555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def generate_submission(model, test_loader, submission_file):\n    model.eval()\n    predictions = []\n    study_ids = []\n\n    with torch.no_grad():\n        for data, ids in test_loader:  # Assuming the dataset now returns study_ids\n            data = [img.to(device) for img in data]\n            output = model(data)\n            probs = torch.sigmoid(output).cpu().numpy()\n            predictions.append(probs)\n            study_ids.extend(ids)\n    \n    predictions = np.concatenate(predictions, axis=0)\n    \n    # Create submission DataFrame\n    sample_submission = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv')\n    conditions = sample_submission['row_id'].str.split('_', n=1).str[1].unique()\n    \n    submission_data = []\n    for i, study_id in enumerate(study_ids):\n        for j, condition in enumerate(conditions):\n            row = {\n                'row_id': f\"{study_id}_{condition}\",\n                'normal_mild': predictions[i, j*3],\n                'moderate': predictions[i, j*3 + 1],\n                'severe': predictions[i, j*3 + 2]\n            }\n            submission_data.append(row)\n    \n    submission = pd.DataFrame(submission_data)\n    \n    # Ensure probabilities sum to 1 for each condition\n    submission[['normal_mild', 'moderate', 'severe']] = submission[['normal_mild', 'moderate', 'severe']].div(\n        submission[['normal_mild', 'moderate', 'severe']].sum(axis=1), axis=0\n    )\n    \n    submission.to_csv(submission_file, index=False)\n\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ])\n\n# Create test dataset and dataloader\ntest_dataset = LumbarSpineDataset('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv', '/kaggle/working/model_extracted_crops_test', transform=transform)\ntest_loader = torch.utils.data.DataLoader(test_dataset, batch_size=8, shuffle=False, num_workers=4)\n\n# Load the trained model\nmodel = LumbarSpineModel().to(device)\nmodel.load_state_dict(torch.load('/kaggle/input/best_lumbar_spine_model/pytorch/default/1/best_lumbar_spine_model.pth'))\n\n# Generate submission\ngenerate_submission(model, test_loader, '/kaggle/working/submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-09-13T11:38:41.683188Z","iopub.execute_input":"2024-09-13T11:38:41.683812Z","iopub.status.idle":"2024-09-13T11:38:44.477413Z","shell.execute_reply.started":"2024-09-13T11:38:41.683762Z","shell.execute_reply":"2024-09-13T11:38:44.476419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-09-13T11:38:44.479152Z","iopub.execute_input":"2024-09-13T11:38:44.479587Z","iopub.status.idle":"2024-09-13T11:38:44.486960Z","shell.execute_reply.started":"2024-09-13T11:38:44.479529Z","shell.execute_reply":"2024-09-13T11:38:44.485939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2024-09-13T11:38:44.489663Z","iopub.execute_input":"2024-09-13T11:38:44.489996Z","iopub.status.idle":"2024-09-13T11:38:44.512306Z","shell.execute_reply.started":"2024-09-13T11:38:44.489962Z","shell.execute_reply":"2024-09-13T11:38:44.511394Z"},"trusted":true},"execution_count":null,"outputs":[]}]}