{"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":"# RSNA Breast Baseline - Inference\n\nThis notebook shows a simple inference pipeline based on the pregenerated datasets.\n\n**Training Datasets :**\n\n- LB 0.09  -> https://www.kaggle.com/datasets/theoviel/rsna-breast-cancer-512-pngs\n- LB 0.06~?  -> https://www.kaggle.com/datasets/theoviel/rsna-breast-cancer-256-pngs\n\n**Changes :**\n- v1 : LB 0.09\n- v2 : Trying PP from https://www.kaggle.com/competitions/rsna-breast-cancer-detection/discussion/369886","metadata":{}},{"cell_type":"code","source":"DEBUG = False","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:48:26.639768Z","iopub.execute_input":"2022-12-01T22:48:26.640394Z","iopub.status.idle":"2022-12-01T22:48:26.667746Z","shell.execute_reply.started":"2022-12-01T22:48:26.640288Z","shell.execute_reply":"2022-12-01T22:48:26.666893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Initialization","metadata":{}},{"cell_type":"code","source":"try:\n    import pylibjpeg\nexcept:\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}","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-01T22:48:26.669557Z","iopub.execute_input":"2022-12-01T22:48:26.669969Z","iopub.status.idle":"2022-12-01T22:48:59.3711Z","shell.execute_reply.started":"2022-12-01T22:48:26.669934Z","shell.execute_reply":"2022-12-01T22:48:59.369968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nimport cv2\nimport glob\nimport gdcm\nimport json\nimport pydicom\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nfrom tqdm.notebook import tqdm\nfrom joblib import Parallel, delayed","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-01T22:48:59.37266Z","iopub.execute_input":"2022-12-01T22:48:59.372971Z","iopub.status.idle":"2022-12-01T22:49:00.254825Z","shell.execute_reply.started":"2022-12-01T22:48:59.37294Z","shell.execute_reply":"2022-12-01T22:49:00.25383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Preparation\n\n- I use the same strategy as in https://www.kaggle.com/code/theoviel/dicom-resized-png-jpg","metadata":{}},{"cell_type":"code","source":"test_images = glob.glob(\"/kaggle/input/rsna-breast-cancer-detection/test_images/*/*.dcm\")\n\nif DEBUG:\n    test_images = glob.glob(\"/kaggle/input/rsna-breast-cancer-detection/train_images/10042/*.dcm\")\n    \nprint(\"Number of images :\", len(test_images))","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:00.257476Z","iopub.execute_input":"2022-12-01T22:49:00.257816Z","iopub.status.idle":"2022-12-01T22:49:00.268291Z","shell.execute_reply.started":"2022-12-01T22:49:00.25778Z","shell.execute_reply":"2022-12-01T22:49:00.267237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SAVE_FOLDER = \"/kaggle/tmp/output/\"\nSIZE = 512\nEXTENSION = \"png\"\n\nos.makedirs(SAVE_FOLDER, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:00.269783Z","iopub.execute_input":"2022-12-01T22:49:00.270205Z","iopub.status.idle":"2022-12-01T22:49:00.275977Z","shell.execute_reply.started":"2022-12-01T22:49:00.270168Z","shell.execute_reply":"2022-12-01T22:49:00.275005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process(f, size=512, save_folder=\"\", extension=\"png\"):\n    patient = f.split('/')[-2]\n    image = f.split('/')[-1][:-4]\n\n    dicom = pydicom.dcmread(f)\n    img = dicom.pixel_array\n\n    img = (img - img.min()) / (img.max() - img.min())\n\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n\n    img = cv2.resize(img, (size, size))\n\n    cv2.imwrite(save_folder + f\"{patient}_{image}.{extension}\", (img * 255).astype(np.uint8))","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:00.277255Z","iopub.execute_input":"2022-12-01T22:49:00.278269Z","iopub.status.idle":"2022-12-01T22:49:00.287037Z","shell.execute_reply.started":"2022-12-01T22:49:00.278233Z","shell.execute_reply":"2022-12-01T22:49:00.286093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_ = Parallel(n_jobs=4)(\n    delayed(process)(uid, size=SIZE, save_folder=SAVE_FOLDER, extension=EXTENSION)\n    for uid in tqdm(test_images)\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:00.289734Z","iopub.execute_input":"2022-12-01T22:49:00.290467Z","iopub.status.idle":"2022-12-01T22:49:05.740661Z","shell.execute_reply.started":"2022-12-01T22:49:00.290422Z","shell.execute_reply":"2022-12-01T22:49:05.739523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training\n","metadata":{}},{"cell_type":"markdown","source":"\n**There are already public codes available, I will not share my training code but here are the settings & logs :**","metadata":{}},{"cell_type":"code","source":"class Config:\n    \"\"\"\n    Parameters used for training\n    \"\"\"\n    # General\n    seed = 42\n    verbose = 1\n    device = \"cuda\"\n    save_weights = True\n\n    # Images\n    size = 256\n    size = 512\n\n    # k-fold\n    k = 4  # Stratified GKF\n\n    # Model\n    name = \"tf_efficientnetv2_s\"\n    pretrained_weights = None\n    num_classes = 1\n    n_channels = 3\n\n    # Training    \n    loss_config = {\n        \"name\": \"bce\",  # dice, ce, bce\n        \"smoothing\": 0.,  # 0.01\n        \"activation\": \"sigmoid\",  # \"sigmoid\", \"softmax\"\n        \"aux_loss_weight\": 0,\n    }\n\n    data_config = {\n        \"batch_size\": 32,\n        \"val_bs\": 32,\n    }\n\n    optimizer_config = {\n        \"name\": \"AdamW\",\n        \"lr\": 3e-4,\n        \"warmup_prop\": 0.1,\n        \"betas\": (0.9, 0.999),\n        \"max_grad_norm\": 10.,\n    }\n\n    epochs = 4\n    use_fp16 = True\n    \n    ## Other stuff\n    # Augmentations : Only HorizontalFlip","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:05.742606Z","iopub.execute_input":"2022-12-01T22:49:05.743181Z","iopub.status.idle":"2022-12-01T22:49:05.761882Z","shell.execute_reply.started":"2022-12-01T22:49:05.743136Z","shell.execute_reply":"2022-12-01T22:49:05.76036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"```\n-------------   Fold 1 / 4  -------------\n\n    -> 41092 training images\n    -> 13614 validation images\n    -> 21459769 trainable parameters\n\nEpoch 01/04 (step 0500) \tlr=2.9e-04 \t t=292s \t loss=0.181\t val_loss=0.095    pf1=0.014    auc=0.552\nEpoch 01/04 (step 1000) \tlr=2.7e-04 \t t=290s \t loss=0.101\t val_loss=0.097    pf1=0.018    auc=0.567\nEpoch 02/04 (step 1500) \tlr=2.4e-04 \t t=293s \t loss=0.103\t val_loss=0.093    pf1=0.016    auc=0.627\nEpoch 02/04 (step 2000) \tlr=2.0e-04 \t t=291s \t loss=0.099\t val_loss=0.093    pf1=0.021    auc=0.645\nEpoch 02/04 (step 2500) \tlr=1.7e-04 \t t=293s \t loss=0.104\t val_loss=0.088    pf1=0.029    auc=0.710\nEpoch 03/04 (step 3000) \tlr=1.4e-04 \t t=296s \t loss=0.094\t val_loss=0.090    pf1=0.032    auc=0.699\nEpoch 03/04 (step 3500) \tlr=1.1e-04 \t t=293s \t loss=0.094\t val_loss=0.086    pf1=0.035    auc=0.732\nEpoch 04/04 (step 4000) \tlr=7.4e-05 \t t=296s \t loss=0.094\t val_loss=0.088    pf1=0.044    auc=0.721\nEpoch 04/04 (step 4500) \tlr=4.2e-05 \t t=296s \t loss=0.085\t val_loss=0.087    pf1=0.065    auc=0.740\nEpoch 04/04 (step 5137) \tlr=1.9e-07 \t t=352s \t loss=0.076\t val_loss=0.089    pf1=0.083    auc=0.755\n```","metadata":{}},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"markdown","source":"### Utils","metadata":{}},{"cell_type":"code","source":"import cv2\nimport torch\nfrom torch.utils.data import Dataset\n\n\nclass BreastDataset(Dataset):\n    \"\"\"\n    Image torch Dataset.\n    \"\"\"\n    def __init__(\n        self,\n        df,\n        transforms=None,\n    ):\n        \"\"\"\n        Constructor\n\n        Args:\n            paths (list): Path to images.\n            transforms (albumentation transforms, optional): Transforms to apply. Defaults to None.\n        \"\"\"\n        self.paths = df['path'].values\n        self.transforms = transforms\n        self.targets = df['cancer'].values\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, idx):\n        \"\"\"\n        Item accessor\n\n        Args:\n            idx (int): Index.\n\n        Returns:\n            np array [H x W x C]: Image.\n            torch tensor [1]: Label.\n            torch tensor [1]: Sample weight.\n        \"\"\"\n        image = cv2.imread(self.paths[idx])\n\n        if self.transforms:\n            image = self.transforms(image=image)[\"image\"]\n\n        y = torch.tensor([self.targets[idx]], dtype=torch.float)\n        w = torch.tensor([1])\n\n        return image, y, w\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-01T22:49:05.76599Z","iopub.execute_input":"2022-12-01T22:49:05.766774Z","iopub.status.idle":"2022-12-01T22:49:07.301691Z","shell.execute_reply.started":"2022-12-01T22:49:05.766736Z","shell.execute_reply":"2022-12-01T22:49:07.300682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport albumentations as albu\nfrom albumentations.pytorch import ToTensorV2\n\nMEAN = np.array([0.66437738, 0.50478148, 0.70114894])\nSTD = np.array([0.15825711, 0.24371008, 0.13832686])\n\ndef get_transfos(augment=True, visualize=False, mean=MEAN, std=STD):\n    \"\"\"\n    Returns transformations.\n\n    Args:\n        augment (bool, optional): Whether to apply augmentations. Defaults to True.\n        visualize (bool, optional): Whether to use transforms for visualization. Defaults to False.\n        mean (np array, optional): Mean for normalization. Defaults to MEAN.\n        std (np array, optional): Standard deviation for normalization. Defaults to STD.\n\n    Returns:\n        albumentation transforms: transforms.\n    \"\"\"\n    return albu.Compose(\n        [\n            albu.Normalize(mean=0, std=1),\n            ToTensorV2(),\n        ],\n        p=1,\n    )\n\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-01T22:49:07.30554Z","iopub.execute_input":"2022-12-01T22:49:07.30636Z","iopub.status.idle":"2022-12-01T22:49:08.439126Z","shell.execute_reply.started":"2022-12-01T22:49:07.306322Z","shell.execute_reply":"2022-12-01T22:49:08.438182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nimport random\nimport numpy as np\n\n\ndef load_model_weights(model, filename, verbose=1, cp_folder=\"\", strict=True):\n    \"\"\"\n    Loads the weights of a PyTorch model. The exception handles cpu/gpu incompatibilities.\n\n    Args:\n        model (torch model): Model to load the weights to.\n        filename (str): Name of the checkpoint.\n        verbose (int, optional): Whether to display infos. Defaults to 1.\n        cp_folder (str, optional): Folder to load from. Defaults to \"\".\n\n    Returns:\n        torch model: Model with loaded weights.\n    \"\"\"\n    state_dict = torch.load(os.path.join(cp_folder, filename), map_location=\"cpu\")\n\n    try:\n        model.load_state_dict(state_dict, strict=strict)\n    except BaseException:\n        try:\n            del state_dict['logits.weight'], state_dict['logits.bias']\n            model.load_state_dict(state_dict, strict=strict)\n        except BaseException:\n            del state_dict['encoder.conv_stem.weight']\n            model.load_state_dict(state_dict, strict=strict)\n\n    if verbose:\n        print(f\"\\n -> Loading encoder weights from {os.path.join(cp_folder,filename)}\\n\")\n\n    return model\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-01T22:49:08.440661Z","iopub.execute_input":"2022-12-01T22:49:08.441023Z","iopub.status.idle":"2022-12-01T22:49:08.451214Z","shell.execute_reply.started":"2022-12-01T22:49:08.440987Z","shell.execute_reply":"2022-12-01T22:49:08.449053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport numpy as np\nfrom torch.utils.data import DataLoader\n\nNUM_WORKERS = 2\n\nFLIPS = [None, [-1]]\n\n\ndef predict(model, dataset, loss_config, batch_size=64, device=\"cuda\"):\n    \"\"\"\n    Torch predict function.\n\n    Args:\n        model (torch model): Model to predict with.\n        dataset (CustomDataset): Dataset to predict on.\n        loss_config (dict): Loss config, used for activation functions.\n        batch_size (int, optional): Batch size. Defaults to 64.\n        device (str, optional): Device for torch. Defaults to \"cuda\".\n\n    Returns:\n        numpy array [len(dataset) x num_classes]: Predictions.\n    \"\"\"\n    model.eval()\n    preds = np.empty((0,  model.num_classes))\n\n    loader = DataLoader(\n        dataset, batch_size=batch_size, shuffle=False, num_workers=2\n    )\n\n    with torch.no_grad():\n        for batch in loader:\n            x = batch[0].to(device)\n\n            # Forward\n            pred, pred_aux = model(x)\n\n            # Get probabilities\n            if loss_config['activation'] == \"sigmoid\":\n                pred = pred.sigmoid()\n            elif loss_config['activation'] == \"softmax\":\n                pred = pred.softmax(-1)\n\n            preds = np.concatenate([preds, pred.cpu().numpy()])\n\n    return preds\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-01T22:49:08.452982Z","iopub.execute_input":"2022-12-01T22:49:08.453338Z","iopub.status.idle":"2022-12-01T22:49:08.464224Z","shell.execute_reply.started":"2022-12-01T22:49:08.453302Z","shell.execute_reply":"2022-12-01T22:49:08.463352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/timm-0-6-9/pytorch-image-models-master')\n\nimport timm\nimport torch\nimport torch.nn as nn\n\n\ndef define_model(\n    name,\n    num_classes=1,\n    num_classes_aux=0,\n    n_channels=1,\n    pretrained_weights=\"\",\n    pretrained=True,\n):\n    \"\"\"\n    Loads a pretrained model & builds the architecture.\n    Supports timm models.\n\n    Args:\n        name (str): Model name\n        num_classes (int, optional): Number of classes. Defaults to 1.\n        num_classes_aux (int, optional): Number of aux classes. Defaults to 0.\n        n_channels (int, optional): Number of image channels. Defaults to 3.\n        pretrained_weights (str, optional): Path to pretrained encoder weights. Defaults to ''.\n        pretrained (bool, optional): Whether to load timm pretrained weights.\n\n    Returns:\n        torch model -- Pretrained model.\n    \"\"\"\n    # Load pretrained model\n    encoder = getattr(timm.models, name)(pretrained=pretrained)\n    encoder.name = name\n\n    # Tile Model\n    model = ClsModel(\n        encoder,\n        num_classes=num_classes,\n        num_classes_aux=num_classes_aux,\n        n_channels=n_channels,\n    )\n\n    if pretrained_weights:\n        model = load_model_weights(model, pretrained_weights, verbose=1, strict=False)\n\n    return model\n\n\nclass ClsModel(nn.Module):\n    \"\"\"\n    Model with an attention mechanism.\n    \"\"\"\n    def __init__(\n        self,\n        encoder,\n        num_classes=1,\n        num_classes_aux=0,\n        n_channels=3,\n    ):\n        \"\"\"\n        Constructor.\n\n        Args:\n            encoder (timm model): Encoder.\n            num_classes (int, optional): Number of classes. Defaults to 1.\n            num_classes_aux (int, optional): Number of aux classes. Defaults to 0.\n            n_channels (int, optional): Number of image channels. Defaults to 3.\n        \"\"\"\n        super().__init__()\n\n        self.encoder = encoder\n        self.nb_ft = encoder.num_features\n\n        self.num_classes = num_classes\n        self.num_classes_aux = num_classes_aux\n        self.n_channels = n_channels\n\n        self.logits = nn.Linear(self.nb_ft, num_classes)\n        if self.num_classes_aux:\n            self.logits_aux = nn.Linear(self.nb_ft, num_classes_aux)\n\n        self._update_num_channels()\n\n    def _update_num_channels(self):\n        if self.n_channels != 3:\n            for n, m in self.encoder.named_modules():\n                if n:\n                    # print(\"Replacing\", n)\n                    old_conv = getattr(self.encoder, n)\n                    new_conv = nn.Conv2d(\n                        self.n_channels,\n                        old_conv.out_channels,\n                        kernel_size=old_conv.kernel_size,\n                        stride=old_conv.stride,\n                        padding=old_conv.padding,\n                        bias=old_conv.bias is not None,\n                    )\n                    setattr(self.encoder, n, new_conv)\n                    break\n\n    def extract_features(self, x):\n        \"\"\"\n        Extract features function.\n\n        Args:\n            x (torch tensor [batch_size x 3 x w x h]): Input batch.\n\n        Returns:\n            torch tensor [batch_size x num_features]: Features.\n        \"\"\"\n        fts = self.encoder.forward_features(x)\n\n        while len(fts.size()) > 2:\n            fts = fts.mean(-1)\n\n        return fts\n\n    def get_logits(self, fts):\n        \"\"\"\n        Computes logits.\n\n        Args:\n            fts (torch tensor [batch_size x num_features]): Features.\n\n        Returns:\n            torch tensor [batch_size x num_classes]: logits.\n            torch tensor [batch_size x num_classes_aux]: logits aux.\n        \"\"\"\n        logits = self.logits(fts)\n\n        if self.num_classes_aux:\n            logits_aux = self.logits_aux(fts)\n        else:\n            logits_aux = torch.zeros((fts.size(0)))\n\n        return logits, logits_aux\n\n    def forward(self, x, return_fts=False):\n        \"\"\"\n        Forward function.\n\n        Args:\n            x (torch tensor [batch_size x n_frames x h x w]): Input batch.\n\n        Returns:\n            torch tensor [batch_size x num_classes]: logits.\n            torch tensor [batch_size x num_classes_aux]: logits aux.\n        \"\"\"\n        fts = self.extract_features(x)\n\n        logits, logits_aux = self.get_logits(fts)\n\n        return logits, logits_aux\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-01T22:49:08.46576Z","iopub.execute_input":"2022-12-01T22:49:08.46616Z","iopub.status.idle":"2022-12-01T22:49:10.256443Z","shell.execute_reply.started":"2022-12-01T22:49:08.466124Z","shell.execute_reply":"2022-12-01T22:49:10.255378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Main","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(\"/kaggle/input/rsna-breast-cancer-detection/test.csv\")\ndf['cancer'] = 0\n\nif DEBUG:\n    df = pd.read_csv(\"/kaggle/input/rsna-breast-cancer-detection/train.csv\")\n    df = df[df['patient_id'] == 10042].reset_index(drop=True)\n    \ndf['path'] = SAVE_FOLDER + df[\"patient_id\"].astype(str) + \"_\" + df[\"image_id\"].astype(str) + \".png\"","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:10.26034Z","iopub.execute_input":"2022-12-01T22:49:10.260891Z","iopub.status.idle":"2022-12-01T22:49:10.28276Z","shell.execute_reply.started":"2022-12-01T22:49:10.26085Z","shell.execute_reply":"2022-12-01T22:49:10.281917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"USE_TTA = False\npredict_fct = predict_tta if USE_TTA else predict","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:10.284005Z","iopub.execute_input":"2022-12-01T22:49:10.284964Z","iopub.status.idle":"2022-12-01T22:49:10.290523Z","shell.execute_reply.started":"2022-12-01T22:49:10.284928Z","shell.execute_reply":"2022-12-01T22:49:10.288478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EXP_FOLDERS = [\n    \"/kaggle/input/rsna-breast-weights-public/\"\n]","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:10.292162Z","iopub.execute_input":"2022-12-01T22:49:10.292538Z","iopub.status.idle":"2022-12-01T22:49:10.299674Z","shell.execute_reply.started":"2022-12-01T22:49:10.292501Z","shell.execute_reply":"2022-12-01T22:49:10.298728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Main loop","metadata":{}},{"cell_type":"code","source":"all_preds = []\nfor exp_folder in EXP_FOLDERS:\n    config = Config\n\n    model = define_model(\n        config.name,\n        num_classes=config.num_classes,\n        num_classes_aux=0,\n        n_channels=3,\n        pretrained=False\n    )\n    model = model.cuda().eval()\n\n    dataset = BreastDataset(\n        df,\n        transforms=get_transfos(augment=False),\n    )\n    \n    weights = sorted(glob.glob(exp_folder + f\"*.pt\"))\n\n    preds = []\n    for fold, weight in enumerate(weights):\n        model = load_model_weights(model, weight, verbose=1)\n\n        pred = predict_fct(model, dataset, config.loss_config, batch_size=64)\n        preds.append(pred)\n\n    preds = np.mean(preds, 0)\n    all_preds.append(preds)","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:10.301441Z","iopub.execute_input":"2022-12-01T22:49:10.301863Z","iopub.status.idle":"2022-12-01T22:49:21.609864Z","shell.execute_reply.started":"2022-12-01T22:49:10.301831Z","shell.execute_reply":"2022-12-01T22:49:21.608646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Blend, PP & Submit","metadata":{}},{"cell_type":"code","source":"THRESHOLD = 0.16\n\npreds = np.mean(all_preds, 0)\npreds = (preds > THRESHOLD).astype(int)\n\ndf[\"cancer\"] = preds","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:21.611878Z","iopub.execute_input":"2022-12-01T22:49:21.612288Z","iopub.status.idle":"2022-12-01T22:49:21.621238Z","shell.execute_reply.started":"2022-12-01T22:49:21.612245Z","shell.execute_reply":"2022-12-01T22:49:21.620291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['prediction_id'] = df['patient_id'].astype(str) + \"_\" + df['laterality']\n\nsub = df[['prediction_id', 'cancer']].groupby(\"prediction_id\").mean().reset_index()\n\nsub.to_csv('/kaggle/working/submission.csv', index=False)\n\nsub.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-01T22:49:21.623452Z","iopub.execute_input":"2022-12-01T22:49:21.625181Z","iopub.status.idle":"2022-12-01T22:49:21.65398Z","shell.execute_reply.started":"2022-12-01T22:49:21.625115Z","shell.execute_reply":"2022-12-01T22:49:21.653131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Done !","metadata":{}}]}