{"metadata":{"colab":{"provenance":[],"authorship_tag":"ABX9TyPn+OzgmlMeV9Ew+CIzro6X"},"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":97254,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":81582,"modelId":105908}],"dockerImageVersionId":30746,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **Import Libraries**","metadata":{"id":"omF87Q9kXvdK"}},{"cell_type":"code","source":"# !pip install torchsummary\n\n# imports for neural network\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, random_split, Subset\n\n# imports for vision tasks\nimport torchvision\nimport torchvision.models as models\nimport torchvision.transforms as transforms\nfrom torchvision.datasets import ImageFolder\nfrom torchvision.utils import make_grid\n# from torchsummary import summary\nimport pydicom\n\n# imports for preparing dataset\nimport os\nimport shutil\nimport zipfile\nimport pandas as pd\nfrom skimage import io\nfrom PIL import Image\nimport numpy as np\n\n# imports for visualizations\n%matplotlib inline\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg","metadata":{"execution":{"iopub.status.busy":"2024-08-20T04:31:45.161427Z","iopub.execute_input":"2024-08-20T04:31:45.161763Z","iopub.status.idle":"2024-08-20T04:31:50.980340Z","shell.execute_reply.started":"2024-08-20T04:31:45.161736Z","shell.execute_reply":"2024-08-20T04:31:50.979387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Preprocess Dataset and Device**","metadata":{"id":"KUd_8LBQ2RUV"}},{"cell_type":"code","source":"class_mapping_1 = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n\nimg_sizes_coords = pd.DataFrame(columns=['image_size', 'spinal_canal_stenosis_l1_l2', 'spinal_canal_stenosis_l2_l3', 'spinal_canal_stenosis_l3_l4', 'spinal_canal_stenosis_l4_l5', 'spinal_canal_stenosis_l5_s1',\n 'left_neural_foraminal_narrowing_l1_l2', 'left_neural_foraminal_narrowing_l2_l3', 'left_neural_foraminal_narrowing_l3_l4', 'left_neural_foraminal_narrowing_l4_l5', 'left_neural_foraminal_narrowing_l5_s1',\n 'right_neural_foraminal_narrowing_l1_l2', 'right_neural_foraminal_narrowing_l2_l3', 'right_neural_foraminal_narrowing_l3_l4', 'right_neural_foraminal_narrowing_l4_l5', 'right_neural_foraminal_narrowing_l5_s1',\n 'left_subarticular_stenosis_l1_l2', 'left_subarticular_stenosis_l2_l3', 'left_subarticular_stenosis_l3_l4', 'left_subarticular_stenosis_l4_l5', 'left_subarticular_stenosis_l5_s1',\n 'right_subarticular_stenosis_l1_l2', 'right_subarticular_stenosis_l2_l3', 'right_subarticular_stenosis_l3_l4', 'right_subarticular_stenosis_l4_l5', 'right_subarticular_stenosis_l5_s1'])\n","metadata":{"executionInfo":{"elapsed":368,"status":"ok","timestamp":1723762182346,"user":{"displayName":"Nancy Xiao","userId":"03828519108036202118"},"user_tz":420},"id":"9F6IqNg9x-QT","execution":{"iopub.status.busy":"2024-08-20T04:32:01.083873Z","iopub.execute_input":"2024-08-20T04:32:01.084549Z","iopub.status.idle":"2024-08-20T04:32:01.095933Z","shell.execute_reply.started":"2024-08-20T04:32:01.084518Z","shell.execute_reply":"2024-08-20T04:32:01.094852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"base_dir = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\ntrain_labels = pd.read_csv(os.path.join(base_dir, 'train.csv'))\ntrain_coords = pd.read_csv(os.path.join(base_dir, 'train_label_coordinates.csv'))\n# print(train_coords.head())\n# print(train_labels.head())\n# print(f'Image sizes: {img_sizes_coords.head()}')\n\ntrain_images = []\ntrain_classes = []\n\ntrain_dir = os.path.join(base_dir, 'train_images')\n# test_dir = os.path.join(base_dir, 'test_images')\n\ntranform=transforms.Compose([\n  transforms.Resize((128, 128)),\n  transforms.ToTensor()\n])\n\ndef load_dicom_image(file_path):\n    ds = pydicom.dcmread(file_path)\n    img = ds.pixel_array\n    if img.ndim == 2:\n        img = img.astype(np.uint8)\n        pil_image = Image.fromarray(img, mode='L')\n    else:\n        pil_image = Image.fromarray(img).convert('gray')\n    return pil_image\n\nround_num = 0\nfor pat_id in os.listdir(train_dir):\n    for ser_id in os.listdir(os.path.join(train_dir, pat_id)):\n        for img_id in os.listdir(os.path.join(train_dir, pat_id, ser_id)):\n            coords = train_coords[(train_coords['study_id'] == int(pat_id)) &\n                                  (train_coords['series_id'] == int(ser_id)) &\n                                  (train_coords['instance_number'] == int(img_id[:-4]))]\n            # print(coords.head())\n            img_path = os.path.join(train_dir, pat_id, ser_id, img_id)\n\n            if coords.empty:\n# #                 print('No coords found for image: ', pat_id, ser_id, img_id)\n                continue\n            else:\n                if os.path.exists(img_path):\n                    image = load_dicom_image(img_path)\n                else:\n                    print('Image not found: ', img_path)\n                    continue\n\n            for index, row in coords.iterrows():\n                # print('list of row', index)\n                cond = row['condition'].replace(' ', '_').lower()\n                lev = row['level'].replace('/', '_').lower()\n                col = cond + '_' + lev\n                sev = train_labels[(train_labels['study_id'] == int(pat_id))][col].values[0]\n                if not pd.isna(sev):\n                    if sev in class_mapping_1:\n                        label = torch.tensor(int(class_mapping_1[sev]))\n                        train_classes.append(label)\n\n                    cpy_image = image.copy()\n                    width, height = cpy_image.size\n                    x = row['x']\n                    y = row['y']\n                    if img_sizes_coords[(img_sizes_coords['image_size'] == (width, height))].empty:\n                        new_row = pd.DataFrame({'image_size': [(width, height)]})  # , col: [[(x, y)]\n                        img_sizes_coords = pd.concat([img_sizes_coords, new_row], ignore_index=True)\n                        idx = img_sizes_coords[img_sizes_coords['image_size'] == (width, height)].index[0]\n                        img_sizes_coords.at[idx, col] = ([(x, y)])\n                    else:\n                        idx = img_sizes_coords[img_sizes_coords['image_size'] == (width, height)].index[0]\n                        val = img_sizes_coords.at[idx, col]\n                        if (isinstance(img_sizes_coords.at[idx, col], float)):\n                            img_sizes_coords.at[idx, col] = [(x, y)]\n                        # else:\n                        #     img_sizes_coords.at[idx, col].append((x, y))\n\n                    window_size = int(0.1 * min(cpy_image.size))\n                    start_x = int(x - window_size / 2)\n                    end_x = int(x + window_size / 2)\n                    start_y = int(y - window_size / 2)\n                    end_y = int(y + window_size / 2)\n                    start_x = max(0, start_x)\n                    end_x = min(cpy_image.size[0], end_x)\n                    start_y = max(0, start_y)\n                    end_y = min(cpy_image.size[1], end_y)\n                    cropped_image = cpy_image.crop((start_x, start_y, end_x, end_y))\n                    cropped_image = tranform(cropped_image)\n                    train_images.append(cropped_image)\n                    round_num += 1\n                else:\n#                     print('Severity not found: ', sev)\n                    continue\n#     if round_num > 12000:\n#         break\n#     else:\n#         continue\n                    \nprint(f'Shape of images: {np.array(train_images).shape}')\nprint(f'Shape of classes: {np.array(train_classes).shape}')\n# print(f'Image sizes: {img_sizes_coords}')\n\ntrain_dataset = list(zip(train_images, train_classes))\n\ntrain_set, valid_set = random_split(train_dataset, [0.7, 0.3])\nbatch_size = 32\ntrain_data = DataLoader(dataset=train_set, batch_size=batch_size, shuffle=True)\nval_data = DataLoader(dataset=valid_set, batch_size=batch_size, shuffle=False)\nprint('Length of train dataset:', len(train_data))\nprint('Length of validation dataset:', len(val_data))\n\n# Set device to GPU if it's available\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"executionInfo":{"elapsed":309674,"status":"ok","timestamp":1723762504767,"user":{"displayName":"Nancy Xiao","userId":"03828519108036202118"},"user_tz":420},"id":"ipAB_CTb9gIx","outputId":"ee82d115-3bf6-48d2-e3c0-63e730a94fc7","execution":{"iopub.status.busy":"2024-08-20T04:32:07.523613Z","iopub.execute_input":"2024-08-20T04:32:07.524271Z","iopub.status.idle":"2024-08-20T04:44:30.432547Z","shell.execute_reply.started":"2024-08-20T04:32:07.524242Z","shell.execute_reply":"2024-08-20T04:44:30.431603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coord_idx = img_sizes_coords[img_sizes_coords['image_size'] == (640, 640)].index[0]\nimg_sizes_coords.at[coord_idx, 'left_neural_foraminal_narrowing_l1_l2'] = [(416.1303462321792, 180.85539714867616)]\nimg_sizes_coords.at[coord_idx, 'left_neural_foraminal_narrowing_l2_l3'] = [(393.9714867617108, 257.7596741344195)]\nimg_sizes_coords.at[coord_idx, 'left_neural_foraminal_narrowing_l3_l4'] = [(378.3299389002037, 321.62932790224033)]\nimg_sizes_coords.at[coord_idx, 'left_neural_foraminal_narrowing_l4_l5'] = [(373.1160896130346, 382.89205702647655)]\nimg_sizes_coords.at[coord_idx, 'left_neural_foraminal_narrowing_l5_s1'] = [(378.3299389002037, 448.0651731160896)]\nimg_sizes_coords.at[coord_idx, 'right_neural_foraminal_narrowing_l1_l2'] = [(422.6560121765601, 181.18721461187212)]\nimg_sizes_coords.at[coord_idx, 'right_neural_foraminal_narrowing_l2_l3'] = [(407.07001522070016, 263.01369863013696)]\nimg_sizes_coords.at[coord_idx, 'right_neural_foraminal_narrowing_l3_l4'] = [(395.3805175038052, 347.7625570776256)]\nimg_sizes_coords.at[coord_idx, 'right_neural_foraminal_narrowing_l4_l5'] = [(368.10502283105023, 408.1582952815829)]\nimg_sizes_coords.at[coord_idx, 'right_neural_foraminal_narrowing_l5_s1'] = [(384.6651445966514, 456.8645357686453)]\n\n# output_file = 'img_sizes_coords.csv'\n# img_sizes_coords.to_csv(output_file, index=False)","metadata":{"executionInfo":{"elapsed":239,"status":"ok","timestamp":1723762549418,"user":{"displayName":"Nancy Xiao","userId":"03828519108036202118"},"user_tz":420},"id":"_7Qcne2di7HO","execution":{"iopub.status.busy":"2024-08-20T04:52:27.047830Z","iopub.execute_input":"2024-08-20T04:52:27.048494Z","iopub.status.idle":"2024-08-20T04:52:27.059295Z","shell.execute_reply.started":"2024-08-20T04:52:27.048462Z","shell.execute_reply":"2024-08-20T04:52:27.058285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Define Model**","metadata":{"id":"TJrsaBdAdL0A"}},{"cell_type":"code","source":"# Googlenet\nos.makedirs('/root/.cache/torch/hub/checkpoints', exist_ok=True)\nshutil.copy('/kaggle/input/googlenet-1378be20.pth/pytorch/default/1/googlenet-1378be20.pth', '/root/.cache/torch/hub/checkpoints/googlenet-1378be20.pth')\n\nmodel = models.googlenet(pretrained=True)\nmodel.conv1 = nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\nmodel.fc=nn.Sequential(\n    nn.Linear(in_features=1024,out_features=512),\n    nn.ReLU(),\n    nn.Linear(in_features=512,out_features=128),\n    nn.ReLU(),\n    nn.Linear(in_features=128,out_features=32,bias=True),\n    nn.ReLU(),\n    nn.Linear(in_features=32,out_features=3,bias=True)\n)\nmodel.to(device);\n# Disable the input transformation since you are providing grayscale images\nmodel.transform_input = False\n# model\n# summary(model,(1, 128,128))","metadata":{"executionInfo":{"elapsed":1858,"status":"ok","timestamp":1723762556004,"user":{"displayName":"Nancy Xiao","userId":"03828519108036202118"},"user_tz":420},"id":"Xf537xlfcbqO","outputId":"b8c379be-722a-40f3-a531-00c78d11bb8d","execution":{"iopub.status.busy":"2024-08-20T04:52:32.787461Z","iopub.execute_input":"2024-08-20T04:52:32.788216Z","iopub.status.idle":"2024-08-20T04:52:34.040874Z","shell.execute_reply.started":"2024-08-20T04:52:32.788176Z","shell.execute_reply":"2024-08-20T04:52:34.040041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Tain Model**","metadata":{"id":"102Xhsw5AzZ_"}},{"cell_type":"code","source":"# Loss function and optimizer\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=0.001)\n\n'''Code to implement learning rate scheduler to decrease lr during training process'''\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\nscheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=5)\n# scheduler.step(epoch_val_loss)  # add this to the end of training loop (after validation)\n\n# Early stopping parameters\nearly_stopping_patience = 5\n\n# Define the number of epochs to train for\nepochs = 50\n\n# Using validation loss as metric\nbest_val_loss = float('inf')\nbest_epoch = 0\nearly_stopping_counter = 0\n\n# Save metrics at each epoch for plotting\nepoch_train_loss_values = []\nepoch_val_loss_values = []\nepoch_train_acc_values = []\nepoch_val_acc_values = []\n\nfor epoch in range(epochs):\n    model.train()  # Set model to training mode\n\n    train_losses, train_accuracies = [], []\n\n    for data, label in train_data:\n        data, label = data.to(device), label.to(device)  # Move data to the same device as the model\n        # data = data.to(device)\n        # if not isinstance(label, torch.Tensor):\n        #     label = torch.tensor(label, dtype=torch.long).to(device)\n        # else:\n        #     label = label.to(device)\n\n        optimizer.zero_grad()  # Clear previous epoch's gradients\n        output = model(data)  # Forward pass\n        loss = criterion(output, label)  # Compute loss\n        loss.backward()  # Backward pass\n        optimizer.step()  # Update weights\n\n        # Accumulate metrics\n        acc = (output.argmax(dim=1) == label).float().mean().item()\n        train_losses.append(loss.item())\n        train_accuracies.append(acc)\n\n    # Average metrics across all training steps\n    epoch_train_loss = sum(train_losses) / len(train_losses)\n    epoch_train_accuracy = sum(train_accuracies) / len(train_accuracies)\n\n    # Save current epochs training metrics\n    epoch_train_loss_values.append(epoch_train_loss)\n    epoch_train_acc_values.append(epoch_train_accuracy)\n\n    # Validation\n    model.eval()  # Set model to evaluation mode\n    val_losses, val_accuracies = [], []\n    with torch.no_grad():  # Disable gradient calculation\n        for data, label in val_data:\n            data, label = data.to(device), label.to(device)\n\n            val_output = model(data)\n            val_loss = criterion(val_output, label)\n\n            # Accumulate metrics\n            acc = (val_output.argmax(dim=1) == label).float().mean().item()\n            val_losses.append(val_loss.item())\n            val_accuracies.append(acc)\n\n    # Average metrics across all validation steps\n    epoch_val_loss = sum(val_losses) / len(val_losses)\n    epoch_val_accuracy = sum(val_accuracies) / len(val_accuracies)\n\n    # Save current epochs validation metrics\n    epoch_val_loss_values.append(epoch_val_loss)\n    epoch_val_acc_values.append(epoch_val_accuracy)\n\n    # Update best model if validation accuracy improves\n    if epoch_val_loss < best_val_loss:\n        torch.save(model.state_dict(), 'rsna_best_model.pth')\n\n        best_val_loss = epoch_val_loss\n        best_epoch = epoch + 1\n        early_stopping_counter = 0\n\n    else:\n        early_stopping_counter += 1\n\n    print(f'Epoch: {epoch + 1}\\n'\n          f'Train Acc: {epoch_train_accuracy:.3f}, Val Acc: {epoch_val_accuracy:.3f} '\n          f'Train Loss: {epoch_train_loss:.3f}, Val Loss: {epoch_val_loss:.3f}')\n    print(f'Best Metric: {best_val_loss:.3f} at epoch: {best_epoch}\\n')\n\n    if early_stopping_counter >= early_stopping_patience:\n        print(f\"Early stopping after {early_stopping_patience} epochs of no improvement.\")\n        break","metadata":{"id":"GVAN_Ur6A7mD","executionInfo":{"status":"ok","timestamp":1723752264137,"user_tz":420,"elapsed":911939,"user":{"displayName":"Nancy Xiao","userId":"03828519108036202118"}},"outputId":"bff0a26e-acc5-4640-a632-19cb2fe6be84","execution":{"iopub.status.busy":"2024-08-20T04:52:41.974883Z","iopub.execute_input":"2024-08-20T04:52:41.975767Z","iopub.status.idle":"2024-08-20T05:00:17.004278Z","shell.execute_reply.started":"2024-08-20T04:52:41.975733Z","shell.execute_reply":"2024-08-20T05:00:17.003239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Test Model**","metadata":{"id":"qsiPwrN-DfMi"}},{"cell_type":"markdown","source":"***Test data***","metadata":{"id":"FnAoDfCBtNPK"}},{"cell_type":"code","source":"test_labels = pd.DataFrame(columns=['study_id', 'series_id', 'instance_number', 'condition_level', 'class', 'normal_mild', 'moderate', 'severe'])\n\ntest_desc = pd.read_csv(os.path.join(base_dir, 'test_series_descriptions.csv'))\nprint(test_desc.head())\n\ntest_images =[]\ntest_severes = []\n\ntest_dir = os.path.join(base_dir, 'test_images')\n\nfor test_id in os.listdir(test_dir):\n  for test_ser_id in os.listdir(os.path.join(test_dir, test_id)):\n    test_ser_desc = test_desc[(test_desc['study_id'] == int(test_id)) & (test_desc['series_id'] == int(test_ser_id))]['series_description'].values[0]\n    for test_img_id in os.listdir(os.path.join(test_dir, test_id, test_ser_id)):\n      test_img_path = os.path.join(test_dir, test_id, test_ser_id, test_img_id)\n      test_image = load_dicom_image(test_img_path)\n\n      test_width, test_height = test_image.size\n      if img_sizes_coords[img_sizes_coords['image_size'] == (test_width, test_height)].empty:\n        print('No coords found for test image: ', test_id, test_ser_id, test_img_id, test_width, test_height)\n        continue\n      else:\n        idx = img_sizes_coords[img_sizes_coords['image_size'] == (test_width, test_height)].index[0]\n        if test_ser_desc == 'Sagittal T2/STIR':\n#           print('Sagittal T2/STIR test_width, test_height', test_width, test_height)\n          for col in img_sizes_coords.columns[1:6]:\n            [(x, y)] = img_sizes_coords.at[idx, col]\n            test_image_cpy = test_image.copy()\n            window_size = int(0.1 * min(test_image_cpy.size))\n            start_x = int(x - window_size / 2)\n            end_x = int(x + window_size / 2)\n            start_y = int(y - window_size / 2)\n            end_y = int(y + window_size / 2)\n            start_x = max(0, start_x)\n            end_x = min(test_image_cpy.size[0], end_x)\n            start_y = max(0, start_y)\n            end_y = min(test_image_cpy.size[1], end_y)\n            cropped_image = test_image_cpy.crop((start_x, start_y, end_x, end_y))\n            # plt.imshow(cropped_image)\n            # plt.title(f'Label: {label}')\n            # plt.show()\n            cropped_image = tranform(cropped_image)\n            test_images.append(cropped_image)\n            # plt.imshow(cropped_image.squeeze(0))\n            # plt.title(f'Label: {label}')\n            # plt.show()\n            # test_labels = pd.DataFrame(columns=['study_id', 'series_id', 'instance_number', 'condition_level', 'severe'])\n            new_row = pd.DataFrame({'study_id': [test_id], 'series_id': [test_ser_id], 'instance_number': [test_img_id], 'condition_level': [col]})\n            test_labels = pd.concat([test_labels, new_row], ignore_index=True)\n            test_severes.append(torch.tensor(0))\n\n\n        if test_ser_desc == 'Sagittal T1':\n#           print('Sagittal T1 test_width, test_height', test_width, test_height)\n          for col in img_sizes_coords.columns[6:16]:\n            [(x, y)] = img_sizes_coords.at[idx, col]\n            test_image_cpy = test_image.copy()\n            window_size = int(0.1 * min(test_image_cpy.size))\n            start_x = int(x - window_size / 2)\n            end_x = int(x + window_size / 2)\n            start_y = int(y - window_size / 2)\n            end_y = int(y + window_size / 2)\n            start_x = max(0, start_x)\n            end_x = min(test_image_cpy.size[0], end_x)\n            start_y = max(0, start_y)\n            end_y = min(test_image_cpy.size[1], end_y)\n            cropped_image = test_image_cpy.crop((start_x, start_y, end_x, end_y))\n            # plt.imshow(cropped_image)\n            # plt.title(f'Label: {label}')\n            # plt.show()\n            cropped_image = tranform(cropped_image)\n            test_images.append(cropped_image)\n            # plt.imshow(cropped_image.squeeze(0))\n            # plt.title(f'Label: {label}')\n            # plt.show()\n            # test_labels = pd.DataFrame(columns=['study_id', 'series_id', 'instance_number', 'condition_level', 'severe'])\n            new_row = pd.DataFrame({'study_id': [test_id], 'series_id': [test_ser_id], 'instance_number': [test_img_id], 'condition_level': [col]})\n            test_labels = pd.concat([test_labels, new_row], ignore_index=True)\n            test_severes.append(torch.tensor(0))\n\n        if test_ser_desc == 'Axial T2':\n#           print('Axial T2 test_width, test_height', test_width, test_height)\n          for col in img_sizes_coords.columns[16:]:\n            [(x, y)] = img_sizes_coords.at[idx, col]\n            test_image_cpy = test_image.copy()\n            window_size = int(0.1 * min(test_image_cpy.size))\n            start_x = int(x - window_size / 2)\n            end_x = int(x + window_size / 2)\n            start_y = int(y - window_size / 2)\n            end_y = int(y + window_size / 2)\n            start_x = max(0, start_x)\n            end_x = min(test_image_cpy.size[0], end_x)\n            start_y = max(0, start_y)\n            end_y = min(test_image_cpy.size[1], end_y)\n            cropped_image = test_image_cpy.crop((start_x, start_y, end_x, end_y))\n            # plt.imshow(cropped_image)\n            # plt.title(f'Label: {label}')\n            # plt.show()\n            cropped_image = tranform(cropped_image)\n            test_images.append(cropped_image)\n            # plt.imshow(cropped_image.squeeze(0))\n            # plt.title(f'Label: {label}')\n            # plt.show()\n            # test_labels = pd.DataFrame(columns=['study_id', 'series_id', 'instance_number', 'condition_level', 'severe'])\n            new_row = pd.DataFrame({'study_id': [test_id], 'series_id': [test_ser_id], 'instance_number': [test_img_id], 'condition_level': [col]})\n            test_labels = pd.concat([test_labels, new_row], ignore_index=True)\n            test_severes.append(torch.tensor(0))\n\n# test_labels.to_csv('test_labels.csv', index=False)\ntest_dataset = list(zip(test_images, test_severes))\ntest_img, test_label = test_dataset[0]\nprint('Image dimensions:', test_img.shape)\nprint('Image label:', test_label)\nprint('Length of test dataset:', len(test_dataset))\ntest_data = DataLoader(dataset=test_dataset, batch_size=batch_size, shuffle=False)\nprint('Length of test dataset:', len(test_data))\n","metadata":{"id":"Mka9Y9eEHaEu","executionInfo":{"status":"ok","timestamp":1723783100948,"user_tz":420,"elapsed":6927,"user":{"displayName":"Nancy Xiao","userId":"03828519108036202118"}},"outputId":"2348ab03-f256-4c37-a147-d5a7e1779da9","execution":{"iopub.status.busy":"2024-08-20T05:05:38.661482Z","iopub.execute_input":"2024-08-20T05:05:38.662031Z","iopub.status.idle":"2024-08-20T05:05:44.838422Z","shell.execute_reply.started":"2024-08-20T05:05:38.661994Z","shell.execute_reply":"2024-08-20T05:05:44.837497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Test model**","metadata":{"id":"9CkfO6catiu3"}},{"cell_type":"code","source":"# Load saved model parameters\nmodel.load_state_dict(torch.load('rsna_best_model.pth'))\n\nmodel.eval() # Set for validation and testing\ncriterion = nn.CrossEntropyLoss()\n# Test Loop\nwith torch.no_grad():  # Disable gradient tracking\n    total_accuracy, total_test_loss = 0.0, 0.0\n    num_batches = len(test_data)\n\n    output_idx = 0\n    for data, label in test_data:\n        data, label = data.to(device), label.to(device)  # Load data to same device as model\n\n        output = model(data)\n        loss = criterion(output, label)\n        accuracy = (output.argmax(dim=1) == label).float().mean()\n\n        total_test_loss += loss.item()\n        total_accuracy += accuracy.item()\n\n        # test_labels = pd.DataFrame(columns=['study_id', 'series_id', 'instance_number', 'condition_level', 'class', 'normal_mild', 'moderate', 'severe'])\n        for array in output.detach().cpu().numpy():\n          test_labels.at[output_idx, 'class'] = np.argmax(array)\n          x, y, z = array\n          test_labels.at[output_idx, 'normal_mild'] = x\n          test_labels.at[output_idx, 'moderate'] = y\n          test_labels.at[output_idx, 'severe'] = z\n          output_idx += 1\n\n    # Calculate the average loss and accuracy over all batches\n    avg_loss = total_test_loss / num_batches\n    avg_accuracy = total_accuracy / num_batches\n\n    print(f'test accuracy : {avg_accuracy:.3f}, test loss : {avg_loss:.3f}')\n    print(test_labels.head())\n#     test_labels.to_csv('test_labels.csv', index=False)","metadata":{"executionInfo":{"elapsed":41853,"status":"ok","timestamp":1723783162849,"user":{"displayName":"Nancy Xiao","userId":"03828519108036202118"},"user_tz":420},"id":"Xgy1C2vtDkG2","outputId":"0279c040-9f37-4128-c632-666d57a22990","execution":{"iopub.status.busy":"2024-08-20T05:05:56.081966Z","iopub.execute_input":"2024-08-20T05:05:56.082689Z","iopub.status.idle":"2024-08-20T05:05:56.646222Z","shell.execute_reply.started":"2024-08-20T05:05:56.082656Z","shell.execute_reply":"2024-08-20T05:05:56.645284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Submission report**","metadata":{"id":"IEdKmQ2ltolM"}},{"cell_type":"code","source":"# output_subm = pd.DataFrame(columns=['row_id', 'normal_mild', 'moderate', 'severe'], dtype=object)\noutput_subm = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv')\n\nfor col in img_sizes_coords.columns[1:]:\n    max_row = -1\n    max_class = -1\n    for idx, row in test_labels[test_labels['condition_level'] == col].iterrows():\n        if row['class'] > max_class:\n            max_class = row['class']\n            max_row = idx\n    if (not max_row < 0) and \\\n       (not pd.isna(test_labels.at[max_row, 'study_id'])) and \\\n       (not pd.isna(test_labels.at[max_row, 'normal_mild'])) and \\\n       (not pd.isna(test_labels.at[max_row, 'moderate'])) and \\\n       (not pd.isna(test_labels.at[max_row, 'severe'])):\n        subm_idx = output_subm[output_subm['row_id'] == (test_labels.at[max_row, 'study_id'] + '_' + col)].index[0]\n        output_subm.at[subm_idx, 'normal_mild'] = test_labels.at[max_row, 'normal_mild']\n        output_subm.at[subm_idx, 'moderate'] = test_labels.at[max_row, 'moderate']                                           \n        output_subm.at[subm_idx, 'severe'] = test_labels.at[max_row, 'severe']  \n    else:\n        print('No output is found', test_labels.at[max_row, 'study_id'] + '_' + col)\n        continue\n        \n#     new_row = pd.DataFrame({\n#             'row_id': [test_labels.at[max_row, 'study_id'] + '_' + col], \n#             'normal_mild': [test_labels.at[max_row, 'normal_mild']], \n#             'moderate': [test_labels.at[max_row, 'moderate']], \n#             'severe': [test_labels.at[max_row, 'severe']]\n#         })\n#     output_subm = pd.concat([output_subm, new_row], ignore_index=True)\n            \nprint(output_subm.head())\noutput_subm.to_csv('submission.csv', index=False)\n","metadata":{"id":"xnk9deHFfxCw","executionInfo":{"status":"ok","timestamp":1723786021596,"user_tz":420,"elapsed":249,"user":{"displayName":"Nancy Xiao","userId":"03828519108036202118"}},"outputId":"4da296a7-e887-434a-a149-3219468a35d7","execution":{"iopub.status.busy":"2024-08-20T05:06:10.391673Z","iopub.execute_input":"2024-08-20T05:06:10.392491Z","iopub.status.idle":"2024-08-20T05:06:10.489800Z","shell.execute_reply.started":"2024-08-20T05:06:10.392454Z","shell.execute_reply":"2024-08-20T05:06:10.488884Z"},"trusted":true},"execution_count":null,"outputs":[]}]}