{"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":"markdown","source":"# Local dependencies install\n\n- timm to use pretrained models\n- pylibjpeg to open all dcm files","metadata":{}},{"cell_type":"code","source":"# !pip install timm pylibjpeg[all] pydicom\n# credits goes to: https://www.kaggle.com/code/dragonzhang/rsna-efficientnetv2-inference-tensorflow\n\n# Source: https://www.kaggle.com/code/remekkinas/fast-dicom-processing-1-6-2x-faster?scriptVersionId=113360473\n\ntry:\n    import pylibjpeg\nexcept:\n   !pip install /kaggle/input/rsna-2022-whl/{pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}\n","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:32:13.475317Z","iopub.execute_input":"2023-01-24T16:32:13.475941Z","iopub.status.idle":"2023-01-24T16:32:48.594428Z","shell.execute_reply.started":"2023-01-24T16:32:13.475827Z","shell.execute_reply":"2023-01-24T16:32:48.592919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls /kaggle/input/timmmaster/","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:32:48.597132Z","iopub.execute_input":"2023-01-24T16:32:48.597517Z","iopub.status.idle":"2023-01-24T16:32:49.733771Z","shell.execute_reply.started":"2023-01-24T16:32:48.59748Z","shell.execute_reply":"2023-01-24T16:32:49.732502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp -r /kaggle/input/timmmaster/ ./\n!pip install /kaggle/working/timmmaster","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:32:49.735075Z","iopub.execute_input":"2023-01-24T16:32:49.735443Z","iopub.status.idle":"2023-01-24T16:33:30.142519Z","shell.execute_reply.started":"2023-01-24T16:32:49.73541Z","shell.execute_reply":"2023-01-24T16:33:30.141334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nfrom path import Path\nimport matplotlib.pyplot as plt\nimport pydicom\nfrom torch.utils.data import DataLoader\nfrom torchvision.transforms import Resize\nimport cv2\nimport timm\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\ni = 0\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n        i += 1\n        if i > 20: break\n    if i > 20: break\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-24T16:33:30.14582Z","iopub.execute_input":"2023-01-24T16:33:30.146358Z","iopub.status.idle":"2023-01-24T16:33:34.542319Z","shell.execute_reply.started":"2023-01-24T16:33:30.146305Z","shell.execute_reply":"2023-01-24T16:33:34.540916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# put timm saved model under ~/.cache/torch/hub/checkpoints/ to load pretrained model\n\n!mkdir --parents ~/.cache/torch/hub/checkpoints/\n!cp  /kaggle/input/convnextbase/convnext_base_1k_224_ema.pth ~/.cache/torch/hub/checkpoints/","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:24.393567Z","iopub.execute_input":"2023-01-24T16:34:24.394173Z","iopub.status.idle":"2023-01-24T16:34:30.193565Z","shell.execute_reply.started":"2023-01-24T16:34:24.394129Z","shell.execute_reply":"2023-01-24T16:34:30.191919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path = Path('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ntest_path = Path('/kaggle/input/rsna-breast-cancer-detection/test.csv')\ntrain_images_path = Path('/kaggle/input/rsna-breast-cancer-detection/train_images')\ntest_images_path = Path('/kaggle/input/rsna-breast-cancer-detection/test_images')\n","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:34.892335Z","iopub.execute_input":"2023-01-24T16:34:34.892887Z","iopub.status.idle":"2023-01-24T16:34:34.901161Z","shell.execute_reply.started":"2023-01-24T16:34:34.892841Z","shell.execute_reply":"2023-01-24T16:34:34.899506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_image(patient_id, image_id):\n    full_path = train_images_path / str(patient_id) / (str(image_id) + '.dcm')\n    assert full_path.exists(), f'{full_path}  doesnt exists'\n    return full_path\n    \n\ndef get_test_image(patient_id, image_id):\n    full_path = test_images_path / str(patient_id) / (str(image_id) + '.dcm')\n    assert full_path.exists(), f'{full_path}  doesnt exists'\n    return full_path","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:35.422031Z","iopub.execute_input":"2023-01-24T16:34:35.422538Z","iopub.status.idle":"2023-01-24T16:34:35.429865Z","shell.execute_reply.started":"2023-01-24T16:34:35.422491Z","shell.execute_reply":"2023-01-24T16:34:35.428707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Preprocessing functions (thanks to https://www.kaggle.com/code/markwijkhuizen/rsna-convnextv2-inference-tensorflow\n)","metadata":{}},{"cell_type":"code","source":"# thanks to https://www.kaggle.com/code/markwijkhuizen/rsna-convnextv2-inference-tensorflow for the preprocessing procedure\n\nIS_INTERACTIVE = os.environ['KAGGLE_KERNEL_RUN_TYPE'] == 'Interactive'\n\nTARGET_HEIGHT = 1344\nTARGET_WIDTH = 768\nN_CHANNELS = 1\nINPUT_SHAPE = (N_CHANNELS, TARGET_HEIGHT, TARGET_WIDTH)\nTARGET_HEIGHT_WIDTH_RATIO = TARGET_HEIGHT / TARGET_WIDTH\nTHRESHOLD_BEST = 0.50\n\nCLAHE = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(32, 32))\n\nCROP_IMAGE = True\nAPPLY_CLAHE = False\nAPPLY_EQ_HIST = False","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:37.668386Z","iopub.execute_input":"2023-01-24T16:34:37.669251Z","iopub.status.idle":"2023-01-24T16:34:37.681387Z","shell.execute_reply.started":"2023-01-24T16:34:37.6692Z","shell.execute_reply":"2023-01-24T16:34:37.68004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def smooth(l):\n    kernel_size = int(len(l) * 0.01)\n    kernel = np.ones(kernel_size) / kernel_size\n    return np.convolve(l, kernel, mode='same')\n\ndef get_x_offset(image, max_col_sum_ratio_threshold=0.05):\n    margin = 0\n    sums = smooth(image.sum(axis=0).squeeze())\n    sums_argmax = sums[:int(image.shape[1] * 0.75)].argmax()\n    sums_threshold = sums.max() * max_col_sum_ratio_threshold\n    first_non_zoro_column_found = False\n    \n    for offset, s in enumerate(sums):\n        if s < sums_threshold and first_non_zoro_column_found:\n            return min(image.shape[1], offset + margin)\n        elif s > sums_threshold and offset > sums_argmax:\n            first_non_zoro_column_found = True\n        \n    return offset\n\ndef get_y_offsets(image, max_row_sum_ratio_threshold=0.10):\n    margin = 0\n    sums = smooth(image.sum(axis=1).squeeze())\n    sums_argmax = int(image.shape[0] * 0.25) + sums[int(image.shape[0] * 0.25):int(image.shape[0] * 0.75)].argmax()\n    sum_threshold = sums.max() * max_row_sum_ratio_threshold\n    offset_bottom = 0\n    offset_top = image.shape[0]\n    offset_top_set = False\n\n    # Bottom offset\n    for offset, s in enumerate(sums):\n        if s < sum_threshold and not offset_top_set:\n            offset_bottom += 1\n        else:\n            break\n            \n    for offset, s in enumerate(reversed(sums)):\n        if s > sum_threshold and not offset_top_set:\n            offset_top = image.shape[0] - (offset + 1)\n            break\n            \n    return max(0, offset_bottom - margin), min(image.shape[0], offset_top + margin)\n\ndef crop(image, debug=False):\n    x_offset = get_x_offset(image)\n    offset_bottom, offset_top = get_y_offsets(image[:,:x_offset])\n    \n    image = image[offset_bottom:offset_top:,:x_offset]\n        \n    return image","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:38.034441Z","iopub.execute_input":"2023-01-24T16:34:38.034955Z","iopub.status.idle":"2023-01-24T16:34:38.061016Z","shell.execute_reply.started":"2023-01-24T16:34:38.034917Z","shell.execute_reply":"2023-01-24T16:34:38.059369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process(dicom, size=(TARGET_WIDTH, TARGET_HEIGHT), crop_image=CROP_IMAGE, apply_clahe=APPLY_CLAHE, apply_eq_hist=APPLY_EQ_HIST, debug=False, save=True):\n    # Read Dicom File\n    image = dicom.pixel_array\n\n    # Normalize [0,1] range\n    image = (image - image.min()) / (image.max() - image.min())\n\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":  \n        image = 1 - image\n\n    # Flip T0 Left/Right Orientation\n    h0, w0 = image.shape\n    if image[:,int(-w0 * 0.10):].sum() > image[:,:int(w0 * 0.10)].sum():\n        image = np.flip(image, axis=1)\n    \n    # Save original image\n    if debug:\n        image0 = np.copy(image)\n    \n    # Always crop 10 pixels for weird border noise/lines\n    image = image[int(h0 * 2e-2):-int(h0 * 2e-2),int(w0 * 2e-2):-int(w0 * 2e-2)]\n    \n    # Crop Image\n    if crop_image:\n        image = crop(image, debug=debug)\n        \n    # Resize\n    if size is not None:\n        # Pad black pixels to make square image\n        h, w = image.shape\n        if (h / w) > TARGET_HEIGHT_WIDTH_RATIO:\n            pad = int(h / TARGET_HEIGHT_WIDTH_RATIO - w)\n            image = np.pad(image, [[0,0], [0, pad]])\n            h, w = image.shape\n        else:\n            pad = int(0.50 * (w * TARGET_HEIGHT_WIDTH_RATIO - h))\n            image = np.pad(image, [[pad, pad], [0,0]])\n            h, w = image.shape\n        # Resize\n        image = cv2.resize(image, size, interpolation=cv2.INTER_AREA)\n        \n    # Apply CLAHE contrast enhancement\n    if apply_clahe:\n        image = CLAHE.apply(image)\n        \n     # Apply Histogram Equalization\n    if apply_eq_hist:\n        image = cv2.equalizeHist(image)\n    return image\n#     # Save Only\n#     if save:\n#         image_id = file_path.split('/')[-1].split('.')[0]\n#         cv2.imwrite(f'{image_id}.png', image)","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:38.518264Z","iopub.execute_input":"2023-01-24T16:34:38.518859Z","iopub.status.idle":"2023-01-24T16:34:38.534505Z","shell.execute_reply.started":"2023-01-24T16:34:38.518806Z","shell.execute_reply":"2023-01-24T16:34:38.533134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pytorch Dataset Class","metadata":{}},{"cell_type":"code","source":"\"\"\"\nDataset is composed by a csv file (train/test) and a folder with images in dicom format.\n\"\"\"\nfrom path import Path\nimport pydicom\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import Dataset\n\n\ntrain_path = Path('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ntest_path = Path('/kaggle/input/rsna-breast-cancer-detection/test.csv')\ntrain_images_path = Path('/kaggle/input/rsna-breast-cancer-detection/train_images')\ntest_images_path = Path('/kaggle/input/rsna-breast-cancer-detection/test_images')\n\ndef get_train_image_path(patient_id, image_id):\n    full_path = train_images_path / str(patient_id) / (str(image_id) + '.dcm')\n    assert full_path.exists(), f'{full_path}  doesnt exists'\n    return full_path\n    \n\ndef get_test_image_path(patient_id, image_id):\n    full_path = test_images_path / str(patient_id) / (str(image_id) + '.dcm')\n    assert full_path.exists(), f'{full_path}  doesnt exists'\n    return full_path\n\ndef load_image(dcm_path):\n    with pydicom.dcmread(dcm_path) as f:\n        return f\n    \ndef load_train_dcm(train: pd.DataFrame, i: int):\n    assert 0 <= i < len(train), f'index {i} is out of range [0, {len(train)})'\n    return load_image(get_train_image_path(train.loc[i, 'patient_id'], train.loc[i, 'image_id']))\n\ndef load_test_dcm(test: pd.DataFrame, i: int):\n    assert 0 <= i < len(test), f'index {i} is out of range [0, {len(test)})'\n    return load_image(get_test_image_path(test.loc[i, 'patient_id'], test.loc[i, 'image_id']))\n\nclass DcmDataset(Dataset):\n\n    def __init__(self, train: bool = True):\n        self.use_train = train\n        self.df_path = train_path if train else test_path\n        self.df = pd.read_csv(self.df_path)\n        self.df['age'] = self.df['age'].astype(np.float32)\n        self.df['view']= self.df['view'].where((self.df['view'] == 'MLO') | (self.df['view'] == 'CC'), 'MLO')\n        self.load_dcm_func = load_train_dcm if train else load_test_dcm\n        self.laterality_map = {'L': 0, 'R': 1}\n        self.view_map = {'MLO': 0, 'CC': 1}\n        self.mean_age = 58.54557093995208 \n        self.std_age = 10.052351179220505\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i: int):\n        dcm = self.load_dcm_func(self.df, i)\n        result = dict(dcm=dcm)\n        if self.use_train:\n            result['target'] = self.df.loc[i, 'cancer']\n        result['laterality'] = self.laterality_map[self.df.loc[i, 'laterality']]\n        result['view'] = self.view_map[self.df.loc[i, 'view']]\n        result['age'] = (self.df.loc[i, 'age'].item() - self.mean_age) / self.std_age\n        return result\n\n    def make_df_submission(self, y_pred):\n        if self.use_train:\n            print('Warning: this method is meant for test only')\n            return\n        assert y_pred.shape[0] == len(self.df)\n        submission = self.df[['prediction_id']].copy()\n        submission['cancer'] = y_pred\n        return submission\n    \nclass DcmImageDataset(DcmDataset):\n    \"\"\"\n    Dataset that returns a dict with the image rather than the dcm object.\n    The image are scaled in min max scaled in [0, 1]\n    \"\"\"\n\n    def __init__(self, train: bool = True):\n        super().__init__(train)\n        self.mean = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1).float()\n        self.std = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1).float()\n        \n    def __getitem__(self, i):\n        result = super().__getitem__(i)\n        image = torch.from_numpy(process(result['dcm'])).float().unsqueeze(0)\n        image = image.expand(3, *image.shape[1:])\n        \n        image = (image - self.mean) / self.std\n        image = F.max_pool2d(image, 2)\n\n        result['image'] = image\n        del result['dcm']\n        return result\n\n    def show_data(self, i):\n        data = self[i]\n        plt.imshow(data['image'][0], cmap='gray')\n        plt.show()\n        del data['image']\n        print(data)\n        \n","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:39.795715Z","iopub.execute_input":"2023-01-24T16:34:39.796216Z","iopub.status.idle":"2023-01-24T16:34:39.823388Z","shell.execute_reply.started":"2023-01-24T16:34:39.796179Z","shell.execute_reply":"2023-01-24T16:34:39.822304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pytorch Lightning model\n\nInputs: _dcm, patient age, view, laterality_\n\n- ConvNext Large to extract image features\n- 2-Layers MLP for age\n- Embedding for _view_ and _laterality_","metadata":{}},{"cell_type":"code","source":"import pytorch_lightning as pl\nimport timm\nfrom torch import nn\nfrom torch.nn import functional as F\nimport torch\nimport numpy as np\n\n\nclass MikedevModel(pl.LightningModule):\n\n    def __init__(self, hidden_features: int = 1024):\n        super().__init__()\n        \n        self.hidden_features = hidden_features\n        self.pretrained = timm.create_model('convnext_base', pretrained=True)\n        for param in self.pretrained.parameters():\n            param.requires_grad = False\n        self.age_mlp = nn.Sequential(\n            nn.Linear(1, 32),\n            nn.ReLU(),\n            nn.Linear(32, hidden_features // 4),\n            nn.ReLU(),\n        )\n\n        self.image_feature_map = nn.Linear(1024 * (42 // 4) * (24 // 4) , hidden_features)\n        self.laterality_embedding = nn.Embedding(2, hidden_features // 4)\n        self.view_embedding = nn.Embedding(2, hidden_features // 4)\n        self.final_mlp = nn.Sequential(\n            nn.Linear(hidden_features + int(hidden_features * 3 / 4),  hidden_features),\n            nn.ReLU(),\n            nn.Linear(hidden_features, 100),\n            nn.ReLU(),\n            nn.Linear(100, 1),\n        )\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')\n        \n    def parameters(self):\n        # all except pretrained model\n        for n, p in super().named_parameters():\n            if not n.startswith('pretrained'):\n                yield p\n\n    def forward(self, img, age, view, laterality):\n        with torch.no_grad():\n            visual_features = self.pretrained.forward_features(img)\n            visual_features = F.max_pool2d(visual_features, 2)\n            visual_features = visual_features.flatten(start_dim=1)\n        visual_features = self.image_feature_map(visual_features)\n        assert visual_features.shape == (img.shape[0], self.hidden_features), f'Expected shape (batch_size, {self.hidden_features}), got {visual_features.shape}'\n        age_features = self.age_mlp(age)\n        laterality_features = self.laterality_embedding(laterality)\n        view_features = self.view_embedding(view)\n        hidden_features = torch.cat([visual_features, age_features, laterality_features, view_features], dim=1)\n        return self.final_mlp(hidden_features)\n\n    def training_step(self, batch, batch_idx):\n        logits = self.predict_step(batch, batch_idx).squeeze(-1)\n        target = batch['target'].float()\n\n        loss = self.bce(logits, target).mean(dim=0).sum()\n        if self.global_step % 1000 == 0:\n            self.log('train_loss', loss, prog_bar=True)\n            print('train loss', loss, 'on step', self.global_step)\n        return loss\n    \n    def predict_step(self, batch, batch_idx):\n        image = batch['image']\n        age = batch['age'].unsqueeze(-1).float()\n        view = batch['view']\n        laterality = batch['laterality']\n        pred = self(image, age, view, laterality)\n        return pred\n\n\n    def configure_optimizers(self):\n        return torch.optim.Adam(self.parameters(), lr=1e-4)\n","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:41.275295Z","iopub.execute_input":"2023-01-24T16:34:41.276093Z","iopub.status.idle":"2023-01-24T16:34:43.849964Z","shell.execute_reply.started":"2023-01-24T16:34:41.276043Z","shell.execute_reply":"2023-01-24T16:34:43.848537Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path = Path('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ntest_path = Path('/kaggle/input/rsna-breast-cancer-detection/test.csv')\ntrain_images_path = Path('/kaggle/input/rsna-breast-cancer-detection/train_images')\ntest_images_path = Path('/kaggle/input/rsna-breast-cancer-detection/test_images')","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:43.85184Z","iopub.execute_input":"2023-01-24T16:34:43.853026Z","iopub.status.idle":"2023-01-24T16:34:43.858785Z","shell.execute_reply.started":"2023-01-24T16:34:43.852982Z","shell.execute_reply":"2023-01-24T16:34:43.857697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(train_path)\ntest = pd.read_csv(test_path)","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:44.291984Z","iopub.execute_input":"2023-01-24T16:34:44.292418Z","iopub.status.idle":"2023-01-24T16:34:44.455896Z","shell.execute_reply.started":"2023-01-24T16:34:44.29238Z","shell.execute_reply":"2023-01-24T16:34:44.454334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# age = pd.concat([train['age'], test['age']])\n# print('age', age.mean(), '+-', age.std())\n\nmean_age = 58.54557093995208 \nstd_age = 10.052351179220505","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:45.051185Z","iopub.execute_input":"2023-01-24T16:34:45.051705Z","iopub.status.idle":"2023-01-24T16:34:45.058495Z","shell.execute_reply.started":"2023-01-24T16:34:45.051662Z","shell.execute_reply":"2023-01-24T16:34:45.056657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head()","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:45.873208Z","iopub.execute_input":"2023-01-24T16:34:45.873769Z","iopub.status.idle":"2023-01-24T16:34:45.907237Z","shell.execute_reply.started":"2023-01-24T16:34:45.873727Z","shell.execute_reply":"2023-01-24T16:34:45.906018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test.head()","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:46.151002Z","iopub.execute_input":"2023-01-24T16:34:46.151472Z","iopub.status.idle":"2023-01-24T16:34:46.166419Z","shell.execute_reply.started":"2023-01-24T16:34:46.151436Z","shell.execute_reply":"2023-01-24T16:34:46.165129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_image(dcm_path):\n    with pydicom.dcmread(dcm_path) as f:\n        return f\n    \ndef load_train_image(i: int):\n    return load_image(get_train_image(train.loc[i, 'patient_id'], train.loc[i, 'image_id']))","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:46.67226Z","iopub.execute_input":"2023-01-24T16:34:46.67276Z","iopub.status.idle":"2023-01-24T16:34:46.681197Z","shell.execute_reply.started":"2023-01-24T16:34:46.672716Z","shell.execute_reply":"2023-01-24T16:34:46.679717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Instantiation and running","metadata":{}},{"cell_type":"code","source":"train_dataset = DcmImageDataset()\ntest_dataset = DcmImageDataset(train=False)\ny = train_dataset[0]['image']","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:48.131806Z","iopub.execute_input":"2023-01-24T16:34:48.132236Z","iopub.status.idle":"2023-01-24T16:34:50.443879Z","shell.execute_reply.started":"2023-01-24T16:34:48.132202Z","shell.execute_reply":"2023-01-24T16:34:50.442741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y.dtype, y.shape","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:50.445787Z","iopub.execute_input":"2023-01-24T16:34:50.446193Z","iopub.status.idle":"2023-01-24T16:34:50.453709Z","shell.execute_reply.started":"2023-01-24T16:34:50.446158Z","shell.execute_reply":"2023-01-24T16:34:50.452813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32\nepochs = 4","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:53.762874Z","iopub.execute_input":"2023-01-24T16:34:53.763816Z","iopub.status.idle":"2023-01-24T16:34:53.771982Z","shell.execute_reply.started":"2023-01-24T16:34:53.763756Z","shell.execute_reply":"2023-01-24T16:34:53.769001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dl = DataLoader(train_dataset, batch_size=batch_size)\ntest_dl = DataLoader(test_dataset, batch_size=batch_size)","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:34:55.332942Z","iopub.execute_input":"2023-01-24T16:34:55.33428Z","iopub.status.idle":"2023-01-24T16:34:55.340336Z","shell.execute_reply.started":"2023-01-24T16:34:55.334218Z","shell.execute_reply":"2023-01-24T16:34:55.33905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = MikedevModel()","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:35:00.324443Z","iopub.execute_input":"2023-01-24T16:35:00.32493Z","iopub.status.idle":"2023-01-24T16:35:03.949489Z","shell.execute_reply.started":"2023-01-24T16:35:00.324891Z","shell.execute_reply":"2023-01-24T16:35:03.948087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"accelerator = 'gpu' if torch.cuda.is_available() else 'cpu'\ntrainer = pl.Trainer(max_epochs=epochs, accelerator=accelerator)","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:35:05.852433Z","iopub.execute_input":"2023-01-24T16:35:05.853394Z","iopub.status.idle":"2023-01-24T16:35:06.063439Z","shell.execute_reply.started":"2023-01-24T16:35:05.853347Z","shell.execute_reply":"2023-01-24T16:35:06.062021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(model, train_dl)","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:35:30.182314Z","iopub.execute_input":"2023-01-24T16:35:30.183008Z","iopub.status.idle":"2023-01-24T16:42:47.48789Z","shell.execute_reply.started":"2023-01-24T16:35:30.18297Z","shell.execute_reply":"2023-01-24T16:42:47.486471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test Prediction","metadata":{}},{"cell_type":"code","source":"with torch.no_grad():\n    y_pred_logits = trainer.predict(model, test_dl)\n    if len(y_pred_logits) == 1:\n        y_pred_logits = y_pred_logits[0]\n    else:\n        y_pred_logits = torch.cat([y_pred_logits], dim=0)\n    y_pred = y_pred_logits.sigmoid()\n","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:35:09.29119Z","iopub.execute_input":"2023-01-24T16:35:09.291721Z","iopub.status.idle":"2023-01-24T16:35:19.53229Z","shell.execute_reply.started":"2023-01-24T16:35:09.291675Z","shell.execute_reply":"2023-01-24T16:35:19.5313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = test_dataset.make_df_submission(y_pred)","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:35:24.793067Z","iopub.execute_input":"2023-01-24T16:35:24.793632Z","iopub.status.idle":"2023-01-24T16:35:24.807181Z","shell.execute_reply.started":"2023-01-24T16:35:24.793567Z","shell.execute_reply":"2023-01-24T16:35:24.80585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-01-24T16:35:25.20459Z","iopub.execute_input":"2023-01-24T16:35:25.205452Z","iopub.status.idle":"2023-01-24T16:35:25.215458Z","shell.execute_reply.started":"2023-01-24T16:35:25.205404Z","shell.execute_reply":"2023-01-24T16:35:25.214214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}