{"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":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! pip install -q matplotlib\n\n!pip install -q \"monai[pillow, tqdm]\"\n!pip install -q \"monai[ignite]\"\n!pip install -q \"monai[gdown]\"\n!pip install -q \"pydicom>=1.4.2\"\n!pip install -q \"highdicom>=0.18.2\"\n!pip install -q \"monai-deploy-app-sdk\"\n%matplotlib inline\n\n# !pip uninstall monai holoscan\n# !pip install monai[all] holoscan","metadata":{"execution":{"iopub.status.busy":"2024-07-11T11:27:45.074773Z","iopub.execute_input":"2024-07-11T11:27:45.075319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport shutil\nimport tempfile\nimport glob\nimport PIL.Image\nimport torch\nimport numpy as np\nimport pandas as pd\n\nfrom ignite.engine import Events\n\nfrom monai.apps import download_and_extract\nfrom monai.config import print_config\nfrom monai.networks.nets import DenseNet121\nfrom monai.engines import SupervisedTrainer\nfrom monai.transforms import (\n    EnsureChannelFirst,\n    Compose,\n    LoadImage,\n    RandFlip,\n    RandRotate,\n    RandZoom,\n    ScaleIntensity,\n    EnsureType,\n    Resize,\n    ToTensor,\n)\nfrom monai.utils import set_determinism\n\nset_determinism(seed=0)\n\nprint_config()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pathlib import Path\nimport pydicom\nimport matplotlib.pyplot as plt","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"directory = os.environ.get(\"/kaggle/working\")\nif directory is not None:\n    os.makedirs(directory, exist_ok=True)\nroot_dir = tempfile.mkdtemp() if directory is None else directory\nprint(root_dir)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set_determinism(seed=0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT = Path('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/')\ndf = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\ndf = df.loc[:,['study_id', 'spinal_canal_stenosis_l1_l2']]\n# df = df[df['spinal_canal_stenosis_l1_l2'] == 'Severe']\ndf1 = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')\ndf1 = df1[df1['series_description'] == 'Sagittal T2/STIR']\ndf2 = df1.merge(df, on=['study_id'])\n# spinal_canal_stenosis\n# print(df)\n# print(df1)\n# print(df2)\n# print(glob.glob(str(ROOT/'train_images'/'*')))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_names = ['Normal/Mild', 'Moderate', 'Severe']\nnum_class = len(class_names)\nROOT = Path('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/')\nTRAIN_IMAGE_PATH = ROOT / 'train_images'\nimage_files = [[],[],[]]\nfor idx,i in enumerate(class_names):\n    class_df = df2[df2['spinal_canal_stenosis_l1_l2'] == i]\n    for j in range(class_df.shape[0]):\n        image_path = TRAIN_IMAGE_PATH / str(class_df.iloc[j,0]) / str(class_df.iloc[j,1]) / '*'\n        image_files[idx] += glob.glob(str(image_path))\n# print(image_files)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# image_files[2]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_each = [len(image_files[i]) for i in range(num_class)]\nimage_files_list = []\nimage_class = []\nfor i in range(num_class):\n    image_files_list.extend(image_files[i])\n    image_class.extend([i] * num_each[i])\nnum_total = len(image_class)\n# image_width, image_height = PIL.Image.open(image_files_list[0]).size\n\nprint(f\"Total image count: {num_total}\")\n# print(f\"Image dimensions: {image_width} x {image_height}\")\nprint(f\"Label names: {class_names}\")\nprint(f\"Label counts: {num_each}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"im = pydicom.dcmread('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/4283570761/453728183/10.dcm').pixel_array\nplt.imshow(im, cmap=\"gray\", vmin=0, vmax=255)\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.subplots(3, 3, figsize=(8, 8))\nfor i, k in enumerate(np.random.randint(num_total, size=9)):\n#     im = PIL.Image.open(image_files_list[k])\n#     print(image_files_list[k])\n    im = pydicom.dcmread(image_files_list[k]).pixel_array\n    arr = np.array(im)\n    plt.subplot(3, 3, i + 1)\n    plt.xlabel(class_names[image_class[k]])\n    plt.imshow(arr, cmap=\"gray\", vmin=0, vmax=255)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_frac = 0.1\ntest_frac = 0.1\nlength = len(image_files_list)\nindices = np.arange(length)\nnp.random.shuffle(indices)\n\ntest_split = int(test_frac * length)\nval_split = int(val_frac * length) + test_split\ntest_indices = indices[:test_split]\nval_indices = indices[test_split:val_split]\ntrain_indices = indices[val_split:]\n\ntrain_x = [image_files_list[i] for i in train_indices]\ntrain_y = [image_class[i] for i in train_indices]\nval_x = [image_files_list[i] for i in val_indices]\nval_y = [image_class[i] for i in val_indices]\ntest_x = [image_files_list[i] for i in test_indices]\ntest_y = [image_class[i] for i in test_indices]\n\nprint(f\"Training count: {len(train_x)}, Validation count: \" f\"{len(val_x)}, Test count: {len(test_x)}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = Compose(\n    [\n        LoadImage(image_only=True),\n        EnsureChannelFirst(channel_dim=\"no_channel\"),\n        ScaleIntensity(),\n        RandRotate(range_x=np.pi / 12, prob=0.5, keep_size=True),\n        RandFlip(spatial_axis=0, prob=0.5),\n        RandZoom(min_zoom=0.9, max_zoom=1.1, prob=0.5),\n        EnsureType(),\n        Resize((256, 256, 256)),  # 调整图像大小\n        ToTensor(),\n    ]\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MedNISTDataset(torch.utils.data.Dataset):\n    def __init__(self, image_files, labels, transforms):\n        self.image_files = image_files\n        self.labels = labels\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.image_files)\n\n    def __getitem__(self, index):\n        return self.transforms(self.image_files[index]), self.labels[index]\n\n# just one dataset and loader, we won't bother with validation or testing \ntrain_ds = MedNISTDataset(image_files_list, image_class, train_transforms)\ntrain_loader = torch.utils.data.DataLoader(train_ds, batch_size=20, shuffle=True, num_workers=0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nnet = DenseNet121(spatial_dims=2, in_channels=1, out_channels=len(class_names)).to(device)\nloss_function = torch.nn.CrossEntropyLoss()\nopt = torch.optim.Adam(net.parameters(), 1e-5)\nmax_epochs = 1","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _prepare_batch(batch, device, non_blocking):\n    return tuple(b.to(device) for b in batch)\n\n\ntrainer = SupervisedTrainer(device, max_epochs, train_loader, net, opt, loss_function, prepare_batch=_prepare_batch)\n\n\n@trainer.on(Events.EPOCH_COMPLETED)\ndef _print_loss(engine):\n    print(f\"Epoch {engine.state.epoch}/{engine.state.max_epochs} Loss: {engine.state.output[0]['loss']}\")\n\n\ntrainer.run()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(\"train\", (12, 6))\nplt.subplot(1, 2, 1)\nplt.title(\"Epoch Average Loss\")\nx = [i + 1 for i in range(len(epoch_loss_values))]\ny = epoch_loss_values\nplt.xlabel(\"epoch\")\nplt.plot(x, y)\nplt.subplot(1, 2, 2)\nplt.title(\"Val AUC\")\nx = [val_interval * (i + 1) for i in range(len(metric_values))]\ny = metric_values\nplt.xlabel(\"epoch\")\nplt.plot(x, y)\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load(os.path.join(root_dir, \"best_metric_model.pth\")))\nmodel.eval()\ny_true = []\ny_pred = []\nwith torch.no_grad():\n    for test_data in test_loader:\n        test_images, test_labels = (\n            test_data[0].to(device),\n            test_data[1].to(device),\n        )\n        pred = model(test_images).argmax(dim=1)\n        for i in range(len(pred)):\n            y_true.append(test_labels[i].item())\n            y_pred.append(pred[i].item())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(classification_report(y_true, y_pred, target_names=class_names, digits=4))","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}