{"metadata":{"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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! pip install monai\n!git clone https://github.com/Project-MONAI/MONAI.git\n!cd MONAI/\n!pip install -e '.[all]\n! PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:33:22.810087Z","iopub.execute_input":"2022-09-26T08:33:22.811029Z","iopub.status.idle":"2022-09-26T08:33:45.128513Z","shell.execute_reply.started":"2022-09-26T08:33:22.810924Z","shell.execute_reply":"2022-09-26T08:33:45.127267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nimport numpy as np\nimport pandas as pd\nimport re\nwarnings.filterwarnings(\"ignore\", category=DeprecationWarning) \nimport os\nimport shutil\nimport tempfile\nimport matplotlib.pyplot as plt\nimport PIL\nimport torch\nimport glob \nimport numpy as np\nimport pydicom\nimport cv2\nimport PIL\nfrom PIL import Image\nfrom sklearn.metrics import classification_report\n\nimport torch\nimport torchvision\nimport torchvision.transforms as tt\n\nfrom monai.apps import download_and_extract\nfrom monai.config import print_config\nfrom monai.data import decollate_batch, DataLoader,Dataset,ImageDataset\nfrom monai.metrics import ROCAUCMetric\nfrom monai.networks.nets import DenseNet121\nfrom monai.transforms import (\n    Activations,\n    EnsureChannelFirst,\n    AsDiscrete,\n    Compose,\n    LoadImage,\n    RandFlip,\n    RandRotate,\n    RandZoom,\n    ScaleIntensity,\n)\nfrom monai.utils import set_determinism\n\nprint_config()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-26T08:33:50.198988Z","iopub.execute_input":"2022-09-26T08:33:50.199374Z","iopub.status.idle":"2022-09-26T08:33:56.609366Z","shell.execute_reply.started":"2022-09-26T08:33:50.199337Z","shell.execute_reply":"2022-09-26T08:33:56.608424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"directory = os.environ.get(\"MONAI_DATA_DIRECTORY\")\nroot_dir = \"./\"\nprint(root_dir)","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:34:03.217928Z","iopub.execute_input":"2022-09-26T08:34:03.219625Z","iopub.status.idle":"2022-09-26T08:34:03.22711Z","shell.execute_reply.started":"2022-09-26T08:34:03.21957Z","shell.execute_reply":"2022-09-26T08:34:03.226081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set_determinism(seed=0)","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:34:06.782668Z","iopub.execute_input":"2022-09-26T08:34:06.783362Z","iopub.status.idle":"2022-09-26T08:34:06.791384Z","shell.execute_reply.started":"2022-09-26T08:34:06.783322Z","shell.execute_reply":"2022-09-26T08:34:06.790386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"image_list=glob.glob(\"../input/rsna-2022-cervical-spine-fracture-detection/train_images/**/*dcm\")\nplt.subplots(3, 3, figsize=(8, 8))\nfor i in range(9):\n    im = pydicom.dcmread(image_list[i])\n    im = im.pixel_array.astype(float)\n    im = cv2.resize(im, (150,150), interpolation=cv2.INTER_LINEAR)\n    rescaled_img = (np.maximum(im,0)/im.max())*255\n    fin_img_test = np.uint8(rescaled_img)\n    fin_img_test = Image.fromarray(fin_img_test)\n    arr = np.array(fin_img_test)\n    plt.subplot(3, 3, i + 1)\n    plt.imshow(arr, cmap= plt.cm.gist_heat)\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:34:10.126226Z","iopub.execute_input":"2022-09-26T08:34:10.126601Z","iopub.status.idle":"2022-09-26T08:37:06.335911Z","shell.execute_reply.started":"2022-09-26T08:34:10.126568Z","shell.execute_reply":"2022-09-26T08:37:06.33501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TEST_IMAGES_PATH= \"../input/rsna-2022-cervical-spine-fracture-detection/test_images\"","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:37:27.672265Z","iopub.execute_input":"2022-09-26T08:37:27.672648Z","iopub.status.idle":"2022-09-26T08:37:27.677615Z","shell.execute_reply.started":"2022-09-26T08:37:27.672615Z","shell.execute_reply":"2022-09-26T08:37:27.676398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_slices = glob.glob(f'{TEST_IMAGES_PATH}/*/*')\ntest_slices = [re.findall(f'{TEST_IMAGES_PATH}/(.*)/(.*).dcm', s)[0] for s in test_slices]\ndf_test_slices = pd.DataFrame(data=test_slices, columns=['StudyInstanceUID', 'Slice']).astype({'Slice': int}).sort_values(['StudyInstanceUID', 'Slice']).reset_index(drop=True)\ndf_test_slices.head()","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:37:32.218523Z","iopub.execute_input":"2022-09-26T08:37:32.21891Z","iopub.status.idle":"2022-09-26T08:37:32.540302Z","shell.execute_reply.started":"2022-09-26T08:37:32.218875Z","shell.execute_reply":"2022-09-26T08:37:32.539092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_jpg_list = glob.glob(\"../input/rsna-150-9204-25-data/test_150_rsna/test_150/*\")\nlen(img_jpg_list)","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:37:39.260133Z","iopub.execute_input":"2022-09-26T08:37:39.260506Z","iopub.status.idle":"2022-09-26T08:37:39.641371Z","shell.execute_reply.started":"2022-09-26T08:37:39.260475Z","shell.execute_reply":"2022-09-26T08:37:39.640403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test_slices[\"image\"]= img_jpg_list\n\ndf_test_slices.head()","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:37:42.757693Z","iopub.execute_input":"2022-09-26T08:37:42.758406Z","iopub.status.idle":"2022-09-26T08:37:42.770897Z","shell.execute_reply.started":"2022-09-26T08:37:42.758368Z","shell.execute_reply":"2022-09-26T08:37:42.769752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = \"../input/rsna-150-9204-25-data/train_150_rsna_9204/train_150_jpg\"","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:37:47.948861Z","iopub.execute_input":"2022-09-26T08:37:47.949538Z","iopub.status.idle":"2022-09-26T08:37:47.954529Z","shell.execute_reply.started":"2022-09-26T08:37:47.949498Z","shell.execute_reply":"2022-09-26T08:37:47.953328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_names = sorted(x for x in os.listdir(data_dir)\n                     if os.path.isdir(os.path.join(data_dir, x)))\nnum_class = len(class_names)\nimage_files = [\n    [\n        os.path.join(data_dir, class_names[i], x)\n        for x in os.listdir(os.path.join(data_dir, class_names[i]))\n    ]\n    for i in range(num_class)\n]\nnum_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)\nimage_width, image_height = PIL.Image.open(image_files_list[0]).size\n\nprint(f\"Total image count: {num_total}\")\nprint(f\"Image dimensions: {image_width} x {image_height}\")\nprint(f\"Label names: {class_names}\")\nprint(f\"Label counts: {num_each}\")","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:37:53.44589Z","iopub.execute_input":"2022-09-26T08:37:53.446507Z","iopub.status.idle":"2022-09-26T08:37:54.76706Z","shell.execute_reply.started":"2022-09-26T08:37:53.446469Z","shell.execute_reply":"2022-09-26T08:37:54.765917Z"},"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    arr = np.array(im)\n    plt.subplot(3, 3, i + 1)\n    plt.xlabel(class_names[image_class[k]])\n    plt.imshow(arr, cmap= plt.cm.gist_heat, vmin=0, vmax=255)\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:38:00.883622Z","iopub.execute_input":"2022-09-26T08:38:00.884037Z","iopub.status.idle":"2022-09-26T08:38:02.138741Z","shell.execute_reply.started":"2022-09-26T08:38:00.884001Z","shell.execute_reply":"2022-09-26T08:38:02.137777Z"},"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(\n    f\"Training count: {len(train_x)}, Validation count: \"\n    f\"{len(val_x)}, Test count: {len(test_x)}\")","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:38:16.912583Z","iopub.execute_input":"2022-09-26T08:38:16.912963Z","iopub.status.idle":"2022-09-26T08:38:16.925942Z","shell.execute_reply.started":"2022-09-26T08:38:16.91293Z","shell.execute_reply":"2022-09-26T08:38:16.924905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = Compose(\n    [\n        LoadImage(image_only=True),\n        EnsureChannelFirst(),\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    ]\n)\n\nval_transforms = Compose(\n    [LoadImage(image_only=True), EnsureChannelFirst(), ScaleIntensity()])\n\ny_pred_trans = Compose([Activations(softmax=True)])\ny_trans = Compose([AsDiscrete(to_onehot=num_class)])","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:38:24.562027Z","iopub.execute_input":"2022-09-26T08:38:24.562451Z","iopub.status.idle":"2022-09-26T08:38:24.579236Z","shell.execute_reply.started":"2022-09-26T08:38:24.562417Z","shell.execute_reply":"2022-09-26T08:38:24.578229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNADataset(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\ntrain_ds = RSNADataset(train_x, train_y, train_transforms)\ntrain_loader = DataLoader(\n    train_ds, batch_size=128, shuffle=True, num_workers=2)\n\nval_ds = RSNADataset(val_x, val_y, val_transforms)\nval_loader = DataLoader(\n    val_ds, batch_size=128, num_workers=2)\n\ntest_ds = RSNADataset(test_x, test_y, val_transforms)\ntest_loader = DataLoader(\n    test_ds, batch_size=1, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:57:32.241301Z","iopub.execute_input":"2022-09-26T08:57:32.241684Z","iopub.status.idle":"2022-09-26T08:57:32.249677Z","shell.execute_reply.started":"2022-09-26T08:57:32.241651Z","shell.execute_reply":"2022-09-26T08:57:32.248709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_images = glob.glob(\"../input/rsna-150-9204-25-data/test_150_rsna/test_150/*\")\n\nlen(test_images)","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:57:38.684239Z","iopub.execute_input":"2022-09-26T08:57:38.684612Z","iopub.status.idle":"2022-09-26T08:57:38.697735Z","shell.execute_reply.started":"2022-09-26T08:57:38.684579Z","shell.execute_reply":"2022-09-26T08:57:38.696789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNADataset_test(torch.utils.data.Dataset):\n    def __init__(self, image_files,transforms):\n        self.image_files = image_files\n        self.transforms = transforms\n        \n       \n\n    def __len__(self):\n        return len(self.image_files)\n\n    def __getitem__(self, index):\n       \n        return self.transforms(torch.Tensor(self.image_files[index]))\n        \ntest_ds_test = RSNADataset_test(test_images, val_transforms) \ntest_loader_f = DataLoader(\n    test_ds_test, batch_size=1, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:57:48.242545Z","iopub.execute_input":"2022-09-26T08:57:48.242935Z","iopub.status.idle":"2022-09-26T08:57:48.249658Z","shell.execute_reply.started":"2022-09-26T08:57:48.242895Z","shell.execute_reply":"2022-09-26T08:57:48.248584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = DenseNet121(spatial_dims=2, in_channels=1,\n                    out_channels=num_class).to(device)\nloss_function = torch.nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), 1e-5)\nmax_epochs = 10\nval_interval = 1\nauc_metric = ROCAUCMetric()","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:57:51.742128Z","iopub.execute_input":"2022-09-26T08:57:51.742632Z","iopub.status.idle":"2022-09-26T08:57:51.961124Z","shell.execute_reply.started":"2022-09-26T08:57:51.742587Z","shell.execute_reply":"2022-09-26T08:57:51.960157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_metric = -1\nbest_metric_epoch = -1\nepoch_loss_values = []\nmetric_values = []\nroot_dir=\"./\"\nfor epoch in range(max_epochs):\n    print(\"-\" * 10)\n    print(f\"epoch {epoch + 1}/{max_epochs}\")\n    model.train()\n    epoch_loss = 0\n    step = 0\n    for batch_data in train_loader:\n        step += 1\n        inputs, labels = batch_data[0].to(device), batch_data[1].to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = loss_function(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        epoch_loss += loss.item()\n        with torch.no_grad():\n            y_pred = torch.tensor([], dtype=torch.float32, device=device)\n            y = torch.tensor([], dtype=torch.long, device=device)\n            for val_data in val_loader:\n                val_images, val_labels = (\n                    val_data[0].to(device),\n                    val_data[1].to(device),\n                )\n                y_pred = torch.cat([y_pred, model(val_images)], dim=0)\n                y = torch.cat([y, val_labels], dim=0)\n            y_onehot = [y_trans(i) for i in decollate_batch(y)]\n            y_pred_act = [y_pred_trans(i) for i in decollate_batch(y_pred)]\n            auc_metric(y_pred_act, y_onehot)\n            result = auc_metric.aggregate()\n            auc_metric.reset()\n            del y_pred_act, y_onehot\n            metric_values.append(result)\n            acc_value = torch.eq(y_pred.argmax(dim=1), y)\n            acc_metric = acc_value.sum().item() / len(acc_value)\n            if result > best_metric:\n                best_metric = result\n                best_metric_epoch = epoch + 1\n                torch.save(model.state_dict(), os.path.join(\n                    root_dir, \"best_metric_model.pth\"))\nprint(\nf\"current epoch: {epoch + 1} current AUC: {result:.4f}\"\nf\" current accuracy: {acc_metric:.4f}\"\nf\" best AUC: {best_metric:.4f}\"\nf\" at epoch: {best_metric_epoch}\"\n)\n\nprint(\n    f\"train completed, best_metric: {best_metric:.4f} \"\n    f\"at epoch: {best_metric_epoch}\")","metadata":{"execution":{"iopub.status.busy":"2022-09-26T08:57:57.29933Z","iopub.execute_input":"2022-09-26T08:57:57.299698Z","iopub.status.idle":"2022-09-26T09:25:41.636414Z","shell.execute_reply.started":"2022-09-26T08:57:57.299666Z","shell.execute_reply":"2022-09-26T09:25:41.635252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\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.savefig('./monai_15Denset.jpg',bbox_inches='tight', dpi=150)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-26T09:30:27.974536Z","iopub.execute_input":"2022-09-26T09:30:27.974967Z","iopub.status.idle":"2022-09-26T09:30:28.32469Z","shell.execute_reply.started":"2022-09-26T09:30:27.974931Z","shell.execute_reply":"2022-09-26T09:30:28.32373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nmodel.load_state_dict(torch.load(\"../input/model25/best_metric_model _25.pth\",map_location=\"cpu\"))\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())\n        ","metadata":{"execution":{"iopub.status.busy":"2022-09-26T09:27:21.084132Z","iopub.execute_input":"2022-09-26T09:27:21.085142Z","iopub.status.idle":"2022-09-26T09:28:09.029568Z","shell.execute_reply.started":"2022-09-26T09:27:21.085101Z","shell.execute_reply":"2022-09-26T09:28:09.028314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(classification_report(\n    y_true, y_pred, target_names=class_names, digits=4))","metadata":{"execution":{"iopub.status.busy":"2022-09-26T09:30:50.064711Z","iopub.execute_input":"2022-09-26T09:30:50.065789Z","iopub.status.idle":"2022-09-26T09:30:50.080004Z","shell.execute_reply.started":"2022-09-26T09:30:50.065747Z","shell.execute_reply":"2022-09-26T09:30:50.078866Z"},"trusted":true},"execution_count":null,"outputs":[]}]}