{"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":"DEBUG = False\nINPUT = '/kaggle/input/rsna-2023-abdominal-trauma-detection'\nSTAGE_2_MODELS = '/kaggle/input/rsna23-train-stage2-final-5'","metadata":{"execution":{"iopub.status.busy":"2023-10-15T18:46:16.950297Z","iopub.execute_input":"2023-10-15T18:46:16.950812Z","iopub.status.idle":"2023-10-15T18:46:16.977856Z","shell.execute_reply.started":"2023-10-15T18:46:16.950722Z","shell.execute_reply":"2023-10-15T18:46:16.976876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nsys.path = [\n    '../input/covn3d-same',\n    '../input/timm20221011/pytorch-image-models-master',\n    '../input/smp20210127/segmentation_models.pytorch-master/segmentation_models.pytorch-master',\n    '../input/smp20210127/pretrained-models.pytorch-master/pretrained-models.pytorch-master',\n    '../input/smp20210127/EfficientNet-PyTorch-master/EfficientNet-PyTorch-master',\n] + sys.path\n\n!pip -q install ../input/pylibjpeg140py3/pylibjpeg-1.4.0-py3-none-any.whl\n!pip -q install ../input/pylibjpeg140py3/python_gdcm-3.0.17.1-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n\n!cp -r ../input/timm-20220211/pytorch-image-models-master/timm ./timm4smp","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-10-15T18:46:17.366779Z","iopub.execute_input":"2023-10-15T18:46:17.36762Z","iopub.status.idle":"2023-10-15T18:47:29.227426Z","shell.execute_reply.started":"2023-10-15T18:46:17.367585Z","shell.execute_reply":"2023-10-15T18:47:29.226145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nfrom collections import Counter\nimport psutil\nimport shutil\nimport ctypes\nfrom fastcore.all import Path\n\nfrom collections import defaultdict\nimport ast\nimport cv2\nimport time\nimport timm\nimport timm4smp\nimport pickle\nimport random\nimport pydicom\nimport argparse\nimport warnings\nimport threading\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom glob import glob\nimport albumentations\nimport matplotlib.pyplot as plt\nimport segmentation_models_pytorch as smp\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.cuda.amp as amp\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom pylab import rcParams\n\n%matplotlib inline\ndevice = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\ntorch.backends.cudnn.benchmark = True\n\ntimm.__version__, timm4smp.__version__","metadata":{"execution":{"iopub.status.busy":"2023-10-15T18:47:29.230014Z","iopub.execute_input":"2023-10-15T18:47:29.230718Z","iopub.status.idle":"2023-10-15T18:47:36.948077Z","shell.execute_reply.started":"2023-10-15T18:47:29.230677Z","shell.execute_reply":"2023-10-15T18:47:36.947006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_ram_usage(threshold: float=None): \n    \"Returns ram usage in GB, garbage collecting if its over `threshold`.\"\n    process = psutil.Process(os.getpid())\n    memory_usage_bytes = process.memory_info().rss\n    ram_usage = memory_usage_bytes / (1000 ** 3)\n    if threshold and ram_usage > threshold:\n        print(f'ram usage: {ram_usage}, garbage collecting')\n        libc = ctypes.CDLL(\"libc.so.6\")\n        libc.malloc_trim(0)\n        gc.collect()\n        print(f'new ram usage: {process.memory_info().rss / (1000 ** 3)}')\n    return ram_usage","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:00:54.359018Z","iopub.execute_input":"2023-10-14T23:00:54.359744Z","iopub.status.idle":"2023-10-14T23:00:54.367693Z","shell.execute_reply.started":"2023-10-14T23:00:54.3597Z","shell.execute_reply":"2023-10-14T23:00:54.366062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gru = get_ram_usage","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:00:54.370666Z","iopub.execute_input":"2023-10-14T23:00:54.371112Z","iopub.status.idle":"2023-10-14T23:00:54.382159Z","shell.execute_reply.started":"2023-10-14T23:00:54.371067Z","shell.execute_reply":"2023-10-14T23:00:54.381312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Investigating studies per patient","metadata":{}},{"cell_type":"code","source":"if DEBUG: \n    met = pd.read_parquet('/kaggle/input/rsna-2023-abdominal-trauma-detection/train_dicom_tags.parquet')\n    met = met.rename(columns={'PatientID': 'patient_id'})\n\n    x =  met.SeriesInstanceUID.nunique()/ met.patient_id.nunique()\n    print(x, 'scans per patient')\n    print(f'therefore, there should be around {1100 * x} scans in the test set')\n","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:00:54.38326Z","iopub.execute_input":"2023-10-14T23:00:54.384238Z","iopub.status.idle":"2023-10-14T23:00:54.395961Z","shell.execute_reply.started":"2023-10-14T23:00:54.384205Z","shell.execute_reply":"2023-10-14T23:00:54.394912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_dir_seg = '/kaggle/input/rsna23-train-stage1/models'\nimage_size_seg = (128, 128, 128)\nimage_size = 224\nmsk_size = image_size_seg[0]\nimage_size_cls = 224\nn_slice_per_c = 15\nn_ch = 5\n\nbatch_size_seg = 1\nnum_workers = 2","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:00:54.397111Z","iopub.execute_input":"2023-10-14T23:00:54.398015Z","iopub.status.idle":"2023-10-14T23:00:54.408967Z","shell.execute_reply.started":"2023-10-14T23:00:54.397983Z","shell.execute_reply":"2023-10-14T23:00:54.408191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Make dataframe of studies with dicoms","metadata":{}},{"cell_type":"code","source":"def make_df(kind='test'): \n    TEST_IMAGES = Path(f'/kaggle/input/rsna-2023-abdominal-trauma-detection/{kind}_images')\n    patient_ids = []\n    studys = []\n    image_folders = []\n    for patient_id in TEST_IMAGES.ls(): \n        for study in patient_id.ls(): \n            if len(study.ls()) > 0: \n                patient_ids.append(int(patient_id.stem))\n                studys.append(study.stem)\n                image_folders.append(str(study))\n    return pd.DataFrame(dict(zip(('patient_id', 'study', 'image_folder'), \n                               (patient_ids, studys, image_folders))))","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:00:54.410009Z","iopub.execute_input":"2023-10-14T23:00:54.410901Z","iopub.status.idle":"2023-10-14T23:00:54.425347Z","shell.execute_reply.started":"2023-10-14T23:00:54.410859Z","shell.execute_reply":"2023-10-14T23:00:54.424175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = make_df('train') if DEBUG else make_df('test')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:20:35.275925Z","iopub.execute_input":"2023-10-14T23:20:35.276342Z","iopub.status.idle":"2023-10-14T23:20:35.299047Z","shell.execute_reply.started":"2023-10-14T23:20:35.27631Z","shell.execute_reply":"2023-10-14T23:20:35.297538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"def standardize_pixel_array(dcm: pydicom.dataset.FileDataset) -> np.ndarray:\n    \"\"\"\n    Source : https://www.kaggle.com/competitions/rsna-2023-abdominal-trauma-detection/discussion/427217\n    \"\"\"\n    # Correct DICOM pixel_array if PixelRepresentation == 1.\n    pixel_array = dcm.pixel_array\n    if dcm.PixelRepresentation == 1:\n        bit_shift = dcm.BitsAllocated - dcm.BitsStored\n        dtype = pixel_array.dtype \n        pixel_array = (pixel_array << bit_shift).astype(dtype) >>  bit_shift\n#         pixel_array = pydicom.pixel_data_handlers.util.apply_modality_lut(new_array, dcm)\n\n    intercept = float(dcm.RescaleIntercept)\n    slope = float(dcm.RescaleSlope)\n    center = int(dcm.WindowCenter)\n    width = int(dcm.WindowWidth)\n    low = center - width / 2\n    high = center + width / 2    \n    \n    pixel_array = (pixel_array * slope) + intercept\n    pixel_array = np.clip(pixel_array, low, high)\n\n    return pixel_array\n\n\ndef study_path_to_3D_image(path, image_size_seg=(128, 128, 128), plot_image=False, z_df=None):\n    t_paths = sorted(glob(os.path.join(path, \"*\")), key=lambda x: int(x.split('/')[-1].split(\".\")[0]))\n    t_paths = [p for p in t_paths if '3124/5842/514.dcm' not in p]\n    n_scans = len(t_paths)\n#     print(n_scans)\n    indices = np.quantile(list(range(n_scans)), np.linspace(0., 1., image_size_seg[2])).round().astype(int)\n    print(len(indices))\n    t_paths = [t_paths[i] for i in indices]\n\n    imgs = {}\n    pos_zs = []\n    for filename in t_paths:\n        dicom = pydicom.dcmread(filename)\n        pos_z = dicom[(0x20, 0x32)].value[-1]  # to retrieve the order of frames\n#         dicom_numbers.append(filename.split('/')[-1].split(\".\")[0])\n        pos_zs.append(pos_z)\n        img = standardize_pixel_array(dicom)\n        img = cv2.resize(img, (image_size_seg[0], image_size_seg[1]), interpolation = cv2.INTER_AREA)\n        if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n            img = 1 - img\n        imgs[pos_z] = img\n        \n#     print(dicom_numbers, pos_zs)\n    if len(pos_zs) > 1: \n        pos_z_ascending = pos_zs[-1] > pos_zs[0]\n    else: \n        pos_z_ascending = True\n    if z_df is not None: \n        study = path.split('/')[-1]\n        z_df.loc[study, 'pos_z_ascending'] = pos_zs[-1] > pos_zs[0]\n        \n    images = []\n    cnt = Counter(pos_zs)\n    for i, k in enumerate(sorted(imgs.keys())):\n        img = imgs[k]\n#         images.append(img)\n        images.extend([img] * cnt[k]) # to make sure we have same dimensions\n        if not (i % 100) and plot_image:\n            plt.figure(figsize=(5, 5))\n            plt.imshow(img, cmap=\"gray\")\n            plt.title(f\"Patient {patient} - Study {study} - Frame {i}/{len(imgs)}\")\n            plt.axis(False)\n            plt.show()\n    images = np.stack(images, -1)\n    \n    images = images - np.min(images)\n    images = images / (np.max(images) + 1e-4)\n    images = (images * 255).astype(np.uint8)\n    return images","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:01:40.449393Z","iopub.execute_input":"2023-10-14T23:01:40.449813Z","iopub.status.idle":"2023-10-14T23:01:40.468689Z","shell.execute_reply.started":"2023-10-14T23:01:40.449776Z","shell.execute_reply":"2023-10-14T23:01:40.467363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ORGANS = ['liver', 'spleen', 'kidney', 'bowel']\nLABELS = [['liver_healthy', 'liver_low', 'liver_high'], \n          ['spleen_healthy', 'spleen_low', 'spleen_high'], \n          ['kidney_healthy', 'kidney_low', 'kidney_high'], \n          ['bowel_healthy', 'bowel_injury']\n         ]\nABS = [[0, 15], [15, 30], [30, 60], [60, 75]]\nN_ORGANS = len(ORGANS)","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:01:41.132282Z","iopub.execute_input":"2023-10-14T23:01:41.132638Z","iopub.status.idle":"2023-10-14T23:01:41.137991Z","shell.execute_reply.started":"2023-10-14T23:01:41.132608Z","shell.execute_reply":"2023-10-14T23:01:41.136793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_sample(row, has_mask=True):\n    image = study_path_to_3D_image(row.image_folder)\n    if image.ndim < 4:\n        image = np.expand_dims(image, 0) # to 3ch\n    return image.repeat(3, 0) \n\n\nclass SegTestDataset(Dataset):\n\n    def __init__(self, df):\n        self.df = df.reset_index()\n\n    def __len__(self):\n        return self.df.shape[0]\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n\n        image = load_sample(row, has_mask=False)\n        image = image / 255.\n#         gc.collect()\n        return torch.tensor(image).float()\n","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:01:44.256615Z","iopub.execute_input":"2023-10-14T23:01:44.257379Z","iopub.status.idle":"2023-10-14T23:01:44.264382Z","shell.execute_reply.started":"2023-10-14T23:01:44.257339Z","shell.execute_reply":"2023-10-14T23:01:44.263434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_seg = SegTestDataset(df)\ndisplay(df.head())\nloader_seg = torch.utils.data.DataLoader(dataset_seg, batch_size=batch_size_seg, shuffle=False, num_workers=num_workers)","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:02:40.989538Z","iopub.execute_input":"2023-10-14T23:02:40.990005Z","iopub.status.idle":"2023-10-14T23:02:41.010249Z","shell.execute_reply.started":"2023-10-14T23:02:40.989965Z","shell.execute_reply":"2023-10-14T23:02:41.009194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    rcParams['figure.figsize'] = 20,8\n    for i in range(1):\n        f, axarr = plt.subplots(1,3)\n        for p in range(3):\n            idx = i*4+p\n            img = dataset_seg[idx]\n            img = img[:, :, :, 60]\n            axarr[p].imshow(img.transpose(0, 1).transpose(1,2).squeeze())","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:03:14.316893Z","iopub.execute_input":"2023-10-14T23:03:14.31728Z","iopub.status.idle":"2023-10-14T23:03:14.325255Z","shell.execute_reply.started":"2023-10-14T23:03:14.317251Z","shell.execute_reply":"2023-10-14T23:03:14.323765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    rcParams['figure.figsize'] = 20,8\n    for i in range(1):\n        f, axarr = plt.subplots(1,3)\n        for p in range(3):\n            idx = i*4+p\n            img = dataset_seg[idx]\n            img = img[:, :, 60, :]\n            axarr[p].imshow(img.transpose(0, 1).transpose(1,2).squeeze())\n    get_ram_usage(1)","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:03:16.799958Z","iopub.execute_input":"2023-10-14T23:03:16.800337Z","iopub.status.idle":"2023-10-14T23:03:16.807811Z","shell.execute_reply.started":"2023-10-14T23:03:16.800307Z","shell.execute_reply":"2023-10-14T23:03:16.806538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"drop_rate = 0.\ndrop_path_rate = 0.\nout_dim = 5\nclass TimmSegModel(nn.Module):\n    def __init__(self, backbone, segtype='unet', pretrained=False):\n        super(TimmSegModel, self).__init__()\n\n        self.encoder = timm.create_model(\n            backbone,\n            in_chans=3,\n            features_only=True,\n            drop_rate=drop_rate,\n            drop_path_rate=drop_path_rate,\n            pretrained=pretrained\n        )\n        g = self.encoder(torch.rand(1, 3, 64, 64))\n        encoder_channels = [1] + [_.shape[1] for _ in g]\n        decoder_channels = [256, 128, 64, 32, 16]\n        if segtype == 'unet':\n            self.decoder = smp.unet.decoder.UnetDecoder(\n                encoder_channels=encoder_channels[:n_blocks+1],\n                decoder_channels=decoder_channels[:n_blocks],\n                n_blocks=n_blocks,\n            )\n\n        self.segmentation_head = nn.Conv2d(decoder_channels[n_blocks-1], out_dim, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))\n\n    def forward(self,x):\n        global_features = [0] + self.encoder(x)[:n_blocks]\n        seg_features = self.decoder(*global_features)\n        seg_features = self.segmentation_head(seg_features)\n        return seg_features\n\nfrom timm.models.layers.conv2d_same import Conv2dSame\nfrom conv3d_same import Conv3dSame\n\n\ndef convert_3d(module):\n\n    module_output = module\n    if isinstance(module, torch.nn.BatchNorm2d):\n        module_output = torch.nn.BatchNorm3d(\n            module.num_features,\n            module.eps,\n            module.momentum,\n            module.affine,\n            module.track_running_stats,\n        )\n        if module.affine:\n            with torch.no_grad():\n                module_output.weight = module.weight\n                module_output.bias = module.bias\n        module_output.running_mean = module.running_mean\n        module_output.running_var = module.running_var\n        module_output.num_batches_tracked = module.num_batches_tracked\n        if hasattr(module, \"qconfig\"):\n            module_output.qconfig = module.qconfig\n            \n    elif isinstance(module, Conv2dSame):\n        module_output = Conv3dSame(\n            in_channels=module.in_channels,\n            out_channels=module.out_channels,\n            kernel_size=module.kernel_size[0],\n            stride=module.stride[0],\n            padding=module.padding[0],\n            dilation=module.dilation[0],\n            groups=module.groups,\n            bias=module.bias is not None,\n        )\n        module_output.weight = torch.nn.Parameter(module.weight.unsqueeze(-1).repeat(1,1,1,1,module.kernel_size[0]))\n\n    elif isinstance(module, torch.nn.Conv2d):\n        module_output = torch.nn.Conv3d(\n            in_channels=module.in_channels,\n            out_channels=module.out_channels,\n            kernel_size=module.kernel_size[0],\n            stride=module.stride[0],\n            padding=module.padding[0],\n            dilation=module.dilation[0],\n            groups=module.groups,\n            bias=module.bias is not None,\n            padding_mode=module.padding_mode\n        )\n        module_output.weight = torch.nn.Parameter(module.weight.unsqueeze(-1).repeat(1,1,1,1,module.kernel_size[0]))\n\n    elif isinstance(module, torch.nn.MaxPool2d):\n        module_output = torch.nn.MaxPool3d(\n            kernel_size=module.kernel_size,\n            stride=module.stride,\n            padding=module.padding,\n            dilation=module.dilation,\n            ceil_mode=module.ceil_mode,\n        )\n    elif isinstance(module, torch.nn.AvgPool2d):\n        module_output = torch.nn.AvgPool3d(\n            kernel_size=module.kernel_size,\n            stride=module.stride,\n            padding=module.padding,\n            ceil_mode=module.ceil_mode,\n        )\n\n    for name, child in module.named_children():\n        module_output.add_module(\n            name, convert_3d(child)\n        )\n    del module\n\n    return module_output\n\n\n# backbone = 'resnet18d'\n# n_blocks = 4\n# model = TimmSegModel(backbone)\n# model = convert_3d(model)\n# model(torch.rand(1, 3, 128,128,128)).shape\n    ","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:03:18.840271Z","iopub.execute_input":"2023-10-14T23:03:18.841278Z","iopub.status.idle":"2023-10-14T23:03:18.876061Z","shell.execute_reply.started":"2023-10-14T23:03:18.841236Z","shell.execute_reply":"2023-10-14T23:03:18.874928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"    \nclass TimmModel(nn.Module):\n    def __init__(self, backbone, pretrained=False, out_dim=3, h=image_size, w=image_size):\n        super(TimmModel, self).__init__()\n        self.h = h\n        self.w = w\n\n        self.encoder = timm.create_model(\n            backbone,\n            in_chans=in_chans,\n            num_classes=out_dim,\n            features_only=False,\n            drop_rate=drop_rate,\n            drop_path_rate=drop_path_rate,\n            pretrained=pretrained\n        )\n\n        if 'efficient' in backbone:\n            hdim = self.encoder.conv_head.out_channels\n            self.encoder.classifier = nn.Identity()\n        elif 'convnext' in backbone:\n            hdim = self.encoder.head.fc.in_features\n            self.encoder.head.fc = nn.Identity()\n\n\n        self.lstm = nn.LSTM(hdim, 256, num_layers=2, dropout=drop_rate, bidirectional=True, batch_first=True)\n        self.head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.Dropout(drop_rate_last),\n            nn.LeakyReLU(0.1),\n            nn.Linear(256, out_dim), # chacnged\n        )\n\n    def forward(self, x):  # (bs, nslice, ch, sz, sz)\n        bs = x.shape[0]\n        x = x.view(bs * n_slice_per_c, in_chans, self.h, self.w)\n        feat = self.encoder(x)\n        feat = feat.view(bs, n_slice_per_c, -1)\n        feat, _ = self.lstm(feat)\n        feat = feat.contiguous().view(bs * n_slice_per_c, -1)\n        feat = self.head(feat)\n        feat = feat.view(bs, n_slice_per_c, -1).contiguous()\n\n        return feat","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:03:19.520274Z","iopub.execute_input":"2023-10-14T23:03:19.520862Z","iopub.status.idle":"2023-10-14T23:03:19.53146Z","shell.execute_reply.started":"2023-10-14T23:03:19.520813Z","shell.execute_reply":"2023-10-14T23:03:19.530183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Models","metadata":{}},{"cell_type":"code","source":"models_seg = []\n\nkernel_type = 'timm3d_res18d_unet4b_128_128_128_dsv2_flip12_shift333p7_gd1p5_bs4_lr3e4_20x50ep'\nbackbone = 'resnet18d'\nn_blocks = 4\nfor fold in range(1):\n    model = TimmSegModel(backbone, pretrained=False)\n    model = convert_3d(model)\n    load_model_file = '/kaggle/input/rsna23-train-stage1/models/timm3d_res18d_unet4b_128_128_128_dsv2_flip12_shift333p7_gd1p5_bs4_lr3e4_20x50ep_fold0_best.pth'\n    sd = torch.load(load_model_file, map_location=torch.device('cpu'))\n    if 'model_state_dict' in sd.keys():\n        sd = sd['model_state_dict']\n    sd = {k[7:] if k.startswith('module.') else k: sd[k] for k in sd.keys()}\n    model.load_state_dict(sd, strict=True)\n    model = model.to(device)\n    model.eval()\n    models_seg.append(model)\nlen(models_seg)","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:03:25.264678Z","iopub.execute_input":"2023-10-14T23:03:25.265085Z","iopub.status.idle":"2023-10-14T23:03:27.925923Z","shell.execute_reply.started":"2023-10-14T23:03:25.265054Z","shell.execute_reply":"2023-10-14T23:03:27.924939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kernel_type = '0920_1bonev2_effv2s_224_15_6ch_augv2_mixupp5_drl3_rov1p2_bs8_lr23e5_eta23e6_50ep'\nbackbone = 'tf_efficientnetv2_s_in21ft1k'\nmodel_dir_cls = F'{STAGE_2_MODELS}/models'\nin_chans = 6\ndrop_rate_last = 0.3\nmodels = [TimmModel(backbone, pretrained=False), \n             TimmModel(backbone, pretrained=False), \n             TimmModel(backbone, pretrained=False, h=image_size*2), \n             TimmModel(backbone, pretrained=False, out_dim=2), ]\n\nfor model, organ in zip(models, ORGANS):\n    print(organ)\n    load_model_file = f'{STAGE_2_MODELS}/models/{organ}_0920_1bonev2_effv2s_224_15_6ch_augv2_mixupp5_drl3_rov1p2_bs8_lr23e5_eta23e6_50ep_fold0_last.pth'\n    sd = torch.load(load_model_file, map_location=torch.device('cpu'))\n    if 'model_state_dict' in sd.keys():\n        sd = sd['model_state_dict']\n    sd = {k[7:] if k.startswith('module.') else k: sd[k] for k in sd.keys()}\n    model.load_state_dict(sd, strict=True)\n    model = model.to(device)\n    model.eval()\nlen(models)","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:03:27.928483Z","iopub.execute_input":"2023-10-14T23:03:27.928939Z","iopub.status.idle":"2023-10-14T23:03:35.162812Z","shell.execute_reply.started":"2023-10-14T23:03:27.928894Z","shell.execute_reply":"2023-10-14T23:03:35.16192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:03:35.164232Z","iopub.execute_input":"2023-10-14T23:03:35.164672Z","iopub.status.idle":"2023-10-14T23:03:36.350896Z","shell.execute_reply.started":"2023-10-14T23:03:35.16462Z","shell.execute_reply":"2023-10-14T23:03:36.349557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_bone(msk, cid, t_paths, cropped_images):\n    n_scans = len(t_paths)\n    bone = []\n    try:\n        msk_b = msk[cid] > 0.2\n        msk_c = msk[cid] > 0.05\n\n        x = np.where(msk_b.sum(1).sum(1) > 0)[0]\n        y = np.where(msk_b.sum(0).sum(1) > 0)[0]\n        z = np.where(msk_b.sum(0).sum(0) > 0)[0]\n\n        if len(x) == 0 or len(y) == 0 or len(z) == 0:\n            x = np.where(msk_c.sum(1).sum(1) > 0)[0]\n            y = np.where(msk_c.sum(0).sum(1) > 0)[0]\n            z = np.where(msk_c.sum(0).sum(0) > 0)[0]\n\n        x1, x2 = max(0, x[0] - 1), min(msk.shape[1], x[-1] + 1)\n        y1, y2 = max(0, y[0] - 1), min(msk.shape[2], y[-1] + 1)\n        z1, z2 = max(0, z[0] - 1), min(msk.shape[3], z[-1] + 1)\n        zz1, zz2 = int(z1 / msk_size * n_scans), int(z2 / msk_size * n_scans)\n\n        inds = np.linspace(zz1 ,zz2-1 ,n_slice_per_c).astype(int)\n        inds_ = np.linspace(z1 ,z2-1 ,n_slice_per_c).astype(int)\n        for sid, (ind, ind_) in enumerate(zip(inds, inds_)):\n\n            msk_this = msk[cid, :, :, ind_]\n\n            images = []\n            for i in range(-n_ch//2+1, n_ch//2+1):\n                try:\n                    dicom = pydicom.read_file(t_paths[ind+i])\n                    img = standardize_pixel_array(dicom)\n                    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n                        img = 1 - img\n                    images.append(img)\n                except:\n                    images.append(np.zeros((512, 512)))\n\n            data = np.stack(images, -1)\n            data = data - np.min(data)\n            data = data / (np.max(data) + 1e-4)\n            data = (data * 255).astype(np.uint8)\n            msk_this = msk_this[x1:x2, y1:y2]\n            xx1 = int(x1 / msk_size * data.shape[0])\n            xx2 = int(x2 / msk_size * data.shape[0])\n            yy1 = int(y1 / msk_size * data.shape[1])\n            yy2 = int(y2 / msk_size * data.shape[1])\n            data = data[xx1:xx2, yy1:yy2]\n            data = np.stack([cv2.resize(data[:, :, i], (image_size_cls, image_size_cls), interpolation = cv2.INTER_LINEAR) for i in range(n_ch)], -1)\n            msk_this = (msk_this * 255).astype(np.uint8)\n            msk_this = cv2.resize(msk_this, (image_size_cls, image_size_cls), interpolation = cv2.INTER_LINEAR)\n\n            data = np.concatenate([data, msk_this[:, :, np.newaxis]], -1)\n\n            bone.append(torch.tensor(data))\n#             gc.collect()\n\n    except:\n        for sid in range(n_slice_per_c):\n            bone.append(torch.ones((image_size_cls, image_size_cls, n_ch+1)).int())\n\n    cropped_images[cid] = torch.stack(bone, 0)\n\n\ndef load_cropped_images(msk, image_folder, n_ch=n_ch):\n\n    pos_z = []\n    t_paths = sorted(glob(os.path.join(image_folder, \"*\")), key=lambda x: int(x.split('/')[-1].split(\".\")[0]))\n    for filename in t_paths[:2]:\n        dicom = pydicom.dcmread(filename)\n        pos_z.append(dicom[(0x20, 0x32)].value[-1])  # to retrieve the order of frames\n    if len(pos_z) > 1: z_ascending = pos_z[1] > pos_z[0] \n    else: z_ascending = True\n    if not z_ascending: t_paths.reverse()\n    for cid in range(5):\n        threads[cid] = threading.Thread(target=load_bone, args=(msk, cid, t_paths, cropped_images))\n        threads[cid].start()\n    for cid in range(5):\n        threads[cid].join()\n#     gc.collect()\n\n    return torch.cat(cropped_images, 0)\n","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:03:36.353513Z","iopub.execute_input":"2023-10-14T23:03:36.353894Z","iopub.status.idle":"2023-10-14T23:03:36.378318Z","shell.execute_reply.started":"2023-10-14T23:03:36.353859Z","shell.execute_reply":"2023-10-14T23:03:36.377481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"# For exceptions: \ntrain = pd.read_csv('/kaggle/input/rsna-2023-abdominal-trauma-detection/train.csv')\ntar_means = train.mean()\nn_exceptions = 0","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:17:47.241218Z","iopub.execute_input":"2023-10-14T23:17:47.241697Z","iopub.status.idle":"2023-10-14T23:17:47.258713Z","shell.execute_reply.started":"2023-10-14T23:17:47.241659Z","shell.execute_reply":"2023-10-14T23:17:47.257641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_seg = SegTestDataset(df)\ndisplay(df.head())\nloader_seg = torch.utils.data.DataLoader(dataset_seg, batch_size=batch_size_seg, \n                                         shuffle=False, num_workers=num_workers)","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:18:02.164467Z","iopub.execute_input":"2023-10-14T23:18:02.164932Z","iopub.status.idle":"2023-10-14T23:18:02.179458Z","shell.execute_reply.started":"2023-10-14T23:18:02.164892Z","shell.execute_reply":"2023-10-14T23:18:02.178288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\npreds = defaultdict(list)\nbar = tqdm(loader_seg)\nwith torch.no_grad():\n    for batch_id, (images) in enumerate(bar):\n        try:\n            pred_masks = []\n            images = images.to(device)\n            pred_masks = []\n            for model in models_seg:\n                pmask = model(images).sigmoid()\n                pred_masks.append(pmask)\n            pred_masks = torch.stack(pred_masks, 0).mean(0).cpu().numpy()\n\n            # Build cls input\n            cls_inp = []\n            threads = [None] * 5\n            cropped_images = [None] * 5\n            for i in range(pred_masks.shape[0]):\n                row = df.iloc[batch_id*batch_size_seg+i]\n                cropped_images = load_cropped_images(pred_masks[i], row.image_folder)\n                cls_inp.append(cropped_images.permute(0, 3, 1, 2).float() / 255.)\n            cls_inp = torch.stack(cls_inp, 0).to(device)  # (1, 75, 6, 224, 224)\n\n            image_full = cls_inp[0]\n            out = defaultdict(dict)\n            for organ, cols, (a, b) in zip(ORGANS, LABELS, ABS):\n                images = []\n                for image in image_full[a: b]: \n                    images.append(image.cpu())\n                images = np.stack(images, 0)\n                if organ == 'kidney': \n                    images = np.concatenate((images[:15, :, :, :], images[15:, :, :, :]), 2)\n                out[organ]['images'] = torch.tensor(images).float().to(device)\n\n            for organ, model, label_cols in zip(ORGANS, models, LABELS): \n                try: \n                    logits = model(out[organ]['images'].unsqueeze(0)).squeeze()\n                    data = logits.sigmoid()\n                    min_first_column = torch.min(data[:, 0])\n                    max_last_two_columns = torch.max(data[:, 1:], dim=0).values\n                    result = torch.cat((min_first_column.unsqueeze(0), max_last_two_columns), dim=0)\n#                     print('***********', result.shape) ###################\n                    preds[organ].append(result.cpu())\n                except: \n                    print('XXXXXXXX')\n                    preds[organ].append(torch.tensor(tar_means[label_cols]))\n        except: \n            print('problem in loop')\n            n_exceptions += 1\n            for organ, label_cols in zip(ORGANS, LABELS):\n                preds[organ].append(torch.tensor(tar_means[label_cols]))\n        get_ram_usage(1)\n        get_ram_usage(1)\n        if batch_id % 100 == 0:\n            get_ram_usage(1)\n            time.sleep(1)","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:18:04.386668Z","iopub.execute_input":"2023-10-14T23:18:04.387778Z","iopub.status.idle":"2023-10-14T23:19:10.863178Z","shell.execute_reply.started":"2023-10-14T23:18:04.387737Z","shell.execute_reply":"2023-10-14T23:19:10.861912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for organ, label_cols in zip(ORGANS, LABELS):\n    df.loc[:, label_cols] = torch.stack(preds[organ]).numpy()\ndf = df.assign(extravasation_injury = .06355 * 6, \n               extravasation_healthy= 1 - (.06355 * 6))","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:24:39.180439Z","iopub.execute_input":"2023-10-14T23:24:39.180887Z","iopub.status.idle":"2023-10-14T23:24:39.196066Z","shell.execute_reply.started":"2023-10-14T23:24:39.180835Z","shell.execute_reply":"2023-10-14T23:24:39.194862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.read_csv('/kaggle/input/rsna-2023-abdominal-trauma-detection/sample_submission.csv')\ndf = df[sub.columns].groupby('patient_id').mean()\npred_means = df.mean()\ndisplay(df.head(), pred_means)","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:24:44.143016Z","iopub.execute_input":"2023-10-14T23:24:44.144116Z","iopub.status.idle":"2023-10-14T23:24:44.171603Z","shell.execute_reply.started":"2023-10-14T23:24:44.144072Z","shell.execute_reply":"2023-10-14T23:24:44.170232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = sub[['patient_id']]\nsub = sub.merge(df, left_on='patient_id', right_index=True, how='left').fillna(pred_means)\nsub.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:24:54.27839Z","iopub.execute_input":"2023-10-14T23:24:54.278814Z","iopub.status.idle":"2023-10-14T23:24:54.292893Z","shell.execute_reply.started":"2023-10-14T23:24:54.278777Z","shell.execute_reply":"2023-10-14T23:24:54.291683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub","metadata":{"execution":{"iopub.status.busy":"2023-10-14T23:24:55.030449Z","iopub.execute_input":"2023-10-14T23:24:55.030863Z","iopub.status.idle":"2023-10-14T23:24:55.050752Z","shell.execute_reply.started":"2023-10-14T23:24:55.030809Z","shell.execute_reply":"2023-10-14T23:24:55.047307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}