{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"papermill":{"default_parameters":{},"duration":70.264901,"end_time":"2023-09-02T21:24:06.018226","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2023-09-02T21:22:55.753325","version":"2.4.0"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":6268436,"sourceType":"datasetVersion","datasetId":3603236},{"sourceId":9182126,"sourceType":"datasetVersion","datasetId":5296330},{"sourceId":9495419,"sourceType":"datasetVersion","datasetId":5558243},{"sourceId":9497658,"sourceType":"datasetVersion","datasetId":5558262},{"sourceId":9518097,"sourceType":"datasetVersion","datasetId":5587861},{"sourceId":9521808,"sourceType":"datasetVersion","datasetId":5588533},{"sourceId":9524618,"sourceType":"datasetVersion","datasetId":5588550}],"dockerImageVersionId":30733,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport os\nimport pydicom\nimport cv2\n# import einops","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/rsna-2023-whl/einops-0.6.1-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:42:06.555904Z","iopub.execute_input":"2024-09-27T17:42:06.556247Z","iopub.status.idle":"2024-09-27T17:42:40.358242Z","shell.execute_reply.started":"2024-09-27T17:42:06.556195Z","shell.execute_reply":"2024-09-27T17:42:40.356976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# define a fucntion to read dicom file and return it's pixel\ndef read_dicom(dicom_file):\n    try:\n        dcm = pydicom.dcmread(dicom_file) # read dciom files\n\n        pixel_array = dcm.pixel_array\n\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\n        pixel_array = (pixel_array - pixel_array.min()) / (pixel_array.max() - pixel_array.min() + 1e-6)\n\n        if dcm.PhotometricInterpretation == \"MONOCHROME1\":\n            pixel_array = 1 - pixel_array\n\n        return (pixel_array * 255).astype(np.uint8)\n    except:\n        return np.zeros((294, 294)).astype(np.uint8)","metadata":{"papermill":{"duration":0.019642,"end_time":"2023-09-02T21:23:39.43277","exception":false,"start_time":"2023-09-02T21:23:39.413128","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-09-27T17:42:40.360097Z","iopub.execute_input":"2024-09-27T17:42:40.360455Z","iopub.status.idle":"2024-09-27T17:42:40.368692Z","shell.execute_reply.started":"2024-09-27T17:42:40.360424Z","shell.execute_reply":"2024-09-27T17:42:40.367753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nfrom PIL import Image\nimport os\n\ndef resize(img, size):\n    \"\"\"\n    Pad shorter side to square image\n    \"\"\"\n    H, W = img.shape  # (H, W)\n    if H > W:\n        o = (H - W) // 2\n        assert H - W - o >= 0\n        img = np.pad(img, [(0, 0), (o, H - W - o)])\n    elif H < W:\n        o = (W - H) // 2\n        assert W - H - o >= 0\n        img = np.pad(img, [(o, W - H - o), (0, 0)])\n    new_size = max(H, W)\n    assert img.shape == (new_size, new_size)\n\n    img_re = cv2.resize(np.expand_dims(img, 2),dsize=(size,size))\n\n    return img_re\n\ndef dcm_to_png_1(file, study_id, size):\n    image_id = file.split('/')[-1][:-4]\n    dicom_pixel = read_dicom(dicom_file=file)\n    h, w = dicom_pixel.shape\n    resized_img = resize(dicom_pixel, size) # resize it to specific size\n    image_path = \"/kaggle/working/\" + f\"{study_id}_{image_id}.png\"\n    cv2.imwrite(image_path, resized_img.astype(np.uint8))\n    pil_img = Image.open(image_path).convert('RGB')\n    img = np.asarray(pil_img)\n    os.remove(image_path) \n    return img\n\ndef dcm_to_png_2(file, study_id, size):\n    image_id = file.split('/')[-1][:-4]\n    dicom_pixel = read_dicom(dicom_file=file)\n    h, w = dicom_pixel.shape\n    resized_img = cv2.resize(dicom_pixel, (size, size)) # resize it to specific size\n    image_path = \"/kaggle/working/\" + f\"{study_id}_{image_id}.png\"\n    cv2.imwrite(image_path, resized_img.astype(np.uint8))\n    pil_img = Image.open(image_path).convert('RGB')\n    img = np.asarray(pil_img)\n    os.remove(image_path) \n    return img","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:42:40.371479Z","iopub.execute_input":"2024-09-27T17:42:40.372099Z","iopub.status.idle":"2024-09-27T17:42:40.384772Z","shell.execute_reply.started":"2024-09-27T17:42:40.372063Z","shell.execute_reply":"2024-09-27T17:42:40.383789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom_stack_1(dicom_folder, study_id, images, size, plane, reverse_sort=False):\n    dicoms = [pydicom.dcmread(os.path.join(dicom_folder, str(f)+\".dcm\")) for f in images]\n    plane = {\"sagittal\": 0, \"coronal\": 1, \"axial\": 2}[plane.lower()]\n    positions = np.asarray([float(d.ImagePositionPatient[plane]) for d in dicoms])\n    # if reverse_sort=False, then increasing array index will be from RIGHT->LEFT and CAUDAL->CRANIAL\n    # thus we do reverse_sort=True for axial so increasing array index is craniocaudal\n    idx = np.argsort(-positions if reverse_sort else positions)\n    ipp = np.asarray([d.ImagePositionPatient for d in dicoms]).astype(\"float\")[idx]\n    array = np.stack([dcm_to_png_1(os.path.join(dicom_folder, str(f)+\".dcm\"), study_id, size) for f in images])\n    array = array[idx]\n    images_stack = np.array(images)\n    images_stack = images_stack[idx]\n    \n    return array, images_stack\n\ndef load_dicom_stack_2(dicom_folder, study_id, images, size, plane, reverse_sort=False):\n    dicoms = [pydicom.dcmread(os.path.join(dicom_folder, str(f)+\".dcm\")) for f in images]\n    plane = {\"sagittal\": 0, \"coronal\": 1, \"axial\": 2}[plane.lower()]\n    positions = np.asarray([float(d.ImagePositionPatient[plane]) for d in dicoms])\n    # if reverse_sort=False, then increasing array index will be from RIGHT->LEFT and CAUDAL->CRANIAL\n    # thus we do reverse_sort=True for axial so increasing array index is craniocaudal\n    idx = np.argsort(-positions if reverse_sort else positions)\n    ipp = np.asarray([d.ImagePositionPatient for d in dicoms]).astype(\"float\")[idx]\n    array = np.stack([dcm_to_png_2(os.path.join(dicom_folder, str(f)+\".dcm\"), study_id, size) for f in images])\n    array = array[idx]\n    images_stack = np.array(images)\n    images_stack = images_stack[idx]\n    \n    return array, images_stack","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:42:40.38632Z","iopub.execute_input":"2024-09-27T17:42:40.386678Z","iopub.status.idle":"2024-09-27T17:42:40.39791Z","shell.execute_reply.started":"2024-09-27T17:42:40.386646Z","shell.execute_reply":"2024-09-27T17:42:40.396948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nimport timm\nfrom torch import nn\n\nN_EVAL = 8","metadata":{"papermill":{"duration":4.588221,"end_time":"2023-09-02T21:23:44.028696","exception":false,"start_time":"2023-09-02T21:23:39.440475","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-09-27T17:42:40.399183Z","iopub.execute_input":"2024-09-27T17:42:40.399609Z","iopub.status.idle":"2024-09-27T17:42:40.411827Z","shell.execute_reply.started":"2024-09-27T17:42:40.399572Z","shell.execute_reply":"2024-09-27T17:42:40.410926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn\nimport timm\nimport torch.nn.functional as F\nfrom torch.nn.parameter import Parameter\nimport math\nimport numpy as np\nfrom einops import rearrange, reduce, repeat\nfrom torch import Tensor\nimport einops\n\nclass SelfAttentionPooling(nn.Module):\n    def __init__(self, input_dim):\n        super(SelfAttentionPooling, self).__init__()\n        self.W = nn.Linear(input_dim, 1)\n        \n    def forward(self, x):\n        att_w = nn.functional.softmax(self.W(x).squeeze(dim=-1), dim=-1).unsqueeze(dim=-1)\n        x = torch.sum(x * att_w, dim=1)\n        return x\n\nclass Project0(nn.Module):\n    def __init__(self, feats):\n        super(Project0, self).__init__()\n        self.feats = feats\n        self.conv2d_0 = nn.Conv2d(feats, 256, kernel_size=1, stride=1)\n\n    def forward(self, f0):\n        f0 = self.conv2d_0(f0)\n        b = f0.shape[0]//8\n        s = 8\n        f0 = einops.rearrange(f0, '(b s) c h w -> b (s h w) c', b=b, s=s, h=f0.shape[2], w=f0.shape[3])\n\n        return f0\n    \nclass Project1(nn.Module):\n    def __init__(self, feats):\n        super(Project1, self).__init__()\n        self.feats = feats\n        self.conv2d_0 = nn.Conv2d(feats, 256, kernel_size=1, stride=1)\n\n    def forward(self, f0):\n        f0 = self.conv2d_0(f0)\n        b = f0.shape[0]//10\n        s = 10\n        f0 = einops.rearrange(f0, '(b s) c h w -> b (s h w) c', b=b, s=s, h=f0.shape[2], w=f0.shape[3])\n\n        return f0\n\nclass EmbeddingLayer(nn.Module):\n    def __init__(self, emb_size: int = 256, n_eval: int = 10, total_tokens: int = 10*49):\n        super(EmbeddingLayer, self).__init__()\n        self.cls_token = nn.Parameter(torch.randn(1,1, emb_size))\n        self.positions = nn.Parameter(torch.randn(total_tokens + 1, emb_size))\n        self.segment_embedding = nn.Parameter(torch.randn(n_eval, emb_size))\n        self.s = total_tokens//n_eval\n        self.n_eval = n_eval\n\n    def forward(self, x0):\n        b, _, _ = x0.shape\n        cls_tokens = repeat(self.cls_token, '() n e -> b n e', b=b)\n        x = torch.cat((cls_tokens, x0), dim=1)\n        \n        for i in range(self.n_eval):\n            x[:, 1+self.s*i:1+self.s*(i+1)] += self.segment_embedding[i]\n        x += self.positions\n\n        return x\n    \nclass MultiHeadAttention(nn.Module):\n    def __init__(self, emb_size: int = 256, num_heads: int = 8, dropout: float = 0):\n        super().__init__()\n        self.emb_size = emb_size\n        self.num_heads = num_heads\n        self.qkv = nn.Linear(emb_size, emb_size * 3)\n        self.att_drop = nn.Dropout(dropout)\n        self.projection = nn.Linear(emb_size, emb_size)\n\n    def forward(self, x : Tensor, mask: Tensor = None) -> Tensor:\n        qkv = rearrange(self.qkv(x), \"b n (h d qkv) -> (qkv) b h n d\", h=self.num_heads, qkv=3)\n        queries, keys, values = qkv[0], qkv[1], qkv[2]\n        energy = torch.einsum('bhqd, bhkd -> bhqk', queries, keys) \n        if mask is not None:\n            fill_value = torch.finfo(torch.float32).min\n            energy.mask_fill(~mask, fill_value)\n\n        scaling = self.emb_size ** (1/2)\n        att = F.softmax(energy, dim=-1) / scaling\n        att = self.att_drop(att)\n        out = torch.einsum('bhal, bhlv -> bhav ', att, values)\n        out = rearrange(out, \"b h n d -> b n (h d)\")\n        out = self.projection(out)\n        return out\n\n\nclass ResidualAdd(nn.Module):\n    def __init__(self, fn):\n        super().__init__()\n        self.fn = fn\n\n    def forward(self, x, **kwargs):\n        res = x\n        x = self.fn(x, **kwargs)\n        x += res\n        return x\n\n\nclass FeedForwardBlock(nn.Sequential):\n    def __init__(self, emb_size: int, expansion: int = 2, drop_p: float = 0.):\n        super().__init__(\n            nn.Linear(emb_size, expansion * emb_size),\n            nn.GELU(),\n            nn.Dropout(drop_p),\n            nn.Linear(expansion * emb_size, emb_size),\n        )\n\n\n\nclass TransformerEncoderBlock(nn.Sequential):\n    def __init__(self,\n                 emb_size: int = 256,\n                 drop_p: float = 0.,\n                 forward_expansion: int = 2,\n                 forward_drop_p: float = 0.,\n                 ** kwargs):\n        super().__init__(\n            ResidualAdd(nn.Sequential(\n                nn.LayerNorm(emb_size),\n                MultiHeadAttention(emb_size, **kwargs),\n                nn.Dropout(drop_p)\n            )),\n            ResidualAdd(nn.Sequential(\n                nn.LayerNorm(emb_size),\n                FeedForwardBlock(\n                    emb_size, expansion=forward_expansion, drop_p=forward_drop_p),\n                nn.Dropout(drop_p)\n            )\n            ))\n\n\nclass TransformerEncoder(nn.Sequential):\n    def __init__(self, depth: int = 8, **kwargs):\n        super().__init__(*[TransformerEncoderBlock(**kwargs) for _ in range(depth)])\n\nclass ClassificationHead(nn.Module):\n    def __init__(self, emb_size: int = 256, n_classes: int = 3):\n        super().__init__()\n        self.linear = nn.Linear(emb_size, n_classes)\n\n    def forward(self, x):\n        cls_token = x[:, 0]\n        return self.linear(cls_token)\n    \nclass ImageHead(nn.Module):\n    def __init__(self, input_dim):\n        super(ImageHead, self).__init__()\n        self.linear = nn.Linear(input_dim, input_dim)\n        self.object = nn.Linear(input_dim, 2)\n        \n    def forward(self, x):\n        x = self.linear(x)\n        x = self.object(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:42:40.413131Z","iopub.execute_input":"2024-09-27T17:42:40.41358Z","iopub.status.idle":"2024-09-27T17:42:40.4512Z","shell.execute_reply.started":"2024-09-27T17:42:40.413552Z","shell.execute_reply":"2024-09-27T17:42:40.450199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn\nimport timm\nimport torch.nn.functional as F\nfrom timm.layers.adaptive_avgmax_pool import SelectAdaptivePool2d\n\nclass AxialT23DModel1(nn.Module):\n    def __init__(self, back_bone, n_head, device_id):\n        super().__init__()\n        self.back_bone = back_bone\n        \n        self.model = timm.create_model(back_bone, pretrained=False, num_classes=3)\n        self.model.head = nn.Identity()\n        self.model.set_grad_checkpointing()\n        \n        feats = 768\n        drop = 0.0\n        if \"convnext\" in back_bone:\n            self.global_pool = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n        \n        lstm_embed = feats * 1\n        self.project = Project0(feats)\n        self.encoder = nn.Sequential(\n                                    EmbeddingLayer(n_eval=8,total_tokens=8*81),\n                                    TransformerEncoder()\n                                )\n        self.n_head = n_head\n        self.level_heads = nn.Linear(feats, 6).to(device_id)\n        self.volume_heads = nn.ModuleList([nn.Sequential(\n                                    SelfAttentionPooling(256),\n                                    nn.Linear(256, 3),\n                                ) for i in range(n_head)]).to(device_id)\n\n        self.device_id = device_id\n\n    def extract_feature(self, x0):\n        x0 = x0/255.0\n        x0 = x0.transpose(2, 3).transpose(2, 4).contiguous()\n        bs, N_EVAL_0, in_chans, h0, w0 = x0.shape\n        n_slice_per_c = N_EVAL_0\n        x0 = x0.reshape(bs * N_EVAL_0, in_chans, h0, w0)\n        features0 = self.model.forward_features(x0)\n        return features0, bs, n_slice_per_c\n        \n    def forward(self, x0):\n        features0, bs, n_slice_per_c = self.extract_feature(x0)\n        features0 = self.project(features0)\n        features = self.encoder(features0)\n        volume_logits = torch.zeros(self.n_head, bs, 3).to(self.device_id)\n        for i, l in enumerate(self.volume_heads):\n            volume_logits[i] = self.volume_heads[i](features)\n        return volume_logits\n    \nclass AxialT23DModel2(nn.Module):\n    def __init__(self, back_bone, n_head, device_id):\n        super().__init__()\n        self.back_bone = back_bone\n        \n        self.model = timm.create_model(back_bone, pretrained=False, num_classes=3)\n        self.model.head = nn.Identity()\n        self.model.set_grad_checkpointing()\n        \n        feats = 512\n        drop = 0.0\n        if \"convnext\" in back_bone:\n            self.global_pool = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n        \n        lstm_embed = feats * 1\n        self.project = Project0(feats)\n        self.encoder = nn.Sequential(\n                                    EmbeddingLayer(n_eval=8,total_tokens=8*256),\n                                    TransformerEncoder()\n                                )\n        self.n_head = n_head\n        self.level_heads = nn.Linear(feats, 10).to(device_id)\n        self.volume_heads = nn.ModuleList([nn.Sequential(\n                                    SelfAttentionPooling(256),\n                                    nn.Linear(256, 3),\n                                ) for i in range(n_head)]).to(device_id)\n\n        self.device_id = device_id\n\n    def extract_feature(self, x0):\n        x0 = x0/255.0\n        x0 = x0.transpose(2, 3).transpose(2, 4).contiguous()\n        bs, N_EVAL_0, in_chans, h0, w0 = x0.shape\n        n_slice_per_c = N_EVAL_0\n        x0 = x0.reshape(bs * N_EVAL_0, in_chans, h0, w0)\n        features0 = self.model.forward_features(x0)\n        return features0, bs, n_slice_per_c\n        \n    def forward(self, x0):\n        features0, bs, n_slice_per_c = self.extract_feature(x0)\n        features0 = self.project(features0)\n        features = self.encoder(features0)\n        volume_logits = torch.zeros(self.n_head, bs, 3).to(self.device_id)\n        for i, l in enumerate(self.volume_heads):\n            volume_logits[i] = self.volume_heads[i](features)\n        return volume_logits","metadata":{"papermill":{"duration":0.023122,"end_time":"2023-09-02T21:23:44.059764","exception":false,"start_time":"2023-09-02T21:23:44.036642","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-09-27T17:42:40.452608Z","iopub.execute_input":"2024-09-27T17:42:40.452923Z","iopub.status.idle":"2024-09-27T17:42:40.468519Z","shell.execute_reply.started":"2024-09-27T17:42:40.452896Z","shell.execute_reply":"2024-09-27T17:42:40.467464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SagittalT13DModel1(nn.Module):\n    def __init__(self, back_bone, n_head, device_id):\n        super().__init__()\n        self.back_bone = back_bone\n        \n        self.model0 = timm.create_model(back_bone, pretrained=False, num_classes=3)\n        self.model0.head = nn.Identity()\n        self.model0.set_grad_checkpointing()\n        \n        feats = 768\n        drop = 0.0\n        if \"convnext\" in back_bone:\n            self.global_pool0 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n            self.global_pool1 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n            self.global_pool2 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n        \n        lstm_embed = feats * 1\n        self.project = Project1(feats)\n        self.encoder = nn.Sequential(\n                                    EmbeddingLayer(n_eval=10, total_tokens=10*81),\n                                    TransformerEncoder()\n                                )\n        self.n_head = n_head\n        self.image_heads = nn.ModuleList([nn.Sequential(\n                                    nn.Linear(feats, 1),\n                                ) for i in range(10)]).to(device_id)\n        self.volume_heads = nn.ModuleList([nn.Sequential(\n                                    SelfAttentionPooling(256),\n                                    nn.Linear(256, 3),\n                                ) for i in range(n_head)]).to(device_id)\n\n        self.device_id = device_id\n\n    def extract_feature(self, x0):\n        x0 = x0/255.0\n        x0 = x0.transpose(2, 3).transpose(2, 4).contiguous()\n        bs, N_EVAL_0, in_chans, h0, w0 = x0.shape\n        n_slice_per_c = N_EVAL_0\n        x0 = x0.reshape(bs * N_EVAL_0, in_chans, h0, w0)\n        features0 = self.model0.forward_features(x0)\n        return features0, bs, n_slice_per_c\n        \n    def forward(self, x0):\n        features0, bs, n_slice_per_c = self.extract_feature(x0)\n        features0 = self.project(features0)\n        features = self.encoder(features0)\n        volume_logits = torch.zeros(self.n_head, bs, 3).to(self.device_id)\n        for i, l in enumerate(self.volume_heads):\n            volume_logits[i] = self.volume_heads[i](features)\n        return volume_logits\n    \nclass SagittalT13DModel2(nn.Module):\n    def __init__(self, back_bone, n_head, device_id):\n        super().__init__()\n        self.back_bone = back_bone\n        \n        self.model0 = timm.create_model(back_bone, pretrained=False, num_classes=3)\n        self.model0.head = nn.Identity()\n        self.model0.set_grad_checkpointing()\n        \n        feats = 512\n        drop = 0.0\n        if \"convnext\" in back_bone:\n            self.global_pool0 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n            self.global_pool1 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n            self.global_pool2 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n        \n        lstm_embed = feats * 1\n        self.project = Project1(feats)\n        self.encoder = nn.Sequential(\n                                    EmbeddingLayer(n_eval=10, total_tokens=10*256),\n                                    TransformerEncoder()\n                                )\n        self.n_head = n_head\n        self.image_heads = nn.ModuleList([nn.Sequential(\n                                    nn.Linear(feats, 1),\n                                ) for i in range(10)]).to(device_id)\n        self.volume_heads = nn.ModuleList([nn.Sequential(\n                                    SelfAttentionPooling(256),\n                                    nn.Linear(256, 3),\n                                ) for i in range(n_head)]).to(device_id)\n\n        self.device_id = device_id\n\n    def extract_feature(self, x0):\n        x0 = x0/255.0\n        x0 = x0.transpose(2, 3).transpose(2, 4).contiguous()\n        bs, N_EVAL_0, in_chans, h0, w0 = x0.shape\n        n_slice_per_c = N_EVAL_0\n        x0 = x0.reshape(bs * N_EVAL_0, in_chans, h0, w0)\n        features0 = self.model0.forward_features(x0)\n        return features0, bs, n_slice_per_c\n        \n    def forward(self, x0):\n        features0, bs, n_slice_per_c = self.extract_feature(x0)\n        features0 = self.project(features0)\n        features = self.encoder(features0)\n        volume_logits = torch.zeros(self.n_head, bs, 3).to(self.device_id)\n        for i, l in enumerate(self.volume_heads):\n            volume_logits[i] = self.volume_heads[i](features)\n        return volume_logits","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:42:40.472236Z","iopub.execute_input":"2024-09-27T17:42:40.472852Z","iopub.status.idle":"2024-09-27T17:42:40.488261Z","shell.execute_reply.started":"2024-09-27T17:42:40.472822Z","shell.execute_reply":"2024-09-27T17:42:40.487016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SagittalT23DModel1(nn.Module):\n    def __init__(self, back_bone, n_head, device_id):\n        super().__init__()\n        self.back_bone = back_bone\n        \n        self.model0 = timm.create_model(back_bone, pretrained=False, num_classes=3)\n        self.model0.head = nn.Identity()\n        self.model0.set_grad_checkpointing()\n        \n        feats = 768\n        drop = 0.0\n        if \"convnext\" in back_bone:\n            self.global_pool0 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n            self.global_pool1 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n            self.global_pool2 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n        \n        lstm_embed = feats * 1\n        self.project = Project1(feats)\n        self.encoder = nn.Sequential(\n                                    EmbeddingLayer(n_eval=10, total_tokens=10*81),\n                                    TransformerEncoder()\n                                )\n        self.n_head = n_head\n        self.image_heads = nn.ModuleList([nn.Sequential(\n                                    nn.Linear(feats, 1),\n                                ) for i in range(5)]).to(device_id)\n        self.volume_heads = nn.ModuleList([nn.Sequential(\n                                    SelfAttentionPooling(256),\n                                    nn.Linear(256, 3),\n                                ) for i in range(n_head)]).to(device_id)\n\n        self.device_id = device_id\n\n    def extract_feature(self, x0):\n        x0 = x0/255.0\n        x0 = x0.transpose(2, 3).transpose(2, 4).contiguous()\n        bs, N_EVAL_0, in_chans, h0, w0 = x0.shape\n        n_slice_per_c = N_EVAL_0\n        x0 = x0.reshape(bs * N_EVAL_0, in_chans, h0, w0)\n        features0 = self.model0.forward_features(x0)\n        return features0, bs, n_slice_per_c\n        \n    def forward(self, x0):\n        features0, bs, n_slice_per_c = self.extract_feature(x0)\n        features0 = self.project(features0)\n        features = self.encoder(features0)\n        volume_logits = torch.zeros(self.n_head, bs, 3).to(self.device_id)\n        for i, l in enumerate(self.volume_heads):\n            volume_logits[i] = self.volume_heads[i](features)\n        return volume_logits\n    \nclass SagittalT23DModel2(nn.Module):\n    def __init__(self, back_bone, n_head, device_id):\n        super().__init__()\n        self.back_bone = back_bone\n        \n        self.model0 = timm.create_model(back_bone, pretrained=False, num_classes=3)\n        self.model0.head = nn.Identity()\n        self.model0.set_grad_checkpointing()\n        \n        feats = 512\n        drop = 0.0\n        if \"convnext\" in back_bone:\n            self.global_pool0 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n            self.global_pool1 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n            self.global_pool2 = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n        \n        lstm_embed = feats * 1\n        self.project = Project1(feats)\n        self.encoder = nn.Sequential(\n                                    EmbeddingLayer(n_eval=10, total_tokens=10*256),\n                                    TransformerEncoder()\n                                )\n        self.n_head = n_head\n        self.image_heads = nn.ModuleList([nn.Sequential(\n                                    nn.Linear(feats, 1),\n                                ) for i in range(5)]).to(device_id)\n        self.volume_heads = nn.ModuleList([nn.Sequential(\n                                    SelfAttentionPooling(256),\n                                    nn.Linear(256, 3),\n                                ) for i in range(n_head)]).to(device_id)\n\n        self.device_id = device_id\n\n    def extract_feature(self, x0):\n        x0 = x0/255.0\n        x0 = x0.transpose(2, 3).transpose(2, 4).contiguous()\n        bs, N_EVAL_0, in_chans, h0, w0 = x0.shape\n        n_slice_per_c = N_EVAL_0\n        x0 = x0.reshape(bs * N_EVAL_0, in_chans, h0, w0)\n        features0 = self.model0.forward_features(x0)\n        return features0, bs, n_slice_per_c\n        \n    def forward(self, x0):\n        features0, bs, n_slice_per_c = self.extract_feature(x0)\n        features0 = self.project(features0)\n        features = self.encoder(features0)\n        volume_logits = torch.zeros(self.n_head, bs, 3).to(self.device_id)\n        for i, l in enumerate(self.volume_heads):\n            volume_logits[i] = self.volume_heads[i](features)\n        return volume_logits","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:42:40.489699Z","iopub.execute_input":"2024-09-27T17:42:40.490429Z","iopub.status.idle":"2024-09-27T17:42:40.505409Z","shell.execute_reply.started":"2024-09-27T17:42:40.490382Z","shell.execute_reply":"2024-09-27T17:42:40.504402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AxialClassifyModel(nn.Module):\n    def __init__(self, back_bone, n_head, device_id):\n        super().__init__()\n        self.model = timm.create_model(back_bone, pretrained=False, num_classes=3)\n        self.model.head = nn.Identity()\n        self.model.set_grad_checkpointing()\n\n        feats = 768\n        self.n_head = n_head\n        self.global_pool = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n        \n        self.heads = nn.ModuleList([nn.Sequential(\n                                    nn.Linear(feats, 1),\n                                ) for i in range(n_head)]).to(device_id)\n        self.device_id = device_id\n        self.back_bone = back_bone\n        \n    def forward(self, x):\n        x = x/255.0\n        x = x.transpose(1, 2).transpose(1, 3).contiguous()\n        #x = x.expand(-1, 3, -1, -1) \n        bs, _, _, _ = x.shape\n        \n        features = self.model.forward_features(x)\n        features = self.global_pool(features)\n\n        logits = torch.zeros(self.n_head, bs, 1).to(self.device_id)\n        for i, l in enumerate(self.heads):\n            logits[i] = self.heads[i](features)\n        return logits\n\nclass SaggitalT1ClassifyModel(nn.Module):\n    def __init__(self, back_bone, n_head, device_id):\n        super().__init__()\n        print(\"Number head: \", n_head)\n        self.model = timm.create_model(back_bone, pretrained=False, num_classes=3)\n        self.model.head = nn.Identity()\n        self.model.set_grad_checkpointing()\n\n        feats = 768\n        self.n_head = n_head\n        self.global_pool = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n        \n        self.heads = nn.ModuleList([nn.Sequential(\n                                    nn.Linear(feats, 1),\n                                ) for i in range(n_head)]).to(device_id)\n        self.device_id = device_id\n        self.back_bone = back_bone\n        \n    def forward(self, x):\n        x = x/255.0\n        x = x.transpose(1, 2).transpose(1, 3).contiguous()\n        #x = x.expand(-1, 3, -1, -1) \n        bs, _, _, _ = x.shape\n        features = self.model.forward_features(x)\n        features = self.global_pool(features)\n\n        logits = torch.zeros(self.n_head, bs, 3).to(self.device_id)\n        for i, l in enumerate(self.heads):\n            logits[i] = self.heads[i](features)\n        return logits\n    \nclass SaggitalT2ClassifyModel(nn.Module):\n    def __init__(self, back_bone, n_head, device_id):\n        super().__init__()\n        print(\"Number head: \", n_head)\n        self.model = timm.create_model(back_bone, pretrained=False, num_classes=3)\n        self.model.head = nn.Identity()\n        self.model.set_grad_checkpointing()\n\n        feats = 768\n        self.n_head = n_head\n        self.global_pool = SelectAdaptivePool2d(pool_type='avg', flatten=True, input_fmt=\"NCHW\")\n        \n        self.heads = nn.ModuleList([nn.Sequential(\n                                    nn.Linear(feats, 1),\n                                ) for i in range(n_head)]).to(device_id)\n        self.device_id = device_id\n        self.back_bone = back_bone\n        \n    def forward(self, x):\n        x = x/255.0\n        x = x.transpose(1, 2).transpose(1, 3).contiguous()\n        #x = x.expand(-1, 3, -1, -1) \n        bs, _, _, _ = x.shape\n        features = self.model.forward_features(x)\n        features = self.global_pool(features)\n\n        logits = torch.zeros(self.n_head, bs, 3).to(self.device_id)\n        for i, l in enumerate(self.heads):\n            logits[i] = self.heads[i](features)\n        return logits","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:42:40.507277Z","iopub.execute_input":"2024-09-27T17:42:40.507707Z","iopub.status.idle":"2024-09-27T17:42:40.531043Z","shell.execute_reply.started":"2024-09-27T17:42:40.50767Z","shell.execute_reply":"2024-09-27T17:42:40.53016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nseries_df = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv\")\nseries_map = {}\naxis_map = {}\nfor index, row in series_df.iterrows():\n    series_id = row[\"series_id\"]\n    study_id = row[\"study_id\"]\n    series_description = row[\"series_description\"]\n    series_map[str(series_id)] = series_description\n    if series_description == \"Sagittal T2/STIR\":\n        axis = \"sagittal_t2\"\n    elif series_description == \"Sagittal T1\":\n        axis = \"sagittal_t1\"\n    elif series_description == \"Axial T2\":\n        axis = \"axial_t2\"\n    if study_id not in axis_map.keys():\n        axis_map[study_id] = {}\n    axis_map[study_id][series_id] = axis","metadata":{"execution":{"iopub.status.busy":"2024-09-27T17:42:40.532141Z","iopub.execute_input":"2024-09-27T17:42:40.53246Z","iopub.status.idle":"2024-09-27T17:42:41.041629Z","shell.execute_reply.started":"2024-09-27T17:42:40.532435Z","shell.execute_reply":"2024-09-27T17:42:41.040451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_old_weight(model, weight_path):\n    if len(weight_path) > 0:\n        pretrained_dict = torch.load(weight_path)\n        print(\"Load model: \", weight_path)\n        model.load_state_dict(pretrained_dict)\n    return model\n\ndef predict_worker(q_in):\n    level_list = [\"l1_l2\", \"l2_l3\", \"l3_l4\", \"l4_l5\", \"l5_s1\"]\n    cond_list = [\"left_subarticular_stenosis\", \"right_subarticular_stenosis\", \"left_neural_foraminal_narrowing\", \"right_neural_foraminal_narrowing\", \"spinal_canal_stenosis\"]\n    col_names = []\n    preds_result = {}\n    for cond in cond_list:\n        for level in level_list:\n            col_name = cond + \"_\" + level\n            col_names.append(col_name)\n    axial_t2_model_paths = [\"/kaggle/input/rsna-2024-axialt2-3d-models/axial_t2_loss_exp_22_fold0.pt\",\n                            \"/kaggle/input/rsna-2024-axialt2-3d-models/axial_t2_loss_exp_22_fold1.pt\",\n                            \"/kaggle/input/rsna-2024-axialt2-3d-models/axial_t2_loss_exp_22_fold2.pt\",\n                            \"/kaggle/input/rsna-2024-axialt2-3d-models/axial_t2_loss_exp_22_fold3.pt\",\n                            \"/kaggle/input/rsna-2024-axialt2-3d-models/axial_t2_loss_exp_22_fold4.pt\",\n                            \"/kaggle/input/rsna-2024-axialt2-3d-models/axial_t2_loss_exp_24_fold0.pt\",\n                           \"/kaggle/input/rsna-2024-axialt2-3d-models/axial_t2_loss_exp_24_fold1.pt\",\n                           \"/kaggle/input/rsna-2024-axialt2-3d-models/axial_t2_loss_exp_24_fold2.pt\",\n                           \"/kaggle/input/rsna-2024-axialt2-3d-models/axial_t2_loss_exp_24_fold3.pt\",\n                           \"/kaggle/input/rsna-2024-axialt2-3d-models/axial_t2_loss_exp_24_fold4.pt\"]\n    axial_t2_model_types = [1, 1, 1, 1, 1, 2, 2, 2, 2, 2]\n    axial_t2_models = []\n    for axial_t2_model_path, model_type in zip(axial_t2_model_paths, axial_t2_model_types):\n        if model_type == 1:\n            axial_t2_model = AxialT23DModel1(\"convnext_tiny.in12k_ft_in1k\", 10, \"cuda:0\")\n        elif model_type == 2:\n            axial_t2_model = AxialT23DModel2(\"convnext_pico_ols.d1_in1k\", 10, \"cuda:0\")\n        axial_t2_model = load_old_weight(axial_t2_model, axial_t2_model_path)\n        axial_t2_model.to(\"cuda:0\")\n        axial_t2_model.eval()\n        axial_t2_models.append(axial_t2_model)\n    \n    sagittal_t1_model_paths = [\"/kaggle/input/rsna-2024-sagittal-t1-3d-models/sagittal_t1_loss_exp_20_fold0.pt\",\n                               \"/kaggle/input/rsna-2024-sagittal-t1-3d-models/sagittal_t1_loss_exp_20_fold1.pt\",\n                               \"/kaggle/input/rsna-2024-sagittal-t1-3d-models/sagittal_t1_loss_exp_20_fold2.pt\",\n                               \"/kaggle/input/rsna-2024-sagittal-t1-3d-models/sagittal_t1_loss_exp_20_fold3.pt\",\n                               \"/kaggle/input/rsna-2024-sagittal-t1-3d-models/sagittal_t1_loss_exp_20_fold4.pt\",\n                            \"/kaggle/input/rsna-2024-sagittal-t1-3d-models/sagittal_t1_loss_exp_21_fold0.pt\",\n                              \"/kaggle/input/rsna-2024-sagittal-t1-3d-models/sagittal_t1_loss_exp_21_fold1.pt\",\n                              \"/kaggle/input/rsna-2024-sagittal-t1-3d-models/sagittal_t1_loss_exp_21_fold2.pt\",\n                              \"/kaggle/input/rsna-2024-sagittal-t1-3d-models/sagittal_t1_loss_exp_21_fold3.pt\",\n                              \"/kaggle/input/rsna-2024-sagittal-t1-3d-models/sagittal_t1_loss_exp_21_fold4.pt\"]\n    sagittal_t1_model_types = [1, 1, 1, 1, 1, 2, 2, 2, 2, 2]\n    sagittal_t1_models = []\n    for sagittal_t1_model_path, model_type in zip(sagittal_t1_model_paths, sagittal_t1_model_types):\n        if model_type == 1:\n            sagittal_t1_model = SagittalT13DModel1(\"convnext_tiny.in12k_ft_in1k\", 10, \"cuda:0\")\n        elif model_type == 2:\n            sagittal_t1_model = SagittalT13DModel2(\"convnext_pico_ols.d1_in1k\", 10, \"cuda:0\")\n        sagittal_t1_model = load_old_weight(sagittal_t1_model, sagittal_t1_model_path)\n        sagittal_t1_model.to(\"cuda:0\")\n        sagittal_t1_model.eval()\n        sagittal_t1_models.append(sagittal_t1_model)\n        \n    sagittal_t2_model_paths = [\"/kaggle/input/rsna-2024-sagittal-t2-3d-models/sagittal_t2_loss_exp_20_fold0.pt\",\n                               \"/kaggle/input/rsna-2024-sagittal-t2-3d-models/sagittal_t2_loss_exp_20_fold1.pt\",\n                               \"/kaggle/input/rsna-2024-sagittal-t2-3d-models/sagittal_t2_loss_exp_20_fold2.pt\",\n                               \"/kaggle/input/rsna-2024-sagittal-t2-3d-models/sagittal_t2_loss_exp_20_fold3.pt\",\n                               \"/kaggle/input/rsna-2024-sagittal-t2-3d-models/sagittal_t2_loss_exp_20_fold4.pt\",\n                                \"/kaggle/input/rsna-2024-sagittal-t2-3d-models/sagittal_t2_loss_exp_21_fold0.pt\",\n                              \"/kaggle/input/rsna-2024-sagittal-t2-3d-models/sagittal_t2_loss_exp_21_fold1.pt\",\n                              \"/kaggle/input/rsna-2024-sagittal-t2-3d-models/sagittal_t2_loss_exp_21_fold2.pt\",\n                              \"/kaggle/input/rsna-2024-sagittal-t2-3d-models/sagittal_t2_loss_exp_21_fold3.pt\",\n                              \"/kaggle/input/rsna-2024-sagittal-t2-3d-models/sagittal_t2_loss_exp_21_fold4.pt\"]\n    sagittal_t2_model_types = [1, 1, 1, 1, 1, 2, 2, 2, 2, 2]\n    sagittal_t2_models = []\n    for sagittal_t2_model_path, model_type in zip(sagittal_t2_model_paths, sagittal_t2_model_types):\n        if model_type == 1:\n            sagittal_t2_model = SagittalT23DModel1(\"convnext_tiny.in12k_ft_in1k\", 5, \"cuda:0\")\n        elif model_type == 2:\n            sagittal_t2_model = SagittalT23DModel2(\"convnext_pico_ols.d1_in1k\", 5, \"cuda:0\")\n        sagittal_t2_model = load_old_weight(sagittal_t2_model, sagittal_t2_model_path)\n        sagittal_t2_model.to(\"cuda:0\")\n        sagittal_t2_model.eval()\n        sagittal_t2_models.append(sagittal_t2_model)\n        \n    with torch.no_grad():\n        while True:\n            batch_queue = q_in.get()\n            if batch_queue is None:\n                break\n\n            study_id = batch_queue[\"study_id\"]\n            axis = batch_queue[\"axis\"]\n            if axis == 'axial_t2':\n                models = axial_t2_models\n                model_types = axial_t2_model_types\n                ensemble_probs = torch.zeros((10, 3)).cuda()\n                batch_288 = torch.from_numpy(np.array(batch_queue[\"batch_288\"]))\n                batch_512 = torch.from_numpy(np.array(batch_queue[\"batch_512\"]))\n            elif axis == 'sagittal_t1':\n                models = sagittal_t1_models\n                model_types = sagittal_t1_model_types\n                ensemble_probs = torch.zeros((10, 3)).cuda()\n                batch_288 = torch.from_numpy(np.array(batch_queue[\"batch_288\"]))\n                batch_512 = torch.from_numpy(np.array(batch_queue[\"batch_512\"]))\n            elif axis == 'sagittal_t2': \n                models = sagittal_t2_models\n                model_types = sagittal_t2_model_types\n                ensemble_probs = torch.zeros((5, 3)).cuda()\n                batch_288 = torch.from_numpy(np.array(batch_queue[\"batch_288\"]))\n                batch_512 = torch.from_numpy(np.array(batch_queue[\"batch_512\"]))\n            \n            batch_288 = batch_288.to(\"cuda:0\").float()\n            batch_288 = torch.unsqueeze(batch_288, 0)\n            batch_512 = batch_512.to(\"cuda:0\").float()\n            batch_512 = torch.unsqueeze(batch_512, 0)\n            \n            for i in range(len(models)):\n                model_type = model_types[i]\n                if model_type == 1:\n                    logits = models[i](batch_288)[:, 0, :]\n                elif model_type == 2:\n                    logits = models[i](batch_512)[:, 0, :]\n                probs = torch.nn.Softmax(dim=1)(logits)\n                ensemble_probs += probs\n                \n            if len(models) > 0:\n                ensemble_probs /= len(models)\n                ensemble_probs = ensemble_probs.detach().cpu().numpy()\n                r = {\"series_id\": series_id, \"prob\": ensemble_probs, \"axis\": axis}\n    #             print(series_id, ensemble_probs)\n                if study_id not in preds_result.keys():\n                    preds_result[study_id] = {}\n                if axis not in preds_result[study_id].keys():\n                    preds_result[study_id][axis] = []\n                preds_result[study_id][axis].append(r)\n    rows = []\n    for study_id in preds_result.keys():\n        probs = np.zeros((len(col_names), 3))\n        \n        for axis in preds_result[study_id].keys():\n            prob_axis = np.zeros((preds_result[study_id][axis][0][\"prob\"].shape[0], 3)) \n            for i in range(preds_result[study_id][axis][0][\"prob\"].shape[0]):\n                min_normal = 1\n                min_index = 0\n                for index, r in enumerate(preds_result[study_id][axis]):\n                    if r[\"prob\"][i][0] < min_normal:\n                        min_normal = r[\"prob\"][i][0]\n                        min_index = index\n                prob = preds_result[study_id][axis][min_index][\"prob\"][i]      \n                prob_axis[i] = prob\n            if axis == 'axial_t2':\n                probs[0:10,:] += prob_axis\n            elif axis == 'sagittal_t1':\n                probs[10:20,:] += prob_axis\n            elif axis == 'sagittal_t2': \n                probs[20:25,:] += prob_axis\n       \n        for i, col_name in enumerate(col_names):\n            row_id = str(study_id) + \"_\" + col_name\n            rows.append({\"row_id\": row_id, \"normal_mild\": probs[i][0], \"moderate\": probs[i][1], \"severe\": probs[i][2]})\n    if len(rows) > 0:\n        submission = pd.DataFrame.from_dict(rows) \n        submission.to_csv('submission.csv', index=False)\n        submission.head()","metadata":{"papermill":{"duration":0.023905,"end_time":"2023-09-02T21:23:44.091107","exception":false,"start_time":"2023-09-02T21:23:44.067202","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-09-27T17:42:41.043528Z","iopub.execute_input":"2024-09-27T17:42:41.04386Z","iopub.status.idle":"2024-09-27T17:42:41.071715Z","shell.execute_reply.started":"2024-09-27T17:42:41.043831Z","shell.execute_reply":"2024-09-27T17:42:41.070607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.019235,"end_time":"2023-09-02T21:23:46.122514","exception":false,"start_time":"2023-09-02T21:23:46.103279","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nfrom operator import itemgetter\nDEBUG = False\nif DEBUG == True:\n    test_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images\"\nelse:\n    test_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images\"\n\ndef read_worker(q_in_predict, study_id_list):\n    model_axial_paths = [\"/kaggle/input/rsna-2024-axial-classify-model/auc_exp_0_axial_t2_fold0.pt\",\n                        \"/kaggle/input/rsna-2024-axial-classify-model/auc_exp_0_axial_t2_fold1.pt\",\n                        \"/kaggle/input/rsna-2024-axial-classify-model/auc_exp_0_axial_t2_fold2.pt\",\n                        \"/kaggle/input/rsna-2024-axial-classify-model/auc_exp_0_axial_t2_fold3.pt\",\n                        \"/kaggle/input/rsna-2024-axial-classify-model/auc_exp_0_axial_t2_fold4.pt\"]\n    axial_models = []\n    for model_path in model_axial_paths:\n        model = AxialClassifyModel(\"convnext_tiny.in12k_ft_in1k\", 10, \"cuda:0\")\n        model = load_old_weight(model, model_path)\n        model.to(\"cuda:0\")\n        model.eval()\n        axial_models.append(model)\n        \n    model_sagittal_t1_paths = [\"/kaggle/input/rsna-2024-saggital-classify-model/auc_exp_10_sagittal_t1_fold0.pt\",\n                              \"/kaggle/input/rsna-2024-saggital-classify-model/auc_exp_10_sagittal_t1_fold1.pt\",\n                              \"/kaggle/input/rsna-2024-saggital-classify-model/auc_exp_10_sagittal_t1_fold2.pt\",\n                              \"/kaggle/input/rsna-2024-saggital-classify-model/auc_exp_10_sagittal_t1_fold3.pt\",\n                              \"/kaggle/input/rsna-2024-saggital-classify-model/auc_exp_10_sagittal_t1_fold4.pt\"]\n    sagittal_t1_models = []\n    for model_path in model_sagittal_t1_paths:\n        model = SaggitalT1ClassifyModel(\"convnext_tiny.in12k_ft_in1k\", 5, \"cuda:0\")\n        model = load_old_weight(model, model_path)\n        model.to(\"cuda:0\")\n        model.eval()\n        sagittal_t1_models.append(model)\n        \n    model_sagittal_t2_paths = [\"/kaggle/input/rsna-2024-saggital-classify-model/auc_exp_10_sagittal_t2_fold0.pt\",\n                              \"/kaggle/input/rsna-2024-saggital-classify-model/auc_exp_10_sagittal_t2_fold1.pt\",\n                              \"/kaggle/input/rsna-2024-saggital-classify-model/auc_exp_10_sagittal_t2_fold2.pt\",\n                              \"/kaggle/input/rsna-2024-saggital-classify-model/auc_exp_10_sagittal_t2_fold3.pt\",\n                              \"/kaggle/input/rsna-2024-saggital-classify-model/auc_exp_10_sagittal_t2_fold4.pt\"]\n    sagittal_t2_models = []\n    for model_path in model_sagittal_t2_paths:\n        model = SaggitalT2ClassifyModel(\"convnext_tiny.in12k_ft_in1k\", 5, \"cuda:0\")\n        model = load_old_weight(model, model_path)\n        model.to(\"cuda:0\")\n        model.eval()\n        sagittal_t2_models.append(model)\n\n    for study_id in study_id_list:\n        patient_path = os.path.join(test_path, study_id)\n        for series in axis_map[int(study_id)].keys():\n            axis = axis_map[int(study_id)][series]\n            series_path = os.path.join(patient_path, str(series))\n            images = os.listdir(series_path)\n            images = [imagename.replace('.dcm', '') for imagename in images]\n            images = list(map(int, images))\n            images = sorted(images)\n            \n            image_paths = []\n            for imagename in images:\n                image_path = os.path.join(series_path, str(imagename) + \".dcm\")\n                image_paths.append(image_path)\n\n            total_images = len(image_paths)\n            step = 1\n\n            series_288_batch = []\n            series_294_batch = []\n            \n            batch_size_2d = 8\n            batch_2d = torch.zeros((batch_size_2d, 384, 384, 3)).cuda()\n            c_batch = 0\n            total_probs = []\n            total_sum_probs = []\n            if \"axial\" in axis:\n                n_eval = 8\n                n_class = 10 \n                n_split = 1\n            else:\n                n_eval = 10\n                n_class = 5\n                n_split = 2\n            for i in range(0, total_images, step):\n                batch_2d[c_batch] = torch.from_numpy(dcm_to_png_2(image_paths[i], study_id, 384)).cuda()\n                c_batch += 1\n                if c_batch == batch_size_2d or c_batch >= total_images-step:\n                    if \"axial\" in axis:\n                        ensemble_probs = torch.zeros(c_batch, 10).cuda()\n                        models = axial_models\n                    elif \"sagittal_t1\" in axis:\n                        ensemble_probs = torch.zeros(c_batch, 5).cuda()\n                        models = sagittal_t1_models\n                    else:\n                        ensemble_probs = torch.zeros(c_batch, 5).cuda()\n                        models = sagittal_t2_models\n                    with torch.no_grad():\n                        for j in range(len(models)):\n                            logits = models[j](batch_2d[:c_batch])[:, :, 0]\n                            logits = torch.permute(logits, (1, 0))\n                            probs = torch.nn.Sigmoid()(logits)\n                            ensemble_probs += probs\n                        if len(models) > 0:\n                            ensemble_probs /= len(models)\n                            ensemble_probs = ensemble_probs.detach().cpu().numpy()\n                        total_batch_sum_prob =  np.sum(ensemble_probs, axis=1)  \n                        for j in range(total_batch_sum_prob.shape[0]):\n                            total_sum_probs.append((total_batch_sum_prob[j], i+1-total_batch_sum_prob.shape[0]+j))\n                            total_probs.append((ensemble_probs[j], i+1-ensemble_probs.shape[0]+j))\n                    c_batch = 0\n                \n            select_list = set()\n            n_per_split = total_images//n_split\n            for split in range(n_split):\n                for i in range(n_class):\n                    max_prob = 0\n                    max_index = -1\n                    for probs, index in total_probs[split*n_per_split:(split+1)*n_per_split]:\n                        prob = probs[i]\n                        if prob > max_prob:\n                            max_prob = prob\n                            max_index = index\n                    image_path = image_paths[max_index]\n                    select_list.add(os.path.basename(image_path).replace(\".dcm\", \"\"))\n                print(\"select_list: \", axis, split, select_list) \n                \n            batch_list = sorted(total_sum_probs, key=lambda tup: tup[0])\n            n_left = n_eval - len(select_list)\n            cc = len(batch_list) - 1 \n            \n            while n_left > 0 and cc > 0:\n                index = batch_list[cc][1]\n                image_path = image_paths[index]\n                ins_number = os.path.basename(image_path).replace(\".dcm\", \"\")\n                if ins_number not in select_list:\n                    select_list.add(ins_number)\n                    n_left -= 1\n                cc -= 1\n                \n            max_index = len(batch_list) - 1\n            n_left = n_eval - len(select_list)\n            select_list = list(select_list)\n            for i in range(n_left):\n                index = batch_list[max_index][1]\n                image_path = image_paths[index]\n                select_list.append(os.path.basename(image_path).replace(\".dcm\", \"\"))\n            \n#             batch_list = sorted(batch_list, key=lambda tup: tup[1])\n            \n#             print(batch_list)\n            \n#             image_names = []\n#             for _, i in batch_list:\n#                 image_path = image_paths[i]\n#                 image_names.append(os.path.basename(image_path).replace(\".dcm\", \"\"))\n#             print(axis, image_names)\n            \n            if \"axial\" in axis:\n                series_288_batch, images = load_dicom_stack_1(series_path, study_id, select_list, 288, plane=\"axial\", reverse_sort=True)\n                series_512_batch, images = load_dicom_stack_2(series_path, study_id, select_list, 512, plane=\"axial\", reverse_sort=True)\n            else:\n                series_288_batch, images = load_dicom_stack_2(series_path, study_id, select_list, 288, plane=\"sagittal\")\n                series_512_batch, images = load_dicom_stack_2(series_path, study_id, select_list, 512, plane=\"sagittal\")\n            print(axis, images)    \n            series_288_batch = np.array(series_288_batch)\n            series_512_batch = np.array(series_512_batch)\n            batch_queue = {\"study_id\": study_id, \"axis\": axis, \"batch_512\": series_512_batch, \"batch_288\": series_288_batch}\n            q_in_predict.put(batch_queue, block=True, timeout=None)\n\n        q_in_predict.put(batch_queue, block=True, timeout=None)","metadata":{"papermill":{"duration":0.021051,"end_time":"2023-09-02T21:23:46.151085","exception":false,"start_time":"2023-09-02T21:23:46.130034","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-09-27T17:42:41.073328Z","iopub.execute_input":"2024-09-27T17:42:41.073764Z","iopub.status.idle":"2024-09-27T17:42:41.102709Z","shell.execute_reply.started":"2024-09-27T17:42:41.073732Z","shell.execute_reply":"2024-09-27T17:42:41.101721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import multiprocessing as mp\nfrom joblib import Parallel, delayed\n\nfrom multiprocessing import Process, Queue, Manager, Lock\nimport pandas as pd\nimport glob\n\n\nq_in_predict = Queue(maxsize=10)\n\npatient_id_list = os.listdir(test_path)\npatient_id_list = list(map(str, patient_id_list))\n\npredict_process = Process(target=predict_worker, args=(q_in_predict, ))\npredict_process.start()\ntotal_patient = len(patient_id_list)\nread_0_process = Process(target=read_worker, args=(q_in_predict, patient_id_list[:total_patient//3], ))\nread_1_process = Process(target=read_worker, args=(q_in_predict, patient_id_list[total_patient//3:2*total_patient//3], ))\nread_2_process = Process(target=read_worker, args=(q_in_predict, patient_id_list[2*total_patient//3:], ))\nread_0_process.start()\nread_1_process.start()\nread_2_process.start()\n\nread_0_process.join()\nread_1_process.join()\nread_2_process.join()\nq_in_predict.put(None, block=True, timeout=None)\n\npredict_process.join()\n","metadata":{"papermill":{"duration":16.200694,"end_time":"2023-09-02T21:24:02.359101","exception":false,"start_time":"2023-09-02T21:23:46.158407","status":"completed"},"tags":[],"execution":{"iopub.status.idle":"2024-09-27T18:38:38.630872Z","shell.execute_reply.started":"2024-09-27T17:42:41.105102Z","shell.execute_reply":"2024-09-27T18:38:38.629586Z"},"trusted":true},"execution_count":null,"outputs":[]}]}