{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport cv2\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader, Dataset\nfrom transformers import ViTForImageClassification, ViTFeatureExtractor\nfrom sklearn.model_selection import train_test_split\nfrom PIL import Image\nimport timm\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-31T07:56:31.571073Z","iopub.execute_input":"2024-05-31T07:56:31.571338Z","iopub.status.idle":"2024-05-31T07:57:02.472888Z","shell.execute_reply.started":"2024-05-31T07:56:31.571314Z","shell.execute_reply":"2024-05-31T07:57:02.472076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define dataset class\nclass SpineDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None):\n        self.df = df\n        self.img_dir = img_dir\n        self.transform = transform\n        self.class_to_index = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        study_id = row['study_id']\n        series_id = row['series_id']\n        instance_number = row['instance_number']\n        label = row['label']\n        image_path = os.path.join(self.img_dir, f'{study_id}/{series_id}/{instance_number}.dcm')\n        \n        dicom = pydicom.dcmread(image_path)\n        image = dicom.pixel_array\n        image = cv2.resize(image, (224, 224))\n        image = Image.fromarray(image).convert(\"RGB\")\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        label_idx = self.class_to_index[label]\n        label_tensor = torch.tensor(label_idx, dtype=torch.long)\n        \n        return {'image': image, 'label': label_tensor}\n","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.474508Z","iopub.execute_input":"2024-05-31T07:57:02.474803Z","iopub.status.idle":"2024-05-31T07:57:02.483601Z","shell.execute_reply.started":"2024-05-31T07:57:02.474777Z","shell.execute_reply":"2024-05-31T07:57:02.482692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(10),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.488985Z","iopub.execute_input":"2024-05-31T07:57:02.489238Z","iopub.status.idle":"2024-05-31T07:57:02.495218Z","shell.execute_reply.started":"2024-05-31T07:57:02.489216Z","shell.execute_reply":"2024-05-31T07:57:02.494331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load CSV files\ntrain_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\ncoords_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv')\n","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.496534Z","iopub.execute_input":"2024-05-31T07:57:02.496938Z","iopub.status.idle":"2024-05-31T07:57:02.660298Z","shell.execute_reply.started":"2024-05-31T07:57:02.496907Z","shell.execute_reply":"2024-05-31T07:57:02.659485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.661463Z","iopub.execute_input":"2024-05-31T07:57:02.661783Z","iopub.status.idle":"2024-05-31T07:57:02.707995Z","shell.execute_reply.started":"2024-05-31T07:57:02.661733Z","shell.execute_reply":"2024-05-31T07:57:02.707139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.columns.unique()","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.709208Z","iopub.execute_input":"2024-05-31T07:57:02.709561Z","iopub.status.idle":"2024-05-31T07:57:02.716712Z","shell.execute_reply.started":"2024-05-31T07:57:02.709531Z","shell.execute_reply":"2024-05-31T07:57:02.715778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.718075Z","iopub.execute_input":"2024-05-31T07:57:02.718731Z","iopub.status.idle":"2024-05-31T07:57:02.738006Z","shell.execute_reply.started":"2024-05-31T07:57:02.718699Z","shell.execute_reply":"2024-05-31T07:57:02.737176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords_df['level'].values","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.739465Z","iopub.execute_input":"2024-05-31T07:57:02.740284Z","iopub.status.idle":"2024-05-31T07:57:02.751851Z","shell.execute_reply.started":"2024-05-31T07:57:02.740238Z","shell.execute_reply":"2024-05-31T07:57:02.751049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Convert 'level' column to lowercase and replace '/' with '_'\ncoords_df['level'] = coords_df['level'].str.lower().str.replace('/', '_')\n\n# Convert 'condition' column to lowercase and replace spaces with '_'\ncoords_df['condition'] = coords_df['condition'].str.lower().str.replace(' ', '_')\n","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.755099Z","iopub.execute_input":"2024-05-31T07:57:02.755362Z","iopub.status.idle":"2024-05-31T07:57:02.834904Z","shell.execute_reply.started":"2024-05-31T07:57:02.755340Z","shell.execute_reply":"2024-05-31T07:57:02.834156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.836037Z","iopub.execute_input":"2024-05-31T07:57:02.836389Z","iopub.status.idle":"2024-05-31T07:57:02.849648Z","shell.execute_reply.started":"2024-05-31T07:57:02.836357Z","shell.execute_reply":"2024-05-31T07:57:02.848570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords_df.columns","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.851040Z","iopub.execute_input":"2024-05-31T07:57:02.851818Z","iopub.status.idle":"2024-05-31T07:57:02.862760Z","shell.execute_reply.started":"2024-05-31T07:57:02.851785Z","shell.execute_reply":"2024-05-31T07:57:02.861776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords_df['condition'].unique()","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.863693Z","iopub.execute_input":"2024-05-31T07:57:02.864043Z","iopub.status.idle":"2024-05-31T07:57:02.883036Z","shell.execute_reply.started":"2024-05-31T07:57:02.864017Z","shell.execute_reply":"2024-05-31T07:57:02.882053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Add the label column by mapping values from train_df\ncoords_df['label'] = coords_df.apply(\n    lambda x: train_df.loc[train_df['study_id'] == x['study_id'], f\"{x['condition']}_{x['level']}\"].values[0],\n    axis=1\n)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:02.884142Z","iopub.execute_input":"2024-05-31T07:57:02.884435Z","iopub.status.idle":"2024-05-31T07:57:18.760832Z","shell.execute_reply.started":"2024-05-31T07:57:02.884413Z","shell.execute_reply":"2024-05-31T07:57:18.760027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:18.761874Z","iopub.execute_input":"2024-05-31T07:57:18.762118Z","iopub.status.idle":"2024-05-31T07:57:18.774216Z","shell.execute_reply.started":"2024-05-31T07:57:18.762097Z","shell.execute_reply":"2024-05-31T07:57:18.773154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords_df['label'].isna()","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:18.775273Z","iopub.execute_input":"2024-05-31T07:57:18.775524Z","iopub.status.idle":"2024-05-31T07:57:18.792370Z","shell.execute_reply.started":"2024-05-31T07:57:18.775502Z","shell.execute_reply":"2024-05-31T07:57:18.791432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords_df = coords_df.dropna(subset=['label'])","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:18.793308Z","iopub.execute_input":"2024-05-31T07:57:18.793604Z","iopub.status.idle":"2024-05-31T07:57:18.815235Z","shell.execute_reply.started":"2024-05-31T07:57:18.793581Z","shell.execute_reply":"2024-05-31T07:57:18.814370Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Split dataset into training and validation sets\ntrain_data, val_data = train_test_split(coords_df, test_size=0.2, random_state=42)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:18.816235Z","iopub.execute_input":"2024-05-31T07:57:18.816521Z","iopub.status.idle":"2024-05-31T07:57:18.834577Z","shell.execute_reply.started":"2024-05-31T07:57:18.816498Z","shell.execute_reply":"2024-05-31T07:57:18.833766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create datasets and dataloaders\ntrain_dataset = SpineDataset(train_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=transform)\nval_dataset = SpineDataset(val_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=transform)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:18.835762Z","iopub.execute_input":"2024-05-31T07:57:18.836068Z","iopub.status.idle":"2024-05-31T07:57:18.840613Z","shell.execute_reply.started":"2024-05-31T07:57:18.836045Z","shell.execute_reply":"2024-05-31T07:57:18.839575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:18.841731Z","iopub.execute_input":"2024-05-31T07:57:18.842033Z","iopub.status.idle":"2024-05-31T07:57:18.851756Z","shell.execute_reply.started":"2024-05-31T07:57:18.842011Z","shell.execute_reply":"2024-05-31T07:57:18.851049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load pre-trained ResNet50 model from timm and modify the classification head\nmodel = timm.create_model(\"hf_hub:timm/resnet50.a1_in1k\", pretrained=True)\nnum_classes = 3  # Normal/Mild, Moderate, Severe\nmodel.fc = nn.Linear(model.fc.in_features, num_classes)\n\n# Define optimizer, scheduler, and loss function\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\nscheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)\nloss_fn = nn.CrossEntropyLoss()\n\n# Training loop with early stopping\nnum_epochs = 7\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(device)\n\nbest_val_loss = float('inf')\npatience, trials = 2, 0\n\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    correct_predictions = 0\n    total_samples = 0\n\n    for batch in train_loader:\n        images = batch['image'].to(device)\n        labels = batch['label'].to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = loss_fn(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n        _, predicted = torch.max(outputs, 1)\n        correct_predictions += (predicted == labels).sum().item()\n        total_samples += labels.size(0)\n\n    train_loss = running_loss / len(train_loader)\n    train_accuracy = correct_predictions / total_samples\n\n    # Validation loop\n    model.eval()\n    val_running_loss = 0.0\n    val_correct_predictions = 0\n    val_total_samples = 0\n\n    with torch.no_grad():\n        for batch in val_loader:\n            images = batch['image'].to(device)\n            labels = batch['label'].to(device)\n\n            outputs = model(images)\n            loss = loss_fn(outputs, labels)\n\n            val_running_loss += loss.item()\n            _, predicted = torch.max(outputs, 1)\n            val_correct_predictions += (predicted == labels).sum().item()\n            val_total_samples += labels.size(0)\n\n    val_loss = val_running_loss / len(val_loader)\n    val_accuracy = val_correct_predictions / val_total_samples\n\n    print(f\"Epoch {epoch+1}/{num_epochs}, \"\n          f\"Train Loss: {train_loss:.4f}, Train Accuracy: {train_accuracy:.4f}, \"\n          f\"Val Loss: {val_loss:.4f}, Val Accuracy: {val_accuracy:.4f}\")\n\n    scheduler.step()\n\n    # Check for early stopping\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        trials = 0\n        torch.save(model.state_dict(), 'best_spine_condition_classifier.pth')\n    else:\n        trials += 1\n        if trials >= patience:\n            print('Early stopping')\n            break\n","metadata":{"execution":{"iopub.status.busy":"2024-05-31T07:57:18.853046Z","iopub.execute_input":"2024-05-31T07:57:18.853303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save the model\ntorch.save(model.state_dict(), 'spine_condition_classifier.pth')\n","metadata":{},"execution_count":null,"outputs":[]}]}