{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":36363,"databundleVersionId":4050810,"sourceType":"competition"},{"sourceId":52254,"databundleVersionId":6863140,"sourceType":"competition"},{"sourceId":2786089,"sourceType":"datasetVersion","datasetId":1701116},{"sourceId":3951115,"sourceType":"datasetVersion","datasetId":1027206},{"sourceId":6069560,"sourceType":"datasetVersion","datasetId":3473850},{"sourceId":7764471,"sourceType":"datasetVersion","datasetId":4541364}],"dockerImageVersionId":30512,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Setup","metadata":{"papermill":{"duration":0.011208,"end_time":"2022-11-15T04:46:17.301819","exception":false,"start_time":"2022-11-15T04:46:17.290611","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"### Libraries","metadata":{"papermill":{"duration":0.009395,"end_time":"2022-11-15T04:46:17.321037","exception":false,"start_time":"2022-11-15T04:46:17.311642","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install python_gdcm==3.0.14","metadata":{"execution":{"iopub.status.busy":"2024-03-05T15:04:34.838799Z","iopub.execute_input":"2024-03-05T15:04:34.83916Z","iopub.status.idle":"2024-03-05T15:04:48.653807Z","shell.execute_reply.started":"2024-03-05T15:04:34.839129Z","shell.execute_reply":"2024-03-05T15:04:48.652238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install segmentation_models_pytorch==0.3.0 efficientnet_pytorch==0.7.1","metadata":{"execution":{"iopub.status.busy":"2024-03-05T15:04:48.656303Z","iopub.execute_input":"2024-03-05T15:04:48.656646Z","iopub.status.idle":"2024-03-05T15:05:05.347675Z","shell.execute_reply.started":"2024-03-05T15:04:48.656611Z","shell.execute_reply":"2024-03-05T15:05:05.34657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nimport os\n\nsys.path.append(\"../input/pytorch-segmentation-models-lib/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4\")\nsys.path.append(\"../input/timm-pytorch-image-models/pytorch-image-models-master\")\n# sys.path.append(\"../input/segmentationmoodel030/efficientnet_pytorch-0.7.1/efficientnet_pytorch-0.7.1\")\n# sys.path.append(\"../input/segmentationmoodel030/segmentation_models_pytorch-0.3.0/segmentation_models_pytorch-0.3.0\")\n","metadata":{"papermill":{"duration":1.34386,"end_time":"2022-11-15T04:46:29.791008","exception":false,"start_time":"2022-11-15T04:46:28.447148","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:05.349063Z","iopub.execute_input":"2024-03-05T15:05:05.349331Z","iopub.status.idle":"2024-03-05T15:05:05.354969Z","shell.execute_reply.started":"2024-03-05T15:05:05.349305Z","shell.execute_reply":"2024-03-05T15:05:05.354113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"../input/timm-pytorch-image-models/pytorch-image-models-master\")\nimport timm\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n%matplotlib inline\nimport matplotlib.patches as patches\nimport seaborn as sns\nsns.set(style='darkgrid', font_scale=1.6)\nimport cv2\nimport os\nfrom os import listdir\nimport re\nimport gc\nimport random\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nfrom tqdm.auto import tqdm\nfrom pprint import pprint\nfrom time import time\nimport itertools\nfrom skimage import measure\nfrom mpl_toolkits.mplot3d.art3d import Poly3DCollection\nimport nibabel as nib\nfrom glob import glob\nimport warnings\n\nimport zipfile\nfrom scipy import ndimage\nfrom sklearn.model_selection import train_test_split\nfrom joblib import Parallel, delayed\nfrom PIL import Image\nfrom dipy.denoise.nlmeans import nlmeans\nfrom dipy.denoise.noise_estimate import estimate_sigma\nfrom skimage import exposure\n\n# Pytorch\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torch.nn.functional as F\nimport nibabel as nib\nimport pydicom as dicom\nimport gc \nimport segmentation_models_pytorch as smp\n","metadata":{"_kg_hide-input":true,"papermill":{"duration":7.335592,"end_time":"2022-11-15T04:46:37.136762","exception":false,"start_time":"2022-11-15T04:46:29.80117","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:05.357305Z","iopub.execute_input":"2024-03-05T15:05:05.357613Z","iopub.status.idle":"2024-03-05T15:05:12.054018Z","shell.execute_reply.started":"2024-03-05T15:05:05.357577Z","shell.execute_reply":"2024-03-05T15:05:12.053184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuring GPU/TPU","metadata":{}},{"cell_type":"code","source":"class CFG:\n    seed=42\n    device='GPU' # ['TPU', 'GPU']\n    nprocs=1 # [1, 8]\n    num_workers=2\n    valid_bs=32\n    fold_num=5 \n    \n    target_cols=[\"L1\", \"L2\", \"L3\", \"L4\", \"L5\",\"OT\"]\n    num_classes=8 \n    \n    normalize_mean=[0.4824, 0.4824, 0.4824] \n    normalize_std=[0.22, 0.22, 0.22] \n    \n    fold_list=[0]\n\n    model_arch=\"efficientnet-b0\" \n    img_size=512 \n    croped_img_size = 320 # 裁剪后的图片尺寸\n    weight_path = f\"../input/cervical-review/efficientnet-b0_109_fold0_epoch13.pth\" \n    \n# Config device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')    \ndatadir = \"../input/rsna-2022-cervical-spine-fracture-detection\"","metadata":{"papermill":{"duration":0.083718,"end_time":"2022-11-15T04:46:37.230871","exception":false,"start_time":"2022-11-15T04:46:37.147153","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:12.055173Z","iopub.execute_input":"2024-03-05T15:05:12.055468Z","iopub.status.idle":"2024-03-05T15:05:12.087973Z","shell.execute_reply.started":"2024-03-05T15:05:12.05544Z","shell.execute_reply":"2024-03-05T15:05:12.086956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading Dicom files and reading them","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True \n\nseed_everything(CFG.seed)\n\ndef load_dicom(path):\n    \"\"\"\n    This supports loading both regular and compressed JPEG images. \n    See the first sell with `pip install` commands for the necessary dependencies\n    \"\"\"\n    img = dicom.dcmread(path)\n    img.PhotometricInterpretation = 'YBR_FULL'\n    data = img.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    # data = (data * 255).astype(np.uint8)\n    return data","metadata":{"papermill":{"duration":0.029817,"end_time":"2022-11-15T04:46:37.270589","exception":false,"start_time":"2022-11-15T04:46:37.240772","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:12.089297Z","iopub.execute_input":"2024-03-05T15:05:12.089578Z","iopub.status.idle":"2024-03-05T15:05:12.130907Z","shell.execute_reply.started":"2024-03-05T15:05:12.089553Z","shell.execute_reply":"2024-03-05T15:05:12.130032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\")\ntest_df = pd.read_csv('/kaggle/input/spinal-lumbar/test.csv')\n\ndebug = False\nif len(test_df)==3:\n    debug = True\n    \n    # Fix mismatch with test_images folder\n    test_df = pd.DataFrame(columns = ['row_id','StudyInstanceUID','prediction_type'])\n    for i in ['1.2.826.0.1.3680043.22327','1.2.826.0.1.3680043.25399','1.2.826.0.1.3680043.5876']:\n        for j in [\"L1\", \"L2\", \"L3\", \"L4\", \"L5\",'patient_overall']:\n            test_df = test_df.append({'row_id':i+'_'+j,'StudyInstanceUID':i,'prediction_type':j},ignore_index=True)\n    \n    # Sample submission\n    ss = pd.DataFrame(test_df['row_id'])\n    ss['fractured'] = 0.5\n    print(test_df.shape)","metadata":{"papermill":{"duration":0.076552,"end_time":"2022-11-15T04:46:37.357027","exception":false,"start_time":"2022-11-15T04:46:37.280475","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:12.134524Z","iopub.execute_input":"2024-03-05T15:05:12.134798Z","iopub.status.idle":"2024-03-05T15:05:12.184766Z","shell.execute_reply.started":"2024-03-05T15:05:12.134775Z","shell.execute_reply":"2024-03-05T15:05:12.18389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_study_list = np.unique(test_df[\"StudyInstanceUID\"].values).tolist()\ntest_study_list[:3]","metadata":{"papermill":{"duration":0.023823,"end_time":"2022-11-15T04:46:37.391119","exception":false,"start_time":"2022-11-15T04:46:37.367296","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:12.186084Z","iopub.execute_input":"2024-03-05T15:05:12.186374Z","iopub.status.idle":"2024-03-05T15:05:12.194175Z","shell.execute_reply.started":"2024-03-05T15:05:12.186345Z","shell.execute_reply":"2024-03-05T15:05:12.193077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# study_id_list = []\n# slice_num_list = []\n# for study_name in test_study_list:\n#     slice_file_list = os.listdir(f\"{datadir}/test_images/{study_name}\")\n#     slice_cnt = len(slice_file_list)\n    \n#     study_id_list.extend([study_name]*slice_cnt)\n#     slice_num_list.extend([int(x.replace(\".dcm\",\"\")) for x in slice_file_list])\n# print(len(study_id_list), len(slice_num_list))\n\n# all_slice_df = pd.DataFrame({\"StudyInstanceUID\":study_id_list, \"slice_num\":slice_num_list})\n# all_slice_df = all_slice_df.sort_values([\"StudyInstanceUID\", \"slice_num\"]).reset_index(drop=True)\n# all_slice_df.to_csv(f\"./all_slice_df.csv\", index=False)\n# print(all_slice_df.shape)\n# all_slice_df.head(3)","metadata":{"papermill":{"duration":0.017758,"end_time":"2022-11-15T04:46:37.419103","exception":false,"start_time":"2022-11-15T04:46:37.401345","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:12.195342Z","iopub.execute_input":"2024-03-05T15:05:12.195644Z","iopub.status.idle":"2024-03-05T15:05:12.201914Z","shell.execute_reply.started":"2024-03-05T15:05:12.19561Z","shell.execute_reply":"2024-03-05T15:05:12.2012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_slice_list = []\nfor file_name in test_study_list:\n    image_path_list = glob(f\"{datadir}/test_images/{file_name}/*\")\n    image_path_list = sorted(image_path_list, key=lambda x:int(x.split(\"/\")[-1].replace(\".dcm\",\"\")))\n    for path_idx in range(len(image_path_list)):\n        path1 = \"nofile\" if path_idx-1 < 0 else image_path_list[path_idx-1].replace(f\"{datadir}/test_images/\", \"\")\n        path2 = image_path_list[path_idx].replace(f\"{datadir}/test_images/\", \"\")\n        path3 = \"nofile\" if path_idx+1 >= len(image_path_list) else image_path_list[path_idx+1].replace(f\"{datadir}/test_images/\", \"\")\n        slice_num = int(path2.split(\"/\")[-1].replace(\".dcm\",\"\"))\n        all_slice_list.append([f\"{file_name}_{slice_num}\", file_name, slice_num, path1, path2, path3])","metadata":{"papermill":{"duration":0.132564,"end_time":"2022-11-15T04:46:37.561495","exception":false,"start_time":"2022-11-15T04:46:37.428931","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:12.205962Z","iopub.execute_input":"2024-03-05T15:05:12.206212Z","iopub.status.idle":"2024-03-05T15:05:12.547937Z","shell.execute_reply.started":"2024-03-05T15:05:12.20619Z","shell.execute_reply":"2024-03-05T15:05:12.547098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"slice_df = pd.DataFrame(all_slice_list, columns=[\"id\", \"StudyInstanceUID\", \"slice_num\", \"path1\", \"path2\", \"path3\"])\nslice_df","metadata":{"papermill":{"duration":0.034001,"end_time":"2022-11-15T04:46:37.605575","exception":false,"start_time":"2022-11-15T04:46:37.571574","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:12.54916Z","iopub.execute_input":"2024-03-05T15:05:12.549944Z","iopub.status.idle":"2024-03-05T15:05:12.572856Z","shell.execute_reply.started":"2024-03-05T15:05:12.549886Z","shell.execute_reply":"2024-03-05T15:05:12.571971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DataSet","metadata":{"papermill":{"duration":0.010032,"end_time":"2022-11-15T04:46:37.625818","exception":false,"start_time":"2022-11-15T04:46:37.615786","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\nclass VoxelDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        im2 = load_dicom(f\"{datadir}/test_images/{row['path2']}\")   # 512*512  \n        im2h = im2.shape[0]\n        im2w = im2.shape[1]\n\n        im1 = load_dicom(f\"{datadir}/test_images/{row['path1']}\") if row['path1'] != \"nofile\" else np.zeros((im2h, im2w))  # 512*512                                                       \n        im3 = load_dicom(f\"{datadir}/test_images/{row['path3']}\") if row['path3'] != \"nofile\" else np.zeros((im2h, im2w))  # 512*512  \n\n        if im1.shape !=  (im2h, im2w):\n            im1 = cv2.resize(im1, (im2w, im2h))\n        if im3.shape !=  (im2h, im2w):\n            im3 = cv2.resize(im3, (im2w, im2h)) \n        image_list = [im1, im2, im3]\n        image = np.stack(image_list, axis=2) # 512*512*3; 0-1\n\n        # transform\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n        \n        # image = image/255.0\n        image = np.transpose(image, (2, 0, 1)) # 3*img_size*img_size; 0-1\n        return torch.from_numpy(image), row['StudyInstanceUID'], row['slice_num'] ","metadata":{"papermill":{"duration":0.022797,"end_time":"2022-11-15T04:46:37.658778","exception":false,"start_time":"2022-11-15T04:46:37.635981","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:12.574059Z","iopub.execute_input":"2024-03-05T15:05:12.574373Z","iopub.status.idle":"2024-03-05T15:05:12.585191Z","shell.execute_reply.started":"2024-03-05T15:05:12.574345Z","shell.execute_reply":"2024-03-05T15:05:12.58418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from albumentations import CenterCrop, Resize, RandomCrop, GaussianBlur, JpegCompression, Downscale, ElasticTransform, Compose\nimport albumentations\nfrom albumentations.pytorch import ToTensorV2\n\ndef get_transforms(data):\n    if data == 'valid':\n        return Compose([\n            Resize(CFG.img_size, CFG.img_size, interpolation=cv2.INTER_NEAREST),\n        ])","metadata":{"papermill":{"duration":0.855098,"end_time":"2022-11-15T04:46:38.524133","exception":false,"start_time":"2022-11-15T04:46:37.669035","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:12.586349Z","iopub.execute_input":"2024-03-05T15:05:12.58668Z","iopub.status.idle":"2024-03-05T15:05:13.486359Z","shell.execute_reply.started":"2024-03-05T15:05:12.58663Z","shell.execute_reply":"2024-03-05T15:05:13.485367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nif debug:\n    from pylab import rcParams\n    dataset_show = VoxelDataset(\n        slice_df, \n        get_transforms(\"valid\") # None, get_transforms(\"train\")\n        )\n    rcParams['figure.figsize'] = 30,20\n    for i in range(2):\n        f, axarr = plt.subplots(1,3)\n        idx = np.random.randint(0, len(dataset_show))\n        img, file_name, n_slice= dataset_show[idx]\n        # axarr[p].imshow(img) # transform=None\n        axarr[0].imshow(img[0]); plt.axis('OFF');\n        axarr[1].imshow(img[1]); plt.axis('OFF');\n        axarr[2].imshow(img[2]); plt.axis('OFF');","metadata":{"papermill":{"duration":2.152162,"end_time":"2022-11-15T04:46:40.686811","exception":false,"start_time":"2022-11-15T04:46:38.534649","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:13.487885Z","iopub.execute_input":"2024-03-05T15:05:13.488196Z","iopub.status.idle":"2024-03-05T15:05:17.002732Z","shell.execute_reply.started":"2024-03-05T15:05:13.488169Z","shell.execute_reply":"2024-03-05T15:05:17.00173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"papermill":{"duration":0.03502,"end_time":"2022-11-15T04:46:40.758373","exception":false,"start_time":"2022-11-15T04:46:40.723353","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\n\ndef build_model():\n    model = smp.Unet(\n        encoder_name=CFG.model_arch,    # choose encoder, e.g. mobilenet_v2 or efficientnet-b7\n        encoder_weights=\"imagenet\",     # use `imagenet` pre-trained weights for encoder initialization\n        in_channels=3,                  # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n        classes=CFG.num_classes,        # model output channels (number of classes in your dataset)\n        activation=None,\n    )\n    model.to(device)\n    return model\n\ndef load_model(path):\n    model = build_model()\n    model.load_state_dict(torch.load(path)[\"model\"])\n    model.eval()\n    return model","metadata":{"papermill":{"duration":0.044067,"end_time":"2022-11-15T04:46:40.835761","exception":false,"start_time":"2022-11-15T04:46:40.791694","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:17.003942Z","iopub.execute_input":"2024-03-05T15:05:17.004425Z","iopub.status.idle":"2024-03-05T15:05:17.010367Z","shell.execute_reply.started":"2024-03-05T15:05:17.004393Z","shell.execute_reply":"2024-03-05T15:05:17.009572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"slice_class_list = []\nvoxel_crop_list = []\ndef crop_voxel(voxel_mask, last_f_name):\n    area_thr = 10\n    # x\n    x_list = []\n    length = voxel_mask.shape[0]\n    for i in range(length):\n        if torch.count_nonzero(voxel_mask[i]).item() >= area_thr:\n            x_list.append(i)\n            break\n    else:\n        x_list.append(0)\n\n    for i in range(length-1, -1, -1):\n        if torch.count_nonzero(voxel_mask[i]).item() >= area_thr:\n            x_list.append(i)\n            break\n    else:\n        x_list.append(length-1)\n\n    # y\n    y_list = []\n    length = voxel_mask.shape[1]\n    for i in range(length):\n        if torch.count_nonzero(voxel_mask[:, i]).item() >= area_thr:\n            y_list.append(i)\n            break\n    else:\n        y_list.append(0)\n\n    for i in range(length-1, -1, -1):\n        if torch.count_nonzero(voxel_mask[:, i]).item() >= area_thr:\n            y_list.append(i)\n            break\n    else:\n        y_list.append(length-1)\n\n    # z\n    z_list = []\n    length = voxel_mask.shape[2]\n    for i in range(length):\n        if torch.count_nonzero(voxel_mask[:, :, i]).item() >= area_thr:\n            z_list.append(i)\n            break\n    else:\n        z_list.append(0)\n\n    for i in range(length-1, -1, -1):\n        if torch.count_nonzero(voxel_mask[:, :, i]).item() >= area_thr:\n            z_list.append(i)\n            break\n    else:\n        z_list.append(length-1)\n    # croped_voxel = voxels[x_list[0]:x_list[1]+1, y_list[0]:y_list[1]+1, z_list[0]:z_list[1]+1]\n    try:\n        croped_voxel_mask = voxel_mask[x_list[0]:x_list[1]+1, y_list[0]:y_list[1]+1, z_list[0]:z_list[1]+1]\n    except:\n        print(f\"last_f_name:{last_f_name}, voxel_mask.shape:{voxel_mask.shape}, x_list:{x_list}, y_list:{y_list}, z_list:{z_list}\")\n        x_list = [0, voxel_mask.shape[0]-1]; y_list = [0, voxel_mask.shape[1]-1]; z_list = [0, voxel_mask.shape[2]-1]\n        croped_voxel_mask = voxel_mask\n    voxel_crop_list.append([last_f_name, voxel_mask.shape[1], x_list[0], x_list[1]+1, y_list[0], y_list[1]+1, z_list[0], z_list[1]+1])\n\n    # croped_voxel = croped_voxel.to('cpu').numpy() # bs*img_size*img_size; 0-8 classes\n    croped_voxel_mask = croped_voxel_mask.to('cpu').numpy().astype(np.uint8) # bs*img_size*img_size; 0-8 classes\n    for x_idx in range(croped_voxel_mask.shape[0]):\n        slice_mask = croped_voxel_mask[x_idx]\n\n        unique, counts = np.unique(slice_mask, return_counts=True)\n        if len(unique) == 1 and unique[0] == 0:\n            slice_class_list.append([last_f_name, x_idx, x_idx+x_list[0], 0])\n        elif unique[0] == 0:\n            unique = unique[1:]\n            counts = counts[1:]\n            slice_class_list.append([last_f_name, x_idx, x_idx+x_list[0]+1, unique[counts.argmax()]])\n        else:\n            slice_class_list.append([last_f_name, x_idx, x_idx+x_list[0]+1, unique[counts.argmax()]])\n        \n    return None, croped_voxel_mask","metadata":{"papermill":{"duration":0.055375,"end_time":"2022-11-15T04:46:40.924097","exception":false,"start_time":"2022-11-15T04:46:40.868722","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:17.011698Z","iopub.execute_input":"2024-03-05T15:05:17.012048Z","iopub.status.idle":"2024-03-05T15:05:17.0354Z","shell.execute_reply.started":"2024-03-05T15:05:17.012012Z","shell.execute_reply":"2024-03-05T15:05:17.034452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = VoxelDataset(slice_df, transform=get_transforms(\"valid\")) # get_transforms(\"valid\")\ntest_loader = DataLoader(test_dataset, batch_size=CFG.valid_bs, shuffle=False, num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n\nmodel = load_model(CFG.weight_path)\nmodel.eval()\nlast_f_name = \"\"\nvoxel_mask = []\n# voxels = []","metadata":{"papermill":{"duration":3.702128,"end_time":"2022-11-15T04:46:44.65995","exception":false,"start_time":"2022-11-15T04:46:40.957822","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:17.03662Z","iopub.execute_input":"2024-03-05T15:05:17.036975Z","iopub.status.idle":"2024-03-05T15:05:19.533389Z","shell.execute_reply.started":"2024-03-05T15:05:17.036945Z","shell.execute_reply":"2024-03-05T15:05:19.532338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for step, (images, file_names, n_slice) in tqdm(enumerate(test_loader),total=len(test_loader)):\n    images = images.to(device, dtype=torch.float) # bs*3*image_size*image_size\n    batch_size = images.size(0)\n    with torch.no_grad():\n        y_pred = model(images) # [B, 8, H, W]\n    y_pred = y_pred.sigmoid()\n    slice_mask_max = torch.max(y_pred, 1) # bs*img_size*img_size\n    slice_mask = torch.where((slice_mask_max.values)>0.5, slice_mask_max.indices+1, 0) # bs*img_size*img_size; 0-8 classes\n    slice_mask = torch.where(slice_mask==8,0,slice_mask).type(torch.uint8)\n    # slice_mask = slice_mask.to('cpu').numpy().astype(np.uint8) # bs*img_size*img_size; 0-8 classes\n    # slice_image = images[:, 1, :, :] # bs*img_size*img_size\n\n    start_idx = 0\n    for bs_idx in range(batch_size):\n        f_name = file_names[bs_idx]\n        if f_name != last_f_name:\n            voxel_mask.append(slice_mask[start_idx:bs_idx])\n            # voxels.append(slice_image[start_idx:bs_idx])\n            voxel_mask = torch.cat(voxel_mask, dim=0) # n_slice*img_size*img_size; 0-8 classes\n            # voxels = torch.cat(voxels, dim=0) # n_slice*img_size*img_size\n            if len(voxel_mask) > 0:\n                croped_voxel, croped_voxel_mask = crop_voxel(voxel_mask, last_f_name)\n            last_f_name = f_name\n            start_idx = bs_idx\n            voxel_mask = []\n            # voxels = []\n        elif bs_idx == batch_size-1:\n            voxel_mask.append(slice_mask[start_idx:batch_size])\n            # voxels.append(slice_image[start_idx:batch_size])\nvoxel_mask = torch.cat(voxel_mask, dim=0)\nif len(voxel_mask) > 0:\n    croped_voxel, croped_voxel_mask = crop_voxel(voxel_mask, last_f_name)","metadata":{"papermill":{"duration":52.804783,"end_time":"2022-11-15T04:47:37.498584","exception":false,"start_time":"2022-11-15T04:46:44.693801","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:19.534881Z","iopub.execute_input":"2024-03-05T15:05:19.535269Z","iopub.status.idle":"2024-03-05T15:05:58.174118Z","shell.execute_reply.started":"2024-03-05T15:05:19.535232Z","shell.execute_reply":"2024-03-05T15:05:58.172962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"voxel_crop_df = pd.DataFrame(voxel_crop_list, columns=[\"StudyInstanceUID\", \"before_image_size\", \"x0\", \"x1\", \"y0\", \"y1\", \"z0\", \"z1\"]).sort_values(by=[\"StudyInstanceUID\"])\nvoxel_crop_df.to_csv(f\"voxel_crop.csv\", index=False)\nprint(voxel_crop_df.shape)\nvoxel_crop_df.head(3)","metadata":{"papermill":{"duration":0.088776,"end_time":"2022-11-15T04:47:37.640986","exception":false,"start_time":"2022-11-15T04:47:37.55221","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.175833Z","iopub.execute_input":"2024-03-05T15:05:58.176259Z","iopub.status.idle":"2024-03-05T15:05:58.197165Z","shell.execute_reply.started":"2024-03-05T15:05:58.176219Z","shell.execute_reply":"2024-03-05T15:05:58.196306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"slice_class_df = pd.DataFrame(slice_class_list, columns=[\"StudyInstanceUID\", \"new_slice_num\", \"old_slice_num\", \"vertebra_class\"]).sort_values(by=[\"StudyInstanceUID\", \"new_slice_num\"])\nslice_class_df.to_csv(f\"slice_class.csv\", index=False)\nprint(slice_class_df.shape)\nslice_class_df.head(3)","metadata":{"papermill":{"duration":0.088188,"end_time":"2022-11-15T04:47:37.781941","exception":false,"start_time":"2022-11-15T04:47:37.693753","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.198619Z","iopub.execute_input":"2024-03-05T15:05:58.198972Z","iopub.status.idle":"2024-03-05T15:05:58.236678Z","shell.execute_reply.started":"2024-03-05T15:05:58.198943Z","shell.execute_reply":"2024-03-05T15:05:58.235794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"slice_df","metadata":{"papermill":{"duration":0.069153,"end_time":"2022-11-15T04:47:37.903282","exception":false,"start_time":"2022-11-15T04:47:37.834129","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.23794Z","iopub.execute_input":"2024-03-05T15:05:58.238546Z","iopub.status.idle":"2024-03-05T15:05:58.252848Z","shell.execute_reply.started":"2024-03-05T15:05:58.238509Z","shell.execute_reply":"2024-03-05T15:05:58.251835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_df = []\nfor idx, study_id, _, x0, x1, _, _, _, _, in tqdm(voxel_crop_df.itertuples(), total=len(voxel_crop_df)):\n    one_study = slice_df[slice_df[\"StudyInstanceUID\"] == study_id][[\"id\", \"StudyInstanceUID\", \"slice_num\"]].reset_index(drop=True)\n    new_df.append(one_study[x0:x1])\nnew_df = pd.concat(new_df, axis=0).reset_index(drop=True)\nprint(new_df.shape)\nnew_df.head(3)","metadata":{"papermill":{"duration":0.092351,"end_time":"2022-11-15T04:47:38.030042","exception":false,"start_time":"2022-11-15T04:47:37.937691","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.254283Z","iopub.execute_input":"2024-03-05T15:05:58.254966Z","iopub.status.idle":"2024-03-05T15:05:58.29249Z","shell.execute_reply.started":"2024-03-05T15:05:58.25493Z","shell.execute_reply":"2024-03-05T15:05:58.291605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_df = new_df.merge(voxel_crop_df, on=\"StudyInstanceUID\", how=\"left\") # merge study_crop_df\nprint(new_df.shape)\ndisplay(new_df.head(3))\nassert len(slice_class_df) == len(new_df)","metadata":{"papermill":{"duration":0.06074,"end_time":"2022-11-15T04:47:38.126331","exception":false,"start_time":"2022-11-15T04:47:38.065591","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.293634Z","iopub.execute_input":"2024-03-05T15:05:58.293958Z","iopub.status.idle":"2024-03-05T15:05:58.315047Z","shell.execute_reply.started":"2024-03-05T15:05:58.293926Z","shell.execute_reply":"2024-03-05T15:05:58.314009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_slice_df = pd.concat([new_df, slice_class_df[[\"new_slice_num\", \"vertebra_class\"]]], axis=1)\nprint(new_slice_df.shape)\nnew_slice_df.head(3)","metadata":{"papermill":{"duration":0.058809,"end_time":"2022-11-15T04:47:38.21998","exception":false,"start_time":"2022-11-15T04:47:38.161171","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.316376Z","iopub.execute_input":"2024-03-05T15:05:58.316737Z","iopub.status.idle":"2024-03-05T15:05:58.332192Z","shell.execute_reply.started":"2024-03-05T15:05:58.316699Z","shell.execute_reply":"2024-03-05T15:05:58.3312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_num = 24\nvertebrae_df_list = []\nfor study_id in tqdm(np.unique(new_slice_df[\"StudyInstanceUID\"])):\n    one_study = new_slice_df[new_slice_df[\"StudyInstanceUID\"] == study_id].reset_index(drop=True)\n    for cid in range(1, 8):\n        one_study_cid = one_study[one_study[\"vertebra_class\"] == cid].reset_index(drop=True)\n        if len(one_study_cid) >= sample_num:\n            sample_index = np.linspace(0, len(one_study_cid)-1, sample_num, dtype=int)\n            one_study_cid = one_study_cid.iloc[sample_index].reset_index(drop=True)\n        if len(one_study_cid) < 5:\n            continue\n        slice_num_list = one_study_cid[\"slice_num\"].values.tolist()\n        arow = one_study_cid.iloc[0]\n        vertebrae_df_list.append([f\"{study_id}_{cid}\", study_id, cid, slice_num_list, arow[\"before_image_size\"], \\\n            arow[\"x0\"], arow[\"x1\"], arow[\"y0\"], arow[\"y1\"], arow[\"z0\"], arow[\"z1\"]])","metadata":{"papermill":{"duration":0.117096,"end_time":"2022-11-15T04:47:38.371958","exception":false,"start_time":"2022-11-15T04:47:38.254862","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.333365Z","iopub.execute_input":"2024-03-05T15:05:58.333637Z","iopub.status.idle":"2024-03-05T15:05:58.389273Z","shell.execute_reply.started":"2024-03-05T15:05:58.333612Z","shell.execute_reply":"2024-03-05T15:05:58.388381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vertebrae_df = pd.DataFrame(vertebrae_df_list, columns=[\"study_cid\", \"StudyInstanceUID\", \"cid\", \"slice_num_list\", \\\n    \"before_image_size\", \"x0\", \"x1\", \"y0\", \"y1\", \"z0\", \"z1\" ])\nvertebrae_df.to_pickle(f\"vertebrae_df.pkl\")    \nprint(vertebrae_df.shape) #\nvertebrae_df.head(3)","metadata":{"papermill":{"duration":0.057121,"end_time":"2022-11-15T04:47:38.464684","exception":false,"start_time":"2022-11-15T04:47:38.407563","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.390625Z","iopub.execute_input":"2024-03-05T15:05:58.390958Z","iopub.status.idle":"2024-03-05T15:05:58.409043Z","shell.execute_reply.started":"2024-03-05T15:05:58.390926Z","shell.execute_reply":"2024-03-05T15:05:58.408056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del model","metadata":{"execution":{"iopub.status.busy":"2024-03-05T15:05:58.410102Z","iopub.execute_input":"2024-03-05T15:05:58.410405Z","iopub.status.idle":"2024-03-05T15:05:58.416955Z","shell.execute_reply.started":"2024-03-05T15:05:58.410377Z","shell.execute_reply":"2024-03-05T15:05:58.416175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference Class","metadata":{"papermill":{"duration":0.035372,"end_time":"2022-11-15T04:47:38.53505","exception":false,"start_time":"2022-11-15T04:47:38.499678","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import gc\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-03-05T15:05:58.423194Z","iopub.execute_input":"2024-03-05T15:05:58.423472Z","iopub.status.idle":"2024-03-05T15:05:58.734972Z","shell.execute_reply.started":"2024-03-05T15:05:58.423446Z","shell.execute_reply":"2024-03-05T15:05:58.733945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = vertebrae_df\nCFG.img_size = 384\nCFG.valid_bs = 8 # 14\nCFG.seq_len = 24\nCFG.dropout=0.1\nCFG.gpu_parallel=False\n# tf_efficientnetv2_s, resnest50d\nCFG.archs_list=[\n#     \"tf_efficientnetv2_s\",\n#     \"tf_efficientnetv2_s\",\n    \n    \"resnest50d\",\n    \"resnest50d\",\n    \"resnest50d\",\n] \n\n\nCFG.weights_list = [\n#     \"../input/loadmodel/tf_efficientnetv2_s_405_fold0_epoch8.pth\",\n#     \"../input/loadmodel/tf_efficientnetv2_s_405_fold1_epoch9.pth\",\n#     \"../input/loadmodel/tf_efficientnetv2_s_405_fold2_epoch8.pth\",\n    \n    \"../input/cervical-review/resnest50d_406_fold0_epoch13.pth\",\n    \"../input/cervical-review/resnest50d_406_fold1_epoch13.pth\",\n    \"../input/cervical-review/resnest50d_406_fold2_epoch13.pth\",\n]\n\n\n\nCFG.fillna_number = 0.10\n\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"papermill":{"duration":0.045493,"end_time":"2022-11-15T04:47:38.615409","exception":false,"start_time":"2022-11-15T04:47:38.569916","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.736167Z","iopub.execute_input":"2024-03-05T15:05:58.736515Z","iopub.status.idle":"2024-03-05T15:05:58.745384Z","shell.execute_reply.started":"2024-03-05T15:05:58.736483Z","shell.execute_reply":"2024-03-05T15:05:58.744424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True \n\nseed_everything(CFG.seed)\n\ndef load_dicom(path):\n    \"\"\"\n    This supports loading both regular and compressed JPEG images. \n    See the first sell with `pip install` commands for the necessary dependencies\n    \"\"\"\n    img = dicom.dcmread(path)\n    img.PhotometricInterpretation = 'YBR_FULL'\n    data = img.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    # data = (data * 255).astype(np.uint8)\n    return data","metadata":{"papermill":{"duration":0.047373,"end_time":"2022-11-15T04:47:38.698354","exception":false,"start_time":"2022-11-15T04:47:38.650981","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.746535Z","iopub.execute_input":"2024-03-05T15:05:58.746812Z","iopub.status.idle":"2024-03-05T15:05:58.755757Z","shell.execute_reply.started":"2024-03-05T15:05:58.746787Z","shell.execute_reply":"2024-03-05T15:05:58.754811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\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[\"StudyInstanceUID\"]\n        slice_num_list = row['slice_num_list']\n        before_image_size = row[\"before_image_size\"]\n        y0 = row[\"y0\"]; y1 = row[\"y1\"];\n        z0 = row[\"z0\"]; z1 = row[\"z1\"];\n\n        slice_list = []\n        for s_num in slice_num_list:\n            path = f\"{datadir}/test_images/{study_id}/{s_num}.dcm\"\n            img = load_dicom(path)\n            if len(slice_list) == 0:\n                imgh = img.shape[0]\n                imgw = img.shape[1]\n            elif img.shape != (imgh, imgw):\n                img = cv2.resize(img,(imgh,imgw))\n\n            slice_list.append(img)\n        for _ in range(CFG.seq_len - len(slice_list)):\n            slice_list.append(np.zeros((imgh,imgw)))\n\n        image = np.stack(slice_list, axis=2) # 512*512*seq_len; 0-1\n        image = cv2.resize(image, (before_image_size, before_image_size))\n        image = image[y0:y1, z0:z1, :]\n\n        # transform\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n\n        image = np.transpose(image, (2, 0, 1)) # seq_len*img_size*img_size; 0-1\n        return torch.from_numpy(image)","metadata":{"papermill":{"duration":0.051517,"end_time":"2022-11-15T04:47:38.784964","exception":false,"start_time":"2022-11-15T04:47:38.733447","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.757091Z","iopub.execute_input":"2024-03-05T15:05:58.757726Z","iopub.status.idle":"2024-03-05T15:05:58.76897Z","shell.execute_reply.started":"2024-03-05T15:05:58.757692Z","shell.execute_reply":"2024-03-05T15:05:58.768246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from albumentations import Resize, RandomCrop\nimport albumentations\nfrom albumentations.pytorch import ToTensorV2\n\ndef get_transforms(*, data):\n    if data == 'valid':\n        return Compose([\n            Resize(CFG.img_size, CFG.img_size),\n        ])","metadata":{"papermill":{"duration":0.046269,"end_time":"2022-11-15T04:47:38.866167","exception":false,"start_time":"2022-11-15T04:47:38.819898","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.770001Z","iopub.execute_input":"2024-03-05T15:05:58.770259Z","iopub.status.idle":"2024-03-05T15:05:58.784512Z","shell.execute_reply.started":"2024-03-05T15:05:58.770235Z","shell.execute_reply":"2024-03-05T15:05:58.783778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pylab import rcParams\ndataset_show = TestDataset(\n    train_df,\n    transform=get_transforms(data='valid') # None, get_transforms(data='check')\n    )\nrcParams['figure.figsize'] = 30,20\nfor i in range(2):\n    f, axarr = plt.subplots(1,5)\n    idx = np.random.randint(0, len(dataset_show))\n    img = dataset_show[idx]\n    # axarr[p].imshow(img) # transform=None\n    axarr[0].imshow(img[0]); plt.axis('OFF');\n    axarr[1].imshow(img[1]); plt.axis('OFF');\n    axarr[2].imshow(img[2]); plt.axis('OFF');\n    axarr[3].imshow(img[3]); plt.axis('OFF');\n    axarr[4].imshow(img[4]); plt.axis('OFF');","metadata":{"papermill":{"duration":2.695271,"end_time":"2022-11-15T04:47:41.596745","exception":false,"start_time":"2022-11-15T04:47:38.901474","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:05:58.785499Z","iopub.execute_input":"2024-03-05T15:05:58.785765Z","iopub.status.idle":"2024-03-05T15:06:02.928711Z","shell.execute_reply.started":"2024-03-05T15:05:58.785741Z","shell.execute_reply":"2024-03-05T15:06:02.92778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2024-03-05T15:06:02.929842Z","iopub.execute_input":"2024-03-05T15:06:02.930131Z","iopub.status.idle":"2024-03-05T15:06:03.961148Z","shell.execute_reply.started":"2024-03-05T15:06:02.930105Z","shell.execute_reply":"2024-03-05T15:06:03.959882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nfrom itertools import repeat\n\nclass SpatialDropout(nn.Module):\n    def __init__(self, drop=0.5):\n        super(SpatialDropout, self).__init__()\n        self.drop = drop\n        \n    def forward(self, inputs, noise_shape=None):\n        \"\"\"\n        @param: inputs, tensor\n        @param: noise_shape, tuple\n        \"\"\"\n        outputs = inputs.clone()\n        if noise_shape is None:\n            noise_shape = (inputs.shape[0], *repeat(1, inputs.dim()-2), inputs.shape[-1]) \n        \n        self.noise_shape = noise_shape\n        if not self.training or self.drop == 0:\n            return inputs\n        else:\n            noises = self._make_noises(inputs)\n            if self.drop == 1:\n                noises.fill_(0.0)\n            else:\n                noises.bernoulli_(1 - self.drop).div_(1 - self.drop)\n            noises = noises.expand_as(inputs)    \n            outputs.mul_(noises)\n            return outputs\n            \n    def _make_noises(self, inputs):\n        return inputs.new().resize_(self.noise_shape)\n\n\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\n\nfrom typing import Dict, Optional\n \nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom torch import Tensor\n\n\n    \nclass MLPAttentionNetwork(nn.Module):\n \n    def __init__(self, hidden_dim, attention_dim=None):\n        super(MLPAttentionNetwork, self).__init__()\n \n        self.hidden_dim = hidden_dim\n        self.attention_dim = attention_dim\n        if self.attention_dim is None:\n            self.attention_dim = self.hidden_dim\n        # W * x + b\n        self.proj_w = nn.Linear(self.hidden_dim, self.attention_dim, bias=True)\n        # v.T\n        self.proj_v = nn.Linear(self.attention_dim, 1, bias=False)\n \n    def forward(self, x):\n        \"\"\"\n        :param x: seq_len, batch_size, hidden_dim\n        :return: batch_size * seq_len, batch_size * hidden_dim\n        \"\"\"\n        # print(f\"x shape:{x.shape}\")\n        batch_size, seq_len, _ = x.size()\n        # flat_inputs = x.reshape(-1, self.hidden_dim) # (batch_size*seq_len, hidden_dim)\n        # print(f\"flat_inputs shape:{flat_inputs.shape}\")\n        \n        H = torch.tanh(self.proj_w(x)) # (batch_size, seq_len, hidden_dim)\n        # print(f\"H shape:{H.shape}\")\n        \n        att_scores = torch.softmax(self.proj_v(H),axis=1) # (batch_size, seq_len)\n        # print(f\"att_scores shape:{att_scores.shape}\")\n        \n        attn_x = (x * att_scores).sum(1) # (batch_size, hidden_dim)\n        # print(f\"attn_x shape:{attn_x.shape}\")\n        return attn_x","metadata":{"papermill":{"duration":0.076802,"end_time":"2022-11-15T04:47:41.732733","exception":false,"start_time":"2022-11-15T04:47:41.655931","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:06:03.964177Z","iopub.execute_input":"2024-03-05T15:06:03.964949Z","iopub.status.idle":"2024-03-05T15:06:03.981252Z","shell.execute_reply.started":"2024-03-05T15:06:03.964886Z","shell.execute_reply":"2024-03-05T15:06:03.98024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNAClassifier(nn.Module):\n    def __init__(self, model_arch, hidden_dim=256, seq_len=24, pretrained=False):\n        super().__init__()\n        self.seq_len = seq_len\n        self.model = timm.create_model(model_arch, in_chans=1, pretrained=False)\n        self.model_arch = model_arch\n\n        if 'efficientnet' in self.model_arch:\n            cnn_feature = self.model.classifier.in_features\n            self.model.classifier = nn.Identity()\n        elif \"res\" in self.model_arch:\n            cnn_feature = self.model.fc.in_features\n            self.model.global_pool = nn.Identity()\n            self.model.fc = nn.Identity()\n            self.pooling = nn.AdaptiveAvgPool2d(1)\n        \n        self.spatialdropout = SpatialDropout(CFG.dropout)\n        self.gru = nn.GRU(cnn_feature, hidden_dim, 2, batch_first=True, bidirectional=True)\n        self.mlp_attention_layer = MLPAttentionNetwork(2 * hidden_dim)\n        self.logits = nn.Sequential(\n            nn.Linear(hidden_dim*2, 128),\n            nn.ReLU(),\n            nn.Dropout(CFG.dropout),\n            nn.Linear(128, 1)\n        )\n\n        # for n, m in self.named_modules():\n        #     if isinstance(m, nn.GRU):\n        #         print(f\"init {m}\")\n        #         for param in m.parameters():\n        #             if len(param.shape) >= 2:\n        #                 nn.init.orthogonal_(param.data)\n        #             else:\n        #                 nn.init.normal_(param.data)\n\n    def forward(self, x): # (B, seq_len, H, W)\n        bs = x.size(0) \n        x = x.reshape(bs*self.seq_len, 1, x.size(2), x.size(3)) # (B*seq_len, 1, H, W)\n        features = self.model(x)   \n        if \"res\" in self.model_arch:                             \n            features = self.pooling(features).view(bs*self.seq_len, -1) # (B*seq_len, cnn_feature)\n        features = self.spatialdropout(features)                # (B*seq_len, cnn_feature)\n        # print(features.shape)\n        features = features.reshape(bs, self.seq_len, -1)       # (B, seq_len, cnn_feature)\n        features, _ = self.gru(features)                        # (B, seq_len, hidden_dim*2)\n        atten_out = self.mlp_attention_layer(features)          # (B, hidden_dim*2)\n        pred = self.logits(atten_out)                           # (B, 1)\n        pred = pred.view(bs, -1)                                # (B, 1)\n        return pred","metadata":{"papermill":{"duration":0.073318,"end_time":"2022-11-15T04:47:41.859835","exception":false,"start_time":"2022-11-15T04:47:41.786517","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:06:03.98246Z","iopub.execute_input":"2024-03-05T15:06:03.982772Z","iopub.status.idle":"2024-03-05T15:06:03.996738Z","shell.execute_reply.started":"2024-03-05T15:06:03.982745Z","shell.execute_reply":"2024-03-05T15:06:03.99591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = TestDataset(train_df, transform=get_transforms(data='valid'))\ntest_loader = DataLoader(test_dataset, batch_size=CFG.valid_bs, shuffle=False, num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n\ncls_model_list = []\nfor m_arch, m_weight  in zip(CFG.archs_list, CFG.weights_list):\n    model = RSNAClassifier(m_arch, hidden_dim=256, seq_len=24, pretrained=False)\n    model.to(device)\n    model.load_state_dict(torch.load(m_weight)[\"model\"])\n    model.eval()\n    cls_model_list.append(model)","metadata":{"papermill":{"duration":6.628232,"end_time":"2022-11-15T04:47:48.542403","exception":false,"start_time":"2022-11-15T04:47:41.914171","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:06:03.998083Z","iopub.execute_input":"2024-03-05T15:06:03.99876Z","iopub.status.idle":"2024-03-05T15:06:11.354274Z","shell.execute_reply.started":"2024-03-05T15:06:03.998724Z","shell.execute_reply":"2024-03-05T15:06:11.353186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(cls_model_list)","metadata":{"papermill":{"duration":0.099737,"end_time":"2022-11-15T04:47:48.729054","exception":false,"start_time":"2022-11-15T04:47:48.629317","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:06:11.35648Z","iopub.execute_input":"2024-03-05T15:06:11.356876Z","iopub.status.idle":"2024-03-05T15:06:11.364632Z","shell.execute_reply.started":"2024-03-05T15:06:11.356838Z","shell.execute_reply":"2024-03-05T15:06:11.363601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_preds = []\nfor step, (images) in tqdm(enumerate(test_loader), total=len(test_loader)):\n    images = images.to(device, dtype=torch.float) # study-cid:24*img_sz*img_sz\n    models_preds = []\n    for model in cls_model_list:\n        with torch.no_grad(): \n            y_preds = model(images) # (B, 1)\n            y_preds = y_preds.squeeze(1)\n            models_preds.append(y_preds.sigmoid().to('cpu').numpy()) # list,len=model_nums,np(batch)\n    models_preds = np.mean(models_preds, axis=0) # batch, one sample preds\n    all_preds.append(models_preds)    \nall_preds = np.concatenate(all_preds)\n","metadata":{"papermill":{"duration":49.282731,"end_time":"2022-11-15T04:48:38.095203","exception":false,"start_time":"2022-11-15T04:47:48.812472","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:06:11.365966Z","iopub.execute_input":"2024-03-05T15:06:11.366544Z","iopub.status.idle":"2024-03-05T15:07:13.786694Z","shell.execute_reply.started":"2024-03-05T15:06:11.366517Z","shell.execute_reply":"2024-03-05T15:07:13.785744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df[\"fractured\"] = all_preds\nmodel_preds_df = train_df[[\"StudyInstanceUID\", \"cid\", \"fractured\"]]\nprint(model_preds_df.shape)\nmodel_preds_df.head()","metadata":{"papermill":{"duration":0.129204,"end_time":"2022-11-15T04:48:38.325553","exception":false,"start_time":"2022-11-15T04:48:38.196349","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:07:13.788176Z","iopub.execute_input":"2024-03-05T15:07:13.788545Z","iopub.status.idle":"2024-03-05T15:07:13.804318Z","shell.execute_reply.started":"2024-03-05T15:07:13.788513Z","shell.execute_reply":"2024-03-05T15:07:13.803292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def type_to_num(x):\n    type_dict = {\n        \"L1\":1,\n        \"L2\":2,\n        \"L3\":3,\n        \"L4\":4,\n        \"L5\":5,\n        \"patient_overall\":8,\n    }\n    return type_dict[x]\n\ntest_df[\"cid\"] = test_df[\"prediction_type\"].apply(type_to_num)\nprint(test_df.shape)\ntest_df.head(8)","metadata":{"papermill":{"duration":0.125295,"end_time":"2022-11-15T04:48:38.550065","exception":false,"start_time":"2022-11-15T04:48:38.42477","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:07:13.805374Z","iopub.execute_input":"2024-03-05T15:07:13.805669Z","iopub.status.idle":"2024-03-05T15:07:13.82201Z","shell.execute_reply.started":"2024-03-05T15:07:13.805643Z","shell.execute_reply":"2024-03-05T15:07:13.821099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = test_df.merge(model_preds_df, how=\"left\", on=[\"StudyInstanceUID\", \"cid\"])\ntest_df[\"fractured\"] = test_df[\"fractured\"].fillna(CFG.fillna_number)\nprint(test_df.shape)\ntest_df.head(8)","metadata":{"papermill":{"duration":0.127265,"end_time":"2022-11-15T04:48:38.777039","exception":false,"start_time":"2022-11-15T04:48:38.649774","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:07:13.823157Z","iopub.execute_input":"2024-03-05T15:07:13.823496Z","iopub.status.idle":"2024-03-05T15:07:13.843826Z","shell.execute_reply.started":"2024-03-05T15:07:13.823456Z","shell.execute_reply":"2024-03-05T15:07:13.842946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for study_id in test_study_list:\n    # overall_fractured = test_df[test_df[\"StudyInstanceUID\"]==study_id][\"fractured\"].max()\n    overall_fractured = test_df[test_df[\"StudyInstanceUID\"]==study_id][\"fractured\"][:7].agg(lambda x:1-((1-x).prod()))\n    test_df.loc[((test_df[\"StudyInstanceUID\"]==study_id) & (test_df[\"prediction_type\"]==\"patient_overall\")), \"fractured\"] = overall_fractured\nprint(test_df.shape)\ntest_df","metadata":{"papermill":{"duration":0.08577,"end_time":"2022-11-15T04:48:38.950451","exception":false,"start_time":"2022-11-15T04:48:38.864681","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:07:13.845193Z","iopub.execute_input":"2024-03-05T15:07:13.845557Z","iopub.status.idle":"2024-03-05T15:07:13.868343Z","shell.execute_reply.started":"2024-03-05T15:07:13.845523Z","shell.execute_reply":"2024-03-05T15:07:13.867378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_df = test_df[[\"row_id\", \"fractured\"]]\nfinal_df.to_csv(\"submission.csv\", index=False)\nfinal_df","metadata":{"papermill":{"duration":0.074513,"end_time":"2022-11-15T04:48:39.206164","exception":false,"start_time":"2022-11-15T04:48:39.131651","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-05T15:07:13.869714Z","iopub.execute_input":"2024-03-05T15:07:13.870095Z","iopub.status.idle":"2024-03-05T15:07:13.885638Z","shell.execute_reply.started":"2024-03-05T15:07:13.870068Z","shell.execute_reply":"2024-03-05T15:07:13.884691Z"},"trusted":true},"execution_count":null,"outputs":[]}]}