{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Training Notebook\n\n# RSNA 2023 Abdominal Trauma Detection","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setup and Imports\n","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport numpy as np\nimport pandas as pd\nfrom matplotlib import pyplot as plt\nfrom sklearn.model_selection import train_test_split\n","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:05:50.94961Z","iopub.execute_input":"2023-09-07T08:05:50.95Z","iopub.status.idle":"2023-09-07T08:05:50.956343Z","shell.execute_reply.started":"2023-09-07T08:05:50.94997Z","shell.execute_reply":"2023-09-07T08:05:50.955063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    SEED = 42\n    IMAGE_SIZE = [256, 256]\n    BATCH_SIZE = 64\n    EPOCHS = 10\n    TARGET_COLS  = [\n        \"bowel_injury\", \"extravasation_injury\",\n        \"kidney_healthy\", \"kidney_low\", \"kidney_high\",\n        \"liver_healthy\", \"liver_low\", \"liver_high\",\n        \"spleen_healthy\", \"spleen_low\", \"spleen_high\",\n    ]\n    AUTOTUNE = None \n\nconfig = Config()","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:06:45.670583Z","iopub.execute_input":"2023-09-07T08:06:45.67098Z","iopub.status.idle":"2023-09-07T08:06:45.677994Z","shell.execute_reply.started":"2023-09-07T08:06:45.670949Z","shell.execute_reply":"2023-09-07T08:06:45.676789Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Reproducibility\n","metadata":{}},{"cell_type":"code","source":"torch.manual_seed(Config.SEED)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:06:48.929476Z","iopub.execute_input":"2023-09-07T08:06:48.929866Z","iopub.status.idle":"2023-09-07T08:06:48.942475Z","shell.execute_reply.started":"2023-09-07T08:06:48.929835Z","shell.execute_reply":"2023-09-07T08:06:48.941298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset\n\nThe dataset provided in the competition consists of DICOM images. We will not be training on the DICOM images, rather would work on PNG image which are extracted from the DICOM format.\n\n[A helpful resource on the conversion of DICOM to PNG](https://www.kaggle.com/code/radek1/how-to-process-dicom-images-to-pngs)","metadata":{}},{"cell_type":"code","source":"BASE_PATH = f\"/kaggle/input/rsna-atd-512x512-png-v2-dataset\"","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-09-07T08:06:52.239509Z","iopub.execute_input":"2023-09-07T08:06:52.239908Z","iopub.status.idle":"2023-09-07T08:06:52.244872Z","shell.execute_reply.started":"2023-09-07T08:06:52.239874Z","shell.execute_reply":"2023-09-07T08:06:52.243661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Meta Data\n\nThe `train.csv` file contains the following meta information:\n\n- `patient_id`: A unique ID code for each patient.\n- `series_id`: A unique ID code for each scan.\n- `instance_number`: The image number within the scan. The lowest instance number for many series is above zero as the original scans were cropped to the abdomen.\n- `[bowel/extravasation]_[healthy/injury]`: The two injury types with binary targets.\n- `[kidney/liver/spleen]_[healthy/low/high]`: The three injury types with three target levels.\n- `any_injury`: Whether the patient had any injury at all.\n","metadata":{}},{"cell_type":"code","source":"# train\ndataframe = pd.read_csv(f\"{BASE_PATH}/train.csv\")\ndataframe[\"image_path\"] = f\"{BASE_PATH}/train_images\"\\\n                    + \"/\" + dataframe.patient_id.astype(str)\\\n                    + \"/\" + dataframe.series_id.astype(str)\\\n                    + \"/\" + dataframe.instance_number.astype(str) +\".png\"\ndataframe = dataframe.drop_duplicates()\n\ndataframe.head(2)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:07:00.529906Z","iopub.execute_input":"2023-09-07T08:07:00.530306Z","iopub.status.idle":"2023-09-07T08:07:00.71408Z","shell.execute_reply.started":"2023-09-07T08:07:00.530274Z","shell.execute_reply":"2023-09-07T08:07:00.712888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We split the training dataset into train and validation. This is a common practise in the Machine Learning pipelines. We not only want to train our model, but also want to validate it's training.\n\nA small catch here is that the training and validation data should have an aligned data distribution. Here we handle that by grouping the lables and then splitting the dataset. This ensures an aligned data distribution between the training and the validation splits.","metadata":{}},{"cell_type":"code","source":"# Function to handle the split for each group\ndef split_group(group, test_size=0.2):\n    if len(group) == 1:\n        return (group, pd.DataFrame()) if np.random.rand() < test_size else (pd.DataFrame(), group)\n    else:\n        return train_test_split(group, test_size=test_size, random_state=42)\n\n# Initialize the train and validation datasets\ntrain_data = pd.DataFrame()\nval_data = pd.DataFrame()\n\n# Iterate through the groups and split them, handling single-sample groups\nfor _, group in dataframe.groupby(Config.TARGET_COLS):\n    train_group, val_group = split_group(group)\n    train_data = pd.concat([train_data, train_group], ignore_index=True)\n    val_data = pd.concat([val_data, val_group], ignore_index=True)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:07:04.569871Z","iopub.execute_input":"2023-09-07T08:07:04.570699Z","iopub.status.idle":"2023-09-07T08:07:04.730092Z","shell.execute_reply.started":"2023-09-07T08:07:04.570661Z","shell.execute_reply":"2023-09-07T08:07:04.728994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.shape, val_data.shape","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:07:09.059558Z","iopub.execute_input":"2023-09-07T08:07:09.060863Z","iopub.status.idle":"2023-09-07T08:07:09.071008Z","shell.execute_reply.started":"2023-09-07T08:07:09.060814Z","shell.execute_reply":"2023-09-07T08:07:09.069841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Pipeline","metadata":{}},{"cell_type":"code","source":"from PIL import Image\nimport torchvision.transforms as transforms\nimport torch\n\ndef decode_image_and_label(image_path, label):\n    # Open the image using PIL\n    image = Image.open(image_path)\n    \n    # Define transformations for resizing and normalizing the image\n    transform = transforms.Compose([\n        transforms.Resize(config.IMAGE_SIZE),\n        transforms.ToTensor(),  # Converts the image to a PyTorch tensor\n        transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])  # Adjust mean and std as needed\n    ])\n    \n    # Apply the transformations to the image\n    image = transform(image)\n    \n    # Convert label to PyTorch tensor (assuming label is a list or numpy array)\n    label = torch.tensor(label, dtype=torch.float32)\n    \n    # Split the label into the desired segments\n    labels = (label[0:1], label[1:2], label[2:5], label[5:8], label[8:11])\n    \n    return (image, labels)\n","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:07:12.392285Z","iopub.execute_input":"2023-09-07T08:07:12.392695Z","iopub.status.idle":"2023-09-07T08:07:12.789654Z","shell.execute_reply.started":"2023-09-07T08:07:12.392665Z","shell.execute_reply":"2023-09-07T08:07:12.788649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import datasets\nfrom PIL import Image\n\n# Define a custom dataset class\nclass CustomDataset(Dataset):\n    def __init__(self, image_paths, labels, transform=None):\n        self.image_paths = image_paths\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        image_path = self.image_paths[idx]\n        label = self.labels[idx]\n        \n        # Load and preprocess the image\n        image = Image.open(image_path)\n        \n        # Convert grayscale image to RGB if it has only one channel\n        if image.mode == 'L':\n            image = image.convert('RGB')\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        return image, label\n\n# Define data augmentation transformations\ntransform = transforms.Compose([\n    transforms.Resize(Config.IMAGE_SIZE),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomCrop(size=Config.IMAGE_SIZE, padding=10),  # Adjust padding as needed\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),  # Adjust mean and std as needed\n])\n\n# Create custom datasets for training and validation\ntrain_dataset = CustomDataset(image_paths=train_data['image_path'].tolist(), labels=train_data[Config.TARGET_COLS].values, transform=transform)\nval_dataset = CustomDataset(image_paths=val_data['image_path'].tolist(), labels=val_data[Config.TARGET_COLS].values, transform=transform)\n\n# Create data loaders for training and validation\ntrain_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, shuffle=False)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:07:16.299759Z","iopub.execute_input":"2023-09-07T08:07:16.300171Z","iopub.status.idle":"2023-09-07T08:07:16.317904Z","shell.execute_reply.started":"2023-09-07T08:07:16.300139Z","shell.execute_reply":"2023-09-07T08:07:16.316823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Get a batch of data from the DataLoader\nbatch = next(iter(train_loader))\n\n# Unpack the batch into images and labels\nimages, labels = batch\n\n# Now you can access the shapes of images and labels\nprint(images.shape)  # Shape of the batch of images\nprint([label.shape for label in labels])  # Shapes of individual label segments\n","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:07:20.129706Z","iopub.execute_input":"2023-09-07T08:07:20.130117Z","iopub.status.idle":"2023-09-07T08:07:21.265718Z","shell.execute_reply.started":"2023-09-07T08:07:20.130086Z","shell.execute_reply":"2023-09-07T08:07:21.26473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom PIL import Image\n\n# Function to display an image gallery\ndef plot_image_gallery(images, rows, cols, value_range=(0, 1)):\n    fig, axes = plt.subplots(rows, cols, figsize=(12, 8))\n    for i, ax in enumerate(axes.ravel()):\n        if i < len(images):\n            image = images[i]\n            # Reverse the normalization to display images correctly\n            image = (image * (value_range[1] - value_range[0])) + value_range[0]\n            image = Image.fromarray((image * 255).astype('uint8'))\n            ax.imshow(image)\n            ax.axis('off')\n        else:\n            ax.axis('off')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:07:24.625109Z","iopub.execute_input":"2023-09-07T08:07:24.625515Z","iopub.status.idle":"2023-09-07T08:07:24.637922Z","shell.execute_reply.started":"2023-09-07T08:07:24.625485Z","shell.execute_reply":"2023-09-07T08:07:24.636363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Build Model\n","metadata":{}},{"cell_type":"code","source":"# Define a custom model class\nclass CustomModel(nn.Module):\n    def __init__(self, num_classes):\n        super(CustomModel, self).__init__()\n        \n        # Define Backbone (ResNet-50)\n        self.backbone = torchvision.models.resnet50(pretrained=True)\n        self.backbone.fc = nn.Identity()  # Remove the final fully connected layer\n        \n        # Define necks for each head\n        self.neck_bowel = nn.Linear(2048, 32)\n        self.neck_extra = nn.Linear(2048, 32)\n        self.neck_liver = nn.Linear(2048, 32)\n        self.neck_kidney = nn.Linear(2048, 32)\n        self.neck_spleen = nn.Linear(2048, 32)\n        \n        # Define heads\n        self.head_bowel = nn.Linear(32, 1)\n        self.head_extra = nn.Linear(32, 1)\n        self.head_liver = nn.Linear(32, 3)\n        self.head_kidney = nn.Linear(32, 3)\n        self.head_spleen = nn.Linear(32, 3)\n\n    def forward(self, x):\n        # Backbone\n        x = self.backbone(x)\n        \n        # Check if the tensor's shape allows for mean pooling\n        if x.dim() > 2:\n            # Global Average Pooling (GAP)\n            x = torch.mean(x, dim=[2, 3])  # GAP\n        else:\n            # If the tensor is already 2D (e.g., when using adaptive pooling)\n            x = x.view(x.size(0), -1)\n        \n        # Necks\n        x_bowel = self.neck_bowel(x)\n        x_extra = self.neck_extra(x)\n        x_liver = self.neck_liver(x)\n        x_kidney = self.neck_kidney(x)\n        x_spleen = self.neck_spleen(x)\n        \n        # Heads\n        out_bowel = torch.sigmoid(self.head_bowel(x_bowel))  # Sigmoid for binary classification\n        out_extra = torch.sigmoid(self.head_extra(x_extra))  # Sigmoid for binary classification\n        out_liver = torch.softmax(self.head_liver(x_liver), dim=1)  # Softmax for multi-class\n        out_kidney = torch.softmax(self.head_kidney(x_kidney), dim=1)  # Softmax for multi-class\n        out_spleen = torch.softmax(self.head_spleen(x_spleen), dim=1)  # Softmax for multi-class\n        \n        return out_bowel, out_extra, out_liver, out_kidney, out_spleen\n","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:07:29.600856Z","iopub.execute_input":"2023-09-07T08:07:29.601567Z","iopub.status.idle":"2023-09-07T08:07:29.61471Z","shell.execute_reply.started":"2023-09-07T08:07:29.601532Z","shell.execute_reply":"2023-09-07T08:07:29.613261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:08:31.329544Z","iopub.execute_input":"2023-09-07T08:08:31.329967Z","iopub.status.idle":"2023-09-07T08:08:31.335509Z","shell.execute_reply.started":"2023-09-07T08:08:31.329933Z","shell.execute_reply":"2023-09-07T08:08:31.33436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes = 3  # Number of classes for liver, kidney, and spleen heads\nmodel = CustomModel(num_classes)\n\n# Print the model architecture\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:08:33.609487Z","iopub.execute_input":"2023-09-07T08:08:33.609906Z","iopub.status.idle":"2023-09-07T08:08:34.825595Z","shell.execute_reply.started":"2023-09-07T08:08:33.609874Z","shell.execute_reply":"2023-09-07T08:08:34.824464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train the model with \"model.fit\"","metadata":{}},{"cell_type":"code","source":"total_train_steps = len(train_loader) * Config.BATCH_SIZE * Config.EPOCHS\n#warmup steps\nwarmup_steps = int(total_train_steps * 0.10)\n#decay steps\ndecay_steps = total_train_steps - warmup_steps\n\nprint(f\"Total training steps: {total_train_steps}\")\nprint(f\"Warmup steps: {warmup_steps}\")\nprint(f\"Decay steps: {decay_steps}\")","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:08:43.010172Z","iopub.execute_input":"2023-09-07T08:08:43.010559Z","iopub.status.idle":"2023-09-07T08:08:43.017934Z","shell.execute_reply.started":"2023-09-07T08:08:43.010529Z","shell.execute_reply":"2023-09-07T08:08:43.016963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.optim as optim\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\n\n# Assuming you have already defined your model, optimizer, loss function, and data loaders as shown previously\n\ncriterion = {\n    \"bowel\": nn.BCELoss(),\n    \"extra\": nn.BCELoss(),\n    \"liver\": nn.CrossEntropyLoss(),\n    \"kidney\": nn.CrossEntropyLoss(),\n    \"spleen\": nn.CrossEntropyLoss(),\n}\n\n# Define the optimizer\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\n\n# Training loop\ndef train_model(model, train_loader, optimizer, criterion, num_epochs):\n    model.train()  # Set the model to training mode\n\n    for epoch in range(num_epochs):\n        running_loss = 0.0\n        for inputs, labels in train_loader:\n            optimizer.zero_grad()  # Zero the parameter gradients\n            outputs = model(inputs)  # Forward pass\n            losses = [criterion[head](output, label) for head, output, label in zip(Config.TARGET_COLS, outputs, labels)]\n            loss = sum(losses)  # Sum the losses from all heads\n            loss.backward()  # Backpropagation\n            optimizer.step()  # Update the model weights\n\n            running_loss += loss.item()\n\n        # Print the average loss for this epoch\n        print(f\"Epoch [{epoch+1}/{num_epochs}] - Loss: {running_loss / len(train_loader)}\")\n\n# Training parameters\nnum_epochs = Config.EPOCHS\n\n# Start training\ntrain_model(model, train_loader, optimizer, criterion, num_epochs)\n\n# Save the trained model if needed\ntorch.save(model.state_dict(), \"model.pth\")\n\n","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:12:19.500243Z","iopub.execute_input":"2023-09-07T08:12:19.500726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}],"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}}