{"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":"The notebook is a modified version of [this notebook of VLADIMIR SLAYKOVSKIY](https://www.kaggle.com/code/vslaykovsky/infer-pytorch-effnetv2-single-model-lb-0-49)","metadata":{}},{"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-13T22:22:30.704595Z","iopub.execute_input":"2022-10-13T22:22:30.705513Z","iopub.status.idle":"2022-10-13T22:23:42.724863Z","shell.execute_reply.started":"2022-10-13T22:22:30.70541Z","shell.execute_reply":"2022-10-13T22:23:42.723504Z"},"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\n\n# import 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_S_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 = 500\nEFFNET_MAX_EVAL_BATCHES = 50\nONE_CYCLE_MAX_LR = 0.0001\nONE_CYCLE_PCT_START = 0.3\nSAVE_CHECKPOINT_EVERY_STEP = 500\nEFFNET_CHECKPOINTS_PATH = '../input/rsna-2022-base-effnetv2'\nFRAC_LOSS_WEIGHT = 2.\nN_FOLDS = 5\nMETADATA_PATH = '../input/pytorch-effnetv2-vertebrae-detection-acc-0-95'\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\n# os.environ[\"WANDB_MODE\"] = \"online\"\n# if 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 = 16\n    EVAL_BATCH_SIZE = 32\nelse:\n    BATCH_SIZE = 2\n    EVAL_BATCH_SIZE = 16","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-13T23:24:11.636417Z","iopub.execute_input":"2022-10-13T23:24:11.636813Z","iopub.status.idle":"2022-10-13T23:24:11.650905Z","shell.execute_reply.started":"2022-10-13T23:24:11.636778Z","shell.execute_reply":"2022-10-13T23:24:11.649643Z"},"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-13T22:23:44.753217Z","iopub.execute_input":"2022-10-13T22:23:44.753995Z","iopub.status.idle":"2022-10-13T22:23:44.793807Z","shell.execute_reply.started":"2022-10-13T22:23:44.753949Z","shell.execute_reply":"2022-10-13T22:23:44.792612Z"},"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-13T22:23:44.798383Z","iopub.execute_input":"2022-10-13T22:23:44.799255Z","iopub.status.idle":"2022-10-13T22:23:47.321747Z","shell.execute_reply.started":"2022-10-13T22:23:44.79921Z","shell.execute_reply":"2022-10-13T22:23:47.320607Z"},"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-13T22:23:47.323211Z","iopub.execute_input":"2022-10-13T22:23:47.323787Z","iopub.status.idle":"2022-10-13T22:23:47.919096Z","shell.execute_reply.started":"2022-10-13T22:23:47.323746Z","shell.execute_reply":"2022-10-13T22:23:47.917867Z"},"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-13T22:23:47.920946Z","iopub.execute_input":"2022-10-13T22:23:47.921653Z","iopub.status.idle":"2022-10-13T22:23:48.376749Z","shell.execute_reply.started":"2022-10-13T22:23:47.921613Z","shell.execute_reply":"2022-10-13T22:23:48.375752Z"},"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-13T22:23:48.378552Z","iopub.execute_input":"2022-10-13T22:23:48.37934Z","iopub.status.idle":"2022-10-13T22:23:48.39679Z","shell.execute_reply.started":"2022-10-13T22:23:48.37929Z","shell.execute_reply":"2022-10-13T22:23:48.39582Z"},"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-13T22:23:48.398069Z","iopub.execute_input":"2022-10-13T22:23:48.398507Z","iopub.status.idle":"2022-10-13T22:23:48.497982Z","shell.execute_reply.started":"2022-10-13T22:23:48.398462Z","shell.execute_reply":"2022-10-13T22:23:48.496973Z"},"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-13T22:23:48.499372Z","iopub.execute_input":"2022-10-13T22:23:48.49969Z","iopub.status.idle":"2022-10-13T22:23:48.517801Z","shell.execute_reply.started":"2022-10-13T22:23:48.499664Z","shell.execute_reply":"2022-10-13T22:23:48.516801Z"},"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 = 'YBR_FULL'\n    data = img.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return cv2.cvtColor(data, cv2.COLOR_GRAY2RGB), 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-13T22:23:48.522803Z","iopub.execute_input":"2022-10-13T22:23:48.523121Z","iopub.status.idle":"2022-10-13T22:23:49.161303Z","shell.execute_reply.started":"2022-10-13T22:23:48.523085Z","shell.execute_reply":"2022-10-13T22:23:49.160325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EffnetDataSet(torch.utils.data.Dataset):\n    def __init__(self, df, path, transforms=None):\n        super().__init__()\n        self.df = df\n        self.path = path\n        self.transforms = transforms\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        try:\n            img = load_dicom(path)[0]\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        except Exception as ex:\n            print(ex)\n            return None\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-13T22:23:49.165719Z","iopub.execute_input":"2022-10-13T22:23:49.168053Z","iopub.status.idle":"2022-10-13T22:23:49.182711Z","shell.execute_reply.started":"2022-10-13T22:23:49.168015Z","shell.execute_reply":"2022-10-13T22:23:49.181754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train = EffnetDataSet(df_train, TRAIN_IMAGES_PATH, WEIGHTS.transforms())\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-13T22:23:49.187487Z","iopub.execute_input":"2022-10-13T22:23:49.190182Z","iopub.status.idle":"2022-10-13T22:23:49.239952Z","shell.execute_reply.started":"2022-10-13T22:23:49.19014Z","shell.execute_reply":"2022-10-13T22:23:49.238986Z"},"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-13T22:23:49.244066Z","iopub.execute_input":"2022-10-13T22:23:49.246165Z","iopub.status.idle":"2022-10-13T22:24:04.134153Z","shell.execute_reply.started":"2022-10-13T22:23:49.246127Z","shell.execute_reply":"2022-10-13T22:24:04.133145Z"},"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())\nX = ds_test[42]\nX.shape","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-13T22:24:04.13838Z","iopub.execute_input":"2022-10-13T22:24:04.140529Z","iopub.status.idle":"2022-10-13T22:24:04.187389Z","shell.execute_reply.started":"2022-10-13T22:24:04.140489Z","shell.execute_reply":"2022-10-13T22:24:04.186445Z"},"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_s(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-13T22:24:04.190835Z","iopub.execute_input":"2022-10-13T22:24:04.194155Z","iopub.status.idle":"2022-10-13T22:24:06.261182Z","shell.execute_reply.started":"2022-10-13T22:24:04.194118Z","shell.execute_reply":"2022-10-13T22:24:06.260154Z"},"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-13T22:24:06.262707Z","iopub.execute_input":"2022-10-13T22:24:06.263178Z","iopub.status.idle":"2022-10-13T22:24:06.275521Z","shell.execute_reply.started":"2022-10-13T22:24:06.263139Z","shell.execute_reply":"2022-10-13T22:24:06.273441Z"},"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-13T22:24:06.277238Z","iopub.execute_input":"2022-10-13T22:24:06.277594Z","iopub.status.idle":"2022-10-13T22:24:07.982315Z","shell.execute_reply.started":"2022-10-13T22:24:06.277559Z","shell.execute_reply":"2022-10-13T22:24:07.980239Z"},"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-13T22:24:07.990668Z","iopub.execute_input":"2022-10-13T22:24:07.995818Z","iopub.status.idle":"2022-10-13T22:24:08.033237Z","shell.execute_reply.started":"2022-10-13T22:24:07.995759Z","shell.execute_reply":"2022-10-13T22:24:08.031976Z"},"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-13T22:24:08.035432Z","iopub.execute_input":"2022-10-13T22:24:08.03587Z","iopub.status.idle":"2022-10-13T22:24:08.263355Z","shell.execute_reply.started":"2022-10-13T22:24:08.03583Z","shell.execute_reply":"2022-10-13T22:24:08.262285Z"},"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-13T22:24:08.270176Z","iopub.execute_input":"2022-10-13T22:24:08.272143Z","iopub.status.idle":"2022-10-13T22:24:08.327887Z","shell.execute_reply.started":"2022-10-13T22:24:08.272034Z","shell.execute_reply":"2022-10-13T22:24:08.326049Z"},"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-13T22:24:08.334249Z","iopub.execute_input":"2022-10-13T22:24:08.334522Z","iopub.status.idle":"2022-10-13T22:24:08.371076Z","shell.execute_reply.started":"2022-10-13T22:24:08.334497Z","shell.execute_reply":"2022-10-13T22:24:08.369955Z"},"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=EVAL_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-13T23:24:16.786179Z","iopub.execute_input":"2022-10-13T23:24:16.786556Z","iopub.status.idle":"2022-10-13T23:24:21.022387Z","shell.execute_reply.started":"2022-10-13T23:24:16.786522Z","shell.execute_reply":"2022-10-13T23:24:21.021311Z"},"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-13T23:24:21.025636Z","iopub.execute_input":"2022-10-13T23:24:21.026632Z","iopub.status.idle":"2022-10-13T23:24:21.031672Z","shell.execute_reply.started":"2022-10-13T23:24:21.026587Z","shell.execute_reply":"2022-10-13T23:24:21.030668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AverageMeter:\n    \"\"\"Computes and stores the average and current value.\"\"\"\n\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"execution":{"iopub.status.busy":"2022-10-13T23:24:21.033352Z","iopub.execute_input":"2022-10-13T23:24:21.033775Z","iopub.status.idle":"2022-10-13T23:24:21.044909Z","shell.execute_reply.started":"2022-10-13T23:24:21.033739Z","shell.execute_reply":"2022-10-13T23:24:21.044048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%wandb\n# inline wandb diagrams!\n\ndef train_effnet(ds_train, ds_eval, 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    \n    frac_losses = AverageMeter()\n    vert_losses = AverageMeter()\n    losses = AverageMeter()\n    \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\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            \n                print({'eval_frac_loss': frac_loss, 'eval_vert_loss': vert_loss, 'eval_loss': frac_loss + vert_loss})\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                frac_losses.update(frac_loss.item(), X.size(0))\n                vert_losses.update(vert_loss.item(), X.size(0))\n                losses.update(loss.item(), X.size(0))\n                progress.set_postfix(loss=losses.avg, frac_loss=frac_losses.avg, vert_loss=vert_losses.avg)\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#             print({'loss': (loss.item()), 'frac_loss': frac_loss.item(), 'vert_loss': vert_loss.item(), 'lr': scheduler.get_last_lr()[0]})\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(1):\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())\n        ds_eval = EffnetDataSet(df_train.query('split == @fold'), TRAIN_IMAGES_PATH, WEIGHTS.transforms())\n        effnet_models.append(train_effnet(ds_train, ds_eval, f'effnetv2-f{fold}'))\n\n# # \"Main\" model that uses all folds data. Can be used in single-model submissions.\n# if os.path.exists(os.path.join(EFFNET_CHECKPOINTS_PATH, f'effnetv2.tph')):\n#     print(f'Found cached version of effnetv2')\n#     effnet_models.append(load_model(EffnetModel(), f'effnetv2', EFFNET_CHECKPOINTS_PATH))\n# else:\n# #     with wandb.init(project='RSNA-2022', name=f'EffNet-v2') as run:\n#     gc_collect()\n#     ds_train = EffnetDataSet(df_train, TRAIN_IMAGES_PATH, WEIGHTS.transforms())\n#     train_effnet(ds_train, None, f'effnetv2')\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-13T23:32:40.345152Z","iopub.execute_input":"2022-10-13T23:32:40.345519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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 = []\nfor name in tqdm(range(1)):\n    effnet_models.append(load_model(EffnetModel(), f'effnetv2-f{name}', '.'))","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-13T23:30:19.637888Z","iopub.execute_input":"2022-10-13T23:30:19.639063Z","iopub.status.idle":"2022-10-13T23:30:20.658018Z","shell.execute_reply.started":"2022-10-13T23:30:19.639017Z","shell.execute_reply":"2022-10-13T23:30:20.656994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def gen_effnet_predictions(effnet_models, df_train):\n\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            ds_eval = EffnetDataSet(df_train.query('split == @fold')[:1000], TRAIN_IMAGES_PATH, WEIGHTS.transforms())\n\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    df_train_predictions = pd.concat(df_train_predictions)\n    return df_train_predictions","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-13T23:30:25.02532Z","iopub.execute_input":"2022-10-13T23:30:25.025703Z","iopub.status.idle":"2022-10-13T23:30:25.035445Z","shell.execute_reply.started":"2022-10-13T23:30:25.02567Z","shell.execute_reply":"2022-10-13T23:30:25.034474Z"},"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},"execution":{"iopub.status.busy":"2022-10-13T23:30:25.325129Z","iopub.execute_input":"2022-10-13T23:30:25.32549Z","iopub.status.idle":"2022-10-13T23:30:42.572571Z","shell.execute_reply.started":"2022-10-13T23:30:25.325459Z","shell.execute_reply":"2022-10-13T23:30:42.571432Z"},"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"},"execution":{"iopub.status.busy":"2022-10-13T23:30:43.926414Z","iopub.execute_input":"2022-10-13T23:30:43.926807Z","iopub.status.idle":"2022-10-13T23:30:44.490741Z","shell.execute_reply.started":"2022-10-13T23:30:43.926767Z","shell.execute_reply":"2022-10-13T23:30:44.489643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_sample_patient(df_pred)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-13T23:30:44.492691Z","iopub.execute_input":"2022-10-13T23:30:44.493149Z","iopub.status.idle":"2022-10-13T23:30:45.038806Z","shell.execute_reply.started":"2022-10-13T23:30:44.49311Z","shell.execute_reply":"2022-10-13T23:30:45.037872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_sample_patient(df_pred)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-13T23:30:47.834368Z","iopub.execute_input":"2022-10-13T23:30:47.835092Z","iopub.status.idle":"2022-10-13T23:30:48.552811Z","shell.execute_reply.started":"2022-10-13T23:30:47.83505Z","shell.execute_reply":"2022-10-13T23:30:48.55185Z"},"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"},"execution":{"iopub.status.busy":"2022-10-13T23:30:52.385905Z","iopub.execute_input":"2022-10-13T23:30:52.386317Z","iopub.status.idle":"2022-10-13T23:30:52.406098Z","shell.execute_reply.started":"2022-10-13T23:30:52.386283Z","shell.execute_reply":"2022-10-13T23:30:52.405157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_patient_pred","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-13T23:31:06.426515Z","iopub.execute_input":"2022-10-13T23:31:06.426986Z","iopub.status.idle":"2022-10-13T23:31:06.44591Z","shell.execute_reply.started":"2022-10-13T23:31:06.426901Z","shell.execute_reply":"2022-10-13T23:31:06.44498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = np.stack(df_patient_pred.pred.values.tolist())\npredictions","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-13T23:31:08.205744Z","iopub.execute_input":"2022-10-13T23:31:08.206447Z","iopub.status.idle":"2022-10-13T23:31:08.214116Z","shell.execute_reply.started":"2022-10-13T23:31:08.206401Z","shell.execute_reply":"2022-10-13T23:31:08.213122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"targets = df_patient_pred[target_cols].values\ntargets","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2022-10-13T23:31:12.511642Z","iopub.execute_input":"2022-10-13T23:31:12.512044Z","iopub.status.idle":"2022-10-13T23:31:12.519857Z","shell.execute_reply.started":"2022-10-13T23:31:12.512011Z","shell.execute_reply":"2022-10-13T23:31:12.518937Z"},"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"},"execution":{"iopub.status.busy":"2022-10-13T23:31:13.459663Z","iopub.execute_input":"2022-10-13T23:31:13.460585Z","iopub.status.idle":"2022-10-13T23:31:13.470288Z","shell.execute_reply.started":"2022-10-13T23:31:13.460536Z","shell.execute_reply":"2022-10-13T23:31:13.469113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-danger\" style=\"text-align:center; font-size:20px;\">\n    ❤️ Dont forget to ▲upvote▲ if you find this notebook usefull!  ❤️\n</div>","metadata":{"execution":{"iopub.execute_input":"2022-08-20T13:17:18.762083Z","iopub.status.busy":"2022-08-20T13:17:18.761536Z","iopub.status.idle":"2022-08-20T13:17:18.76993Z","shell.execute_reply":"2022-08-20T13:17:18.768312Z","shell.execute_reply.started":"2022-08-20T13:17:18.762038Z"},"pycharm":{"name":"#%% md\n"}}}]}