{"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":"📌 <b>This notebook is based on vslaykovsky's work</b> <a href=\"https://www.kaggle.com/code/vslaykovsky/train-pytorch-effnetv2-baseline-cv-0-49\">Pytorch EfficientNet-v2 single model PL:0.49, ensemble PL:0.47</a>. <b>Thank you so much for publishing the great baseline with us!</b>\n\n📌 <b>What I did is, in addition to augmentation, stacking three monochrome images to make one image with the size of (512,512,3). Since the default EfficientNet accepts three (RGB) channels, I tried cramming more information, neighboring images, into it without changing the NN structure. In other words, a small 2.5D model can be created with the pre-trained weights. </b>\n    \n📌 <b>Stacking three images improved CV score from 0.45 (only w/ augmentation) to 0.42 with single V2M model.</b>\n","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n🦴 1. Imports, constants, dependencies 🦴\n</div>","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"try:\n    import pylibjpeg\nexcept:\n    # The following *.whl files were collected from these pip packages:\n    #!pip install -U \"python-gdcm\" pydicom pylibjpeg    # Required for JPEG decompression. See: https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/341412\n    #!pip install -U torchvision                        # For EfficientNetV2\n\n    # Offline dependencies:\n    !mkdir -p /root/.cache/torch/hub/checkpoints/\n#     !cp ../input/rsna-2022-whl/efficientnet_v2_s-dd5fe13b.pth  /root/.cache/torch/hub/checkpoints/\n    !pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.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    !pip install /kaggle/input/rsna-2022-whl/{torch-1.12.1-cp37-cp37m-manylinux1_x86_64.whl,torchvision-0.13.1-cp37-cp37m-manylinux1_x86_64.whl}","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:24:39.003258Z","iopub.execute_input":"2022-10-27T22:24:39.004064Z","iopub.status.idle":"2022-10-27T22:25:43.826468Z","shell.execute_reply.started":"2022-10-27T22:24:39.003975Z","shell.execute_reply":"2022-10-27T22:25:43.825301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nimport glob\nimport os\nimport re\n\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport pydicom as dicom\nimport torch\nimport torchvision as tv\nfrom sklearn.model_selection import GroupKFold\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torchvision.models.feature_extraction import create_feature_extractor\nfrom tqdm.notebook import tqdm\nfrom PIL import Image\n\nimport wandb\n\nplt.rcParams['figure.figsize'] = (20, 5)\npd.set_option('display.max_rows', 100)\npd.set_option('display.max_columns', 1000)\n\n# Effnet\nWEIGHTS = tv.models.efficientnet.EfficientNet_V2_M_Weights.DEFAULT \nRSNA_2022_PATH = '../input/rsna-2022-cervical-spine-fracture-detection'\nTRAIN_IMAGES_PATH = f'{RSNA_2022_PATH}/train_images'\nTEST_IMAGES_PATH = f'{RSNA_2022_PATH}/test_images'\nEFFNET_MAX_TRAIN_BATCHES = 20000\nEFFNET_MAX_EVAL_BATCHES = 200\nONE_CYCLE_MAX_LR = 0.0001\nONE_CYCLE_PCT_START = 0.3\nSAVE_CHECKPOINT_EVERY_STEP = 1000\nEFFNET_CHECKPOINTS_PATH = \"/kaggle/working\" \nFRAC_LOSS_WEIGHT = 2.\nN_FOLDS = 5\nMETADATA_PATH = '../input/rsna-2022-spine-fracture-detection-metadata'\n\nPREDICT_MAX_BATCHES = 1e9\n\n# Common\ntry:\n    from kaggle_secrets import UserSecretsClient\n    IS_KAGGLE = True\nexcept:\n    IS_KAGGLE = False\n\nos.environ[\"WANDB_MODE\"] = \"online\"\nif os.environ[\"WANDB_MODE\"] == \"online\":\n    if IS_KAGGLE:\n        os.environ['WANDB_API_KEY'] = UserSecretsClient().get_secret(\"WANDB_API_KEY\")\n\nif not IS_KAGGLE:\n    print('Running locally')\n    RSNA_2022_PATH = '/mnt/rsna2022'\n    TRAIN_IMAGES_PATH = '/mnt/rsna2022/train_images'\n    TEST_IMAGES_PATH = '/mnt/rsna2022/test_images'\n    METADATA_PATH = '/home/vslaykovsky/Downloads/'\n    EFFNET_CHECKPOINTS_PATH = 'frac_checkpoints'\n    os.environ['WANDB_API_KEY'] = 'yourkeyhere'\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nif DEVICE == 'cuda':\n    BATCH_SIZE = 8\nelse:\n    BATCH_SIZE = 2\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:43.828645Z","iopub.execute_input":"2022-10-27T22:25:43.829032Z","iopub.status.idle":"2022-10-27T22:25:46.732972Z","shell.execute_reply.started":"2022-10-27T22:25:43.828991Z","shell.execute_reply":"2022-10-27T22:25:46.731829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n    🦴 2. Loading train/eval/test dataframes 🦴\n</div>","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"markdown","source":"### Train data\n\n1. Loading data from competition dataset folder `../input/rsna-2022-cervical-spine-fracture-detection/train.csv`\n2. Joining data with slice information from metadata dataset `../input/rsna-2022-spine-fracture-detection-metadata/meta_train_with_vertebrae.csv`\n3. Adding `Splits` column to facilitate train/eval splits.","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"df_train = pd.read_csv(f'{RSNA_2022_PATH}/train.csv')\ndf_train.sample(2)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:46.735032Z","iopub.execute_input":"2022-10-27T22:25:46.735955Z","iopub.status.idle":"2022-10-27T22:25:46.790604Z","shell.execute_reply.started":"2022-10-27T22:25:46.735911Z","shell.execute_reply":"2022-10-27T22:25:46.789552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# rsna-2022-spine-fracture-detection-metadata contains inference of C1-C7 vertebrae for all training sample (95% accuracy)\ndf_train_slices = pd.read_csv(f'{METADATA_PATH}/train_segmented.csv')\nc1c7 = [f'C{i}' for i in range(1, 8)]\ndf_train_slices[c1c7] = (df_train_slices[c1c7] > 0.5).astype(int)\nprint(df_train_slices.sample(5)[['StudyInstanceUID', 'C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7']].to_markdown())","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:46.794034Z","iopub.execute_input":"2022-10-27T22:25:46.794405Z","iopub.status.idle":"2022-10-27T22:25:50.974594Z","shell.execute_reply.started":"2022-10-27T22:25:46.794377Z","shell.execute_reply":"2022-10-27T22:25:50.973439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = df_train_slices.set_index('StudyInstanceUID').join(df_train.set_index('StudyInstanceUID'),\n                                                              rsuffix='_fracture').reset_index().copy()\ndf_train = df_train.query('StudyInstanceUID != \"1.2.826.0.1.3680043.20574\"').reset_index(drop=True)\ndf_train.sample(2)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:50.976964Z","iopub.execute_input":"2022-10-27T22:25:50.978277Z","iopub.status.idle":"2022-10-27T22:25:51.740307Z","shell.execute_reply.started":"2022-10-27T22:25:50.978231Z","shell.execute_reply":"2022-10-27T22:25:51.739409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"split = GroupKFold(N_FOLDS)\nfor k, (_, test_idx) in enumerate(split.split(df_train, groups=df_train.StudyInstanceUID)):\n    df_train.loc[test_idx, 'split'] = k\ndf_train.sample(2)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:51.744202Z","iopub.execute_input":"2022-10-27T22:25:51.746646Z","iopub.status.idle":"2022-10-27T22:25:52.368899Z","shell.execute_reply.started":"2022-10-27T22:25:51.746609Z","shell.execute_reply":"2022-10-27T22:25:52.367954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Test data\n\n1. Loading data from competition dataset folder `../input/rsna-2022-cervical-spine-fracture-detection/test.csv`\n2. Joining data with slice information collected from test image folders `../input/rsna-2022-cervical-spine-fracture-detection/test_images/*/*`","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"df_test = pd.read_csv(f'{RSNA_2022_PATH}/test.csv')\n\nif df_test.iloc[0].row_id == '1.2.826.0.1.3680043.10197_C1':\n    # test_images and test.csv are inconsistent in the dev dataset, fixing labels for the dev run.\n    df_test = pd.DataFrame({\n        \"row_id\": ['1.2.826.0.1.3680043.22327_C1', '1.2.826.0.1.3680043.25399_C1', '1.2.826.0.1.3680043.5876_C1'],\n        \"StudyInstanceUID\": ['1.2.826.0.1.3680043.22327', '1.2.826.0.1.3680043.25399', '1.2.826.0.1.3680043.5876'],\n        \"prediction_type\": [\"C1\", \"C1\", \"patient_overall\"]}\n    )\n\ndf_test","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:52.373262Z","iopub.execute_input":"2022-10-27T22:25:52.375567Z","iopub.status.idle":"2022-10-27T22:25:52.398148Z","shell.execute_reply.started":"2022-10-27T22:25:52.375528Z","shell.execute_reply":"2022-10-27T22:25:52.397283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_slices = glob.glob(f'{TEST_IMAGES_PATH}/*/*')\ntest_slices = [re.findall(f'{TEST_IMAGES_PATH}/(.*)/(.*).dcm', s)[0] for s in test_slices]\ndf_test_slices = pd.DataFrame(data=test_slices, columns=['StudyInstanceUID', 'Slice'])\ndf_test_slices.sample(2)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:52.402121Z","iopub.execute_input":"2022-10-27T22:25:52.404311Z","iopub.status.idle":"2022-10-27T22:25:52.512036Z","shell.execute_reply.started":"2022-10-27T22:25:52.404276Z","shell.execute_reply":"2022-10-27T22:25:52.511055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test = df_test.set_index('StudyInstanceUID').join(df_test_slices.set_index('StudyInstanceUID')).reset_index()\ndf_test.sample(2)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:52.513425Z","iopub.execute_input":"2022-10-27T22:25:52.513733Z","iopub.status.idle":"2022-10-27T22:25:52.534378Z","shell.execute_reply.started":"2022-10-27T22:25:52.513708Z","shell.execute_reply":"2022-10-27T22:25:52.533551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n    🦴 3. Dataset class 🦴\n</div>\n\n`EffnetDataSet` class returns images of individual slices. It uses a dataframe parameter `df` as a source of slices metadata to locate and load images from `path` folder. It accepts transforms parameter which we set to `WEIGHTS.transforms()`. This is a set of transforms used to pre-train the model on ImageNet dataset.","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"def load_dicom(path):\n    \"\"\"\n    This supports loading both regular and compressed JPEG images. \n    See the first sell with `pip install` commands for the necessary dependencies\n    \"\"\"\n    img = dicom.dcmread(path)\n    img.PhotometricInterpretation = 'MONOCHROME2'\n    data = img.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    #return cv2.cvtColor(data, cv2.COLOR_GRAY2RGB), img\n    # Return a gray image.\n    return data, img\n\n\nim, meta = load_dicom(\n    f'{TRAIN_IMAGES_PATH}/1.2.826.0.1.3680043.10001/1.dcm')\nplt.figure()\nplt.imshow(im)\nplt.title('regular image')\n\nim, meta = load_dicom(\n    f'{TRAIN_IMAGES_PATH}/1.2.826.0.1.3680043.10014/1.dcm')\nplt.figure()\nplt.imshow(im)\nplt.title('jpeg')","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:52.53583Z","iopub.execute_input":"2022-10-27T22:25:52.536194Z","iopub.status.idle":"2022-10-27T22:25:53.054096Z","shell.execute_reply.started":"2022-10-27T22:25:52.536161Z","shell.execute_reply":"2022-10-27T22:25:53.052374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EffnetDataSet(torch.utils.data.Dataset):\n    def __init__(self, df, path, transforms=None, aug = False):\n        super().__init__()\n        self.df = df\n        self.path = path\n        self.transforms = transforms\n        self.aug = aug\n        aug_p = 0.5\n        \n        #Augmenation\n        self.transform_aug = tv.transforms.Compose([\n        tv.transforms.RandomHorizontalFlip(p=aug_p),\n        tv.transforms.RandomRotation(degrees=20), \n        tv.transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n        ])\n\n    def __getitem__(self, i):\n        path = os.path.join(self.path, self.df.iloc[i].StudyInstanceUID, f'{self.df.iloc[i].Slice}.dcm')\n        \n        paths = []\n        max_frame = 3 #Find and include the third neighboring frame over the main one if available.\n        \n        j = max_frame\n        while True:\n            try:\n                check_file = os.path.isfile(os.path.join(self.path, self.df.iloc[i].StudyInstanceUID, f'{self.df.iloc[i+j].Slice}.dcm'))\n                if check_file:\n                    paths.append(os.path.join(self.path, self.df.iloc[i].StudyInstanceUID, f'{self.df.iloc[i+j].Slice}.dcm'))\n                    break\n                else:\n                    j -= 1\n            except:\n                j -= 1\n\n        paths.append(os.path.join(self.path, self.df.iloc[i].StudyInstanceUID, f'{self.df.iloc[i].Slice}.dcm'))\n    \n        j = max_frame\n        while True:\n            try:\n                check_file = os.path.isfile(os.path.join(self.path, self.df.iloc[i].StudyInstanceUID, f'{self.df.iloc[i-j].Slice}.dcm'))\n                if check_file:\n                    paths.append(os.path.join(self.path, self.df.iloc[i].StudyInstanceUID, f'{self.df.iloc[i-j].Slice}.dcm'))\n                    break\n                else:\n                    j -= 1\n            except:\n                j -= 1\n                \n        imgs = []\n        for j in range(3):\n            img = load_dicom(paths[j])[0]\n            imgs.append(img)\n            \n        #img=load_dicom(path)[0]\n        img = np.stack([imgs[0],imgs[1],imgs[2]], axis=2)\n        \n        if self.aug:\n            img = Image.fromarray(img)\n            img = self.transform_aug(img)\n            img = np.array(img)\n\n        # Pytorch uses (batch, channel, height, width) order. Converting (height, width, channel) -> (channel, height, width)\n        img = np.transpose(img, (2, 0, 1))\n        if self.transforms is not None:\n            img = self.transforms(torch.as_tensor(img))\n\n        if 'C1_fracture' in self.df:\n            frac_targets = torch.as_tensor(self.df.iloc[i][['C1_fracture', 'C2_fracture', 'C3_fracture', 'C4_fracture',\n                                                            'C5_fracture', 'C6_fracture', 'C7_fracture']].astype(\n                'float32').values)\n            vert_targets = torch.as_tensor(\n                self.df.iloc[i][['C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7']].astype('float32').values)\n            frac_targets = frac_targets * vert_targets  # we only enable targets that are visible on the current slice\n            return img, frac_targets, vert_targets\n        return img\n\n    def __len__(self):\n        return len(self.df)\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:53.055699Z","iopub.execute_input":"2022-10-27T22:25:53.056072Z","iopub.status.idle":"2022-10-27T22:25:53.073375Z","shell.execute_reply.started":"2022-10-27T22:25:53.056034Z","shell.execute_reply":"2022-10-27T22:25:53.072407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train = EffnetDataSet(df_train, TRAIN_IMAGES_PATH, WEIGHTS.transforms(),aug=True)\nX, y_frac, y_vert = ds_train[42]\nprint(X.shape, y_frac.shape, y_vert.shape)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:53.075108Z","iopub.execute_input":"2022-10-27T22:25:53.075524Z","iopub.status.idle":"2022-10-27T22:25:53.16322Z","shell.execute_reply.started":"2022-10-27T22:25:53.075481Z","shell.execute_reply":"2022-10-27T22:25:53.162255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_sample_patient(df, ds):\n    patient = np.random.choice(df.query('patient_overall > 0').StudyInstanceUID)\n    df = df.query('StudyInstanceUID == @patient')\n    display(df)\n\n    frac = np.stack([ds[i][1] for i in df.index])\n    vert = np.stack([ds[i][2] for i in df.index])\n    ax = plt.subplot(1, 2, 1)\n    ax.plot(frac)\n    ax.set_title(f'Vertebrae with fractures by slice (masked by visible vertebrae). uid:{patient}')\n    ax = plt.subplot(1, 2, 2)\n    ax.set_title(f'Visible vertebrae by slice. uid:{patient}')\n    ax.plot(vert)\n\nplot_sample_patient(df_train, ds_train)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:25:53.169341Z","iopub.execute_input":"2022-10-27T22:25:53.169725Z","iopub.status.idle":"2022-10-27T22:26:43.510018Z","shell.execute_reply.started":"2022-10-27T22:25:53.169687Z","shell.execute_reply":"2022-10-27T22:26:43.50904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Only X values returned by the test dataset\nds_test = EffnetDataSet(df_test, TEST_IMAGES_PATH, WEIGHTS.transforms(),aug=True)\nX = ds_test[42]\nX.shape","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:26:43.511603Z","iopub.execute_input":"2022-10-27T22:26:43.512651Z","iopub.status.idle":"2022-10-27T22:26:43.578417Z","shell.execute_reply.started":"2022-10-27T22:26:43.51261Z","shell.execute_reply":"2022-10-27T22:26:43.577287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n    🦴 4. Model 🦴\n</div>\n\n\nIn Pytorch we use create_feature_extractor to access feature layers of pre-existing models. Final flat layer of `efficientnet_v2_s` model is called `flatten`. We'll build our classification layer on top of it. ","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"class EffnetModel(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n        effnet = tv.models.efficientnet_v2_m(weights=WEIGHTS)\n        self.model = create_feature_extractor(effnet, ['flatten'])\n        self.nn_fracture = torch.nn.Sequential(\n            torch.nn.Linear(1280, 7),\n        )\n        self.nn_vertebrae = torch.nn.Sequential(\n            torch.nn.Linear(1280, 7),\n        )\n\n    def forward(self, x):\n        # returns logits\n        x = self.model(x)['flatten']\n        return self.nn_fracture(x), self.nn_vertebrae(x)\n\n    def predict(self, x):\n        frac, vert = self.forward(x)\n        return torch.sigmoid(frac), torch.sigmoid(vert)\n\nmodel = EffnetModel()\nmodel.predict(torch.randn(1, 3, 512, 512))\ndel model","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:26:43.579977Z","iopub.execute_input":"2022-10-27T22:26:43.580451Z","iopub.status.idle":"2022-10-27T22:26:58.189044Z","shell.execute_reply.started":"2022-10-27T22:26:43.580412Z","shell.execute_reply":"2022-10-27T22:26:58.187937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n    🦴 5.1 Train: loss function 🦴\n</div>\n\nWe use weighted loss here. See definition here: https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/340392\nWeighted loss helps us to optimize the same target that is used in the final scoring.\n\nAuxiliary vertebrae detection loss is added in the training/evaluation loop to improve model's performance.","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"def weighted_loss(y_pred_logit, y, reduction='mean', verbose=False):\n    \"\"\"\n    Weighted loss\n    We reuse torch.nn.functional.binary_cross_entropy_with_logits here. pos_weight and weights combined give us necessary coefficients described in https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/340392\n\n    See also this explanation: https://www.kaggle.com/code/samuelcortinhas/rsna-fracture-detection-in-depth-eda/notebook\n    \"\"\"\n\n    neg_weights = (torch.tensor([7., 1, 1, 1, 1, 1, 1, 1]) if y_pred_logit.shape[-1] == 8 else torch.ones(y_pred_logit.shape[-1])).to(DEVICE)\n    pos_weights = (torch.tensor([14., 2, 2, 2, 2, 2, 2, 2]) if y_pred_logit.shape[-1] == 8 else torch.ones(y_pred_logit.shape[-1]) * 2.).to(DEVICE)\n\n    loss = torch.nn.functional.binary_cross_entropy_with_logits(\n        y_pred_logit,\n        y,\n        reduction='none',\n    )\n\n    if verbose:\n        print('loss', loss)\n\n    pos_weights = y * pos_weights.unsqueeze(0)\n    neg_weights = (1 - y) * neg_weights.unsqueeze(0)\n    all_weights = pos_weights + neg_weights\n\n    if verbose:\n        print('all weights', all_weights)\n\n    loss *= all_weights\n    if verbose:\n        print('weighted loss', loss)\n\n    norm = torch.sum(all_weights, dim=1).unsqueeze(1)\n    if verbose:\n        print('normalization factors', norm)\n\n    loss /= norm\n    if verbose:\n        print('normalized loss', loss)\n\n    loss = torch.sum(loss, dim=1)\n    if verbose:\n        print('summed up over patient_overall-C1-C7 loss', loss)\n\n    if reduction == 'mean':\n        return torch.mean(loss)\n    return loss","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:26:58.190539Z","iopub.execute_input":"2022-10-27T22:26:58.191136Z","iopub.status.idle":"2022-10-27T22:26:58.201673Z","shell.execute_reply.started":"2022-10-27T22:26:58.191079Z","shell.execute_reply":"2022-10-27T22:26:58.200371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Quick test of  patient_overall + C1-C7 loss\nweighted_loss(\n    torch.logit(torch.tensor([\n        [0.1, 0.9, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1],\n        [0.1, 0.9, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1]\n    ])).to(DEVICE),\n    torch.tensor([\n        [1., 1., 0., 0., 0., 0., 0., 0.],\n        [0., 0, 0., 0., 0., 0., 0., 0.]\n    ]).to(DEVICE),\n    reduction=None,\n    verbose=True\n)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:26:58.203109Z","iopub.execute_input":"2022-10-27T22:26:58.203584Z","iopub.status.idle":"2022-10-27T22:26:59.576141Z","shell.execute_reply.started":"2022-10-27T22:26:58.203539Z","shell.execute_reply":"2022-10-27T22:26:59.575049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Quick test of C1-C7 loss\nweighted_loss(\n    torch.logit(torch.tensor([\n        [0.9, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1],\n        [0.9, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1]\n    ])).to(DEVICE),\n    torch.tensor([\n        [1., 0., 0., 0., 0., 0., 0.],\n        [0, 0., 0., 0., 0., 0., 0.]\n    ]).to(DEVICE),\n    reduction=None,\n    verbose=True\n)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:26:59.577602Z","iopub.execute_input":"2022-10-27T22:26:59.578249Z","iopub.status.idle":"2022-10-27T22:26:59.596028Z","shell.execute_reply.started":"2022-10-27T22:26:59.578211Z","shell.execute_reply":"2022-10-27T22:26:59.595015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n    🦴 5.2 Train: training/evaluation loop 🦴\n</div>","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"def filter_nones(b):\n    return torch.utils.data.default_collate([v for v in b if v is not None])","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:26:59.597679Z","iopub.execute_input":"2022-10-27T22:26:59.59933Z","iopub.status.idle":"2022-10-27T22:26:59.604309Z","shell.execute_reply.started":"2022-10-27T22:26:59.599293Z","shell.execute_reply":"2022-10-27T22:26:59.603419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_model(name, model):\n    torch.save(model.state_dict(), f'{name}.tph')","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:26:59.605644Z","iopub.execute_input":"2022-10-27T22:26:59.606024Z","iopub.status.idle":"2022-10-27T22:26:59.6142Z","shell.execute_reply.started":"2022-10-27T22:26:59.605979Z","shell.execute_reply":"2022-10-27T22:26:59.613293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_model(model, name, path='.'):\n    data = torch.load(os.path.join(path, f'{name}.tph'), map_location=DEVICE)\n    model.load_state_dict(data)\n    return model\n\n\n# quick test\nmodel = torch.nn.Linear(2, 1)\nsave_model('testmodel', model)\n\nmodel1 = load_model(torch.nn.Linear(2, 1), 'testmodel')\nassert torch.all(\n    next(iter(model1.parameters())) == next(iter(model.parameters()))\n).item(), \"Loading/saving is inconsistent!\"","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:26:59.615618Z","iopub.execute_input":"2022-10-27T22:26:59.616188Z","iopub.status.idle":"2022-10-27T22:26:59.626698Z","shell.execute_reply.started":"2022-10-27T22:26:59.616153Z","shell.execute_reply":"2022-10-27T22:26:59.62572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def evaluate_effnet(model: EffnetModel, ds, max_batches=PREDICT_MAX_BATCHES, shuffle=False):\n    torch.manual_seed(42)\n    model = model.to(DEVICE)\n    dl_test = torch.utils.data.DataLoader(ds, batch_size=BATCH_SIZE, shuffle=shuffle, num_workers=os.cpu_count(),\n                                          collate_fn=filter_nones)\n    pred_frac = []\n    pred_vert = []\n    with torch.no_grad():\n        model.eval()\n        frac_losses = []\n        vert_losses = []\n        with tqdm(dl_test, desc='Eval', miniters=10) as progress:\n            for i, (X, y_frac, y_vert) in enumerate(progress):\n                with autocast():\n                    y_frac_pred, y_vert_pred = model.forward(X.to(DEVICE))\n                    frac_loss = weighted_loss(y_frac_pred, y_frac.to(DEVICE)).item()\n                    vert_loss = torch.nn.functional.binary_cross_entropy_with_logits(y_vert_pred, y_vert.to(DEVICE)).item()\n                    pred_frac.append(torch.sigmoid(y_frac_pred))\n                    pred_vert.append(torch.sigmoid(y_vert_pred))\n                    frac_losses.append(frac_loss)\n                    vert_losses.append(vert_loss)\n\n                if i >= max_batches:\n                    break\n        return np.mean(frac_losses), np.mean(vert_losses), torch.concat(pred_frac).cpu().numpy(), torch.concat(pred_vert).cpu().numpy()\n\n# quick test\nm = EffnetModel()\nfrac_loss, vert_loss, pred1, pred2 = evaluate_effnet(m, ds_train, max_batches=2)\nfrac_loss, vert_loss, pred1.shape, pred2.shape","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:26:59.628205Z","iopub.execute_input":"2022-10-27T22:26:59.62863Z","iopub.status.idle":"2022-10-27T22:27:03.426184Z","shell.execute_reply.started":"2022-10-27T22:26:59.628595Z","shell.execute_reply":"2022-10-27T22:27:03.425143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def gc_collect():\n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:27:03.427856Z","iopub.execute_input":"2022-10-27T22:27:03.428416Z","iopub.status.idle":"2022-10-27T22:27:03.434771Z","shell.execute_reply.started":"2022-10-27T22:27:03.428373Z","shell.execute_reply":"2022-10-27T22:27:03.433646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%wandb\n# inline wandb diagrams!\n\ndef train_effnet(ds_train, ds_eval, logger, name):\n    torch.manual_seed(42)\n    dl_train = torch.utils.data.DataLoader(ds_train, batch_size=BATCH_SIZE, shuffle=True, num_workers=os.cpu_count(),\n                                           collate_fn=filter_nones)\n\n    model = EffnetModel().to(DEVICE)\n    optim = torch.optim.Adam(model.parameters())\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(optim, max_lr=ONE_CYCLE_MAX_LR, epochs=1,\n                                                    steps_per_epoch=min(EFFNET_MAX_TRAIN_BATCHES, len(dl_train)),\n                                                    pct_start=ONE_CYCLE_PCT_START)\n\n    model.train()\n    scaler = GradScaler()\n    with tqdm(dl_train, desc='Train', miniters=10) as progress:\n        for batch_idx, (X, y_frac, y_vert) in enumerate(progress):\n\n            if ds_eval is not None and batch_idx % SAVE_CHECKPOINT_EVERY_STEP == 0 and EFFNET_MAX_EVAL_BATCHES > 0:\n                frac_loss, vert_loss = evaluate_effnet(\n                    model, ds_eval, max_batches=EFFNET_MAX_EVAL_BATCHES, shuffle=True)[:2]\n                model.train()\n                logger.log(\n                    {'eval_frac_loss': frac_loss, 'eval_vert_loss': vert_loss, 'eval_loss': frac_loss + vert_loss})\n                if batch_idx > 0:  # don't save untrained model\n                    save_model(name, model)\n\n            if batch_idx >= EFFNET_MAX_TRAIN_BATCHES:\n                break\n\n            optim.zero_grad()\n            # Using mixed precision training\n            with autocast():\n                y_frac_pred, y_vert_pred = model.forward(X.to(DEVICE))\n                frac_loss = weighted_loss(y_frac_pred, y_frac.to(DEVICE))\n                vert_loss = torch.nn.functional.binary_cross_entropy_with_logits(y_vert_pred, y_vert.to(DEVICE))\n                loss = FRAC_LOSS_WEIGHT * frac_loss + vert_loss\n\n                if np.isinf(loss.item()) or np.isnan(loss.item()):\n                    print(f'Bad loss, skipping the batch {batch_idx}')\n                    del loss, frac_loss, vert_loss, y_frac_pred, y_vert_pred\n                    gc_collect()\n                    continue\n\n            # scaler is needed to prevent \"gradient underflow\"\n            scaler.scale(loss).backward()\n            scaler.step(optim)\n            scaler.update()\n            scheduler.step()\n\n            progress.set_description(f'Train loss: {loss.item() :.02f}')\n            logger.log({'loss': (loss.item()), 'frac_loss': frac_loss.item(), 'vert_loss': vert_loss.item(),\n                        'lr': scheduler.get_last_lr()[0]})\n\n    save_model(name, model)\n    return model\n\n\n# N-fold models. Can be used to estimate accurate CV score and in ensembled submissions.\neffnet_models = []\nfor fold in range(N_FOLDS):\n    if os.path.exists(os.path.join(EFFNET_CHECKPOINTS_PATH, f'effnetv2-f{fold}.tph')):\n        print(f'Found cached version of effnetv2-f{fold}')\n        effnet_models.append(load_model(EffnetModel(), f'effnetv2-f{fold}', EFFNET_CHECKPOINTS_PATH))\n    else:\n        with wandb.init(project='RSNA-2022', name=f'EffNet-v2-fold{fold}') as run:\n            gc_collect()\n            ds_train = EffnetDataSet(df_train.query('split != @fold'), TRAIN_IMAGES_PATH, WEIGHTS.transforms(),aug=True)\n            ds_eval = EffnetDataSet(df_train.query('split == @fold'), TRAIN_IMAGES_PATH, WEIGHTS.transforms(), aug=True)\n            effnet_models.append(train_effnet(ds_train, ds_eval, run, f'effnetv2-f{fold}'))\n    break #Train only one model here.\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-27T22:27:03.436535Z","iopub.execute_input":"2022-10-27T22:27:03.43695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<img src=\"https://images2.imgbox.com/29/19/ncuwno2X_o.png\" alt=\"image host\"/>","metadata":{}},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\" style=\"font-size:25px\">\n    🦴 6. Evaluation 🦴\n</div>\n\nWe cross-validate our final model here using 5 folds.\n1. We generate prediction for every holdout set for every fold.\n2. Predictions are aggregated using the non-parametric model.\n3. Final results are produced using the `weighted_loss`","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"effnet_models = []\neffnet_models.append(load_model(EffnetModel(), 'effnetv2-f0', EFFNET_CHECKPOINTS_PATH))","metadata":{"pycharm":{"name":"#%%\n"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def gen_effnet_predictions(effnet_models, df_train):\n    if os.path.exists(os.path.join(EFFNET_CHECKPOINTS_PATH, 'train_predictions.csv')):\n        print('Found cached version of train_predictions.csv')\n        df_train_predictions = pd.read_csv(os.path.join(EFFNET_CHECKPOINTS_PATH, 'train_predictions.csv'))\n    else:\n        df_train_predictions = []\n        with tqdm(enumerate(effnet_models), total=len(effnet_models), desc='Folds') as progress:\n            for fold, effnet_model in progress:\n                try:\n                    ds_eval = EffnetDataSet(df_train.query('split == @fold'), TRAIN_IMAGES_PATH, WEIGHTS.transforms(), aug = False)\n                    frac_loss, vert_loss, effnet_pred_frac, effnet_pred_vert = evaluate_effnet(effnet_model, ds_eval, PREDICT_MAX_BATCHES)\n                    progress.set_description(f'Fold score:{frac_loss:.02f}')\n                    df_effnet_pred = pd.DataFrame(data=np.concatenate([effnet_pred_frac, effnet_pred_vert], axis=1),\n                                                  columns=[f'C{i}_effnet_frac' for i in range(1, 8)] +\n                                                          [f'C{i}_effnet_vert' for i in range(1, 8)])\n\n                    df = pd.concat(\n                        [df_train.query('split == @fold').head(len(df_effnet_pred)).reset_index(drop=True), df_effnet_pred],\n                        axis=1\n                    ).sort_values(['StudyInstanceUID', 'Slice'])\n                    df_train_predictions.append(df)\n                except Exception as e:\n                    print(\"An error occurred at fold\", fold)\n                    print(e)\n        df_train_predictions = pd.concat(df_train_predictions)\n    return df_train_predictions","metadata":{"pycharm":{"name":"#%%\n"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_pred = gen_effnet_predictions(effnet_models, df_train)\ndf_pred.to_csv('train_predictions.csv', index=False)\ndf_pred","metadata":{"pycharm":{"name":"#%%\n","is_executing":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_sample_patient(df_pred):\n    patient = np.random.choice(df_pred.StudyInstanceUID)\n    df = df_pred.query('StudyInstanceUID == @patient').reset_index()\n\n    plt.subplot(1, 3, 1).plot((df[[f'C{i}_fracture' for i in range(1, 8)]].values * df[[f'C{i}' for i in range(1, 8)]].values))\n    f'Patient {patient}, fractures'\n\n    df[[f'C{i}_effnet_frac' for i in range(1, 8)]].plot(\n        title=f'Patient {patient}, fracture prediction',\n        ax=(plt.subplot(1, 3, 2)))\n\n    df[[f'C{i}_effnet_vert' for i in range(1, 8)]].plot(\n        title=f'Patient {patient}, vertebrae prediction',\n        ax=plt.subplot(1, 3, 3)\n    )\n\nplot_sample_patient(df_pred)","metadata":{"pycharm":{"name":"#%%\n"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_sample_patient(df_pred)","metadata":{"pycharm":{"name":"#%%\n"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_sample_patient(df_pred)","metadata":{"pycharm":{"name":"#%%\n"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_cols = ['patient_overall'] + [f'C{i}_fracture' for i in range(1, 8)]\nfrac_cols = [f'C{i}_effnet_frac' for i in range(1, 8)]\nvert_cols = [f'C{i}_effnet_vert' for i in range(1, 8)]\n\n\ndef patient_prediction(df):\n    c1c7 = np.average(df[frac_cols].values, axis=0, weights=df[vert_cols].values)\n    pred_patient_overall = 1 - np.prod(1 - c1c7)\n    return np.concatenate([[pred_patient_overall], c1c7])\n\ndf_patient_pred = df_pred.groupby('StudyInstanceUID').apply(lambda df: patient_prediction(df)).to_frame('pred').join(df_pred.groupby('StudyInstanceUID')[target_cols].mean())","metadata":{"pycharm":{"name":"#%%\n"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_patient_pred","metadata":{"pycharm":{"name":"#%%\n"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = np.stack(df_patient_pred.pred.values.tolist())\npredictions","metadata":{"pycharm":{"name":"#%%\n"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"targets = df_patient_pred[target_cols].values\ntargets","metadata":{"pycharm":{"name":"#%%\n"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('CV score:', weighted_loss(torch.logit(torch.as_tensor(predictions)).to(DEVICE), torch.as_tensor(targets).to(DEVICE)))","metadata":{"pycharm":{"name":"#%%\n"},"trusted":true},"execution_count":null,"outputs":[]}]}