{"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":"<a id=\"section-one\"></a>\n# Training a model:\n","metadata":{"execution":{"iopub.execute_input":"2022-12-02T00:49:16.208915Z","iopub.status.busy":"2022-12-02T00:49:16.208242Z","iopub.status.idle":"2022-12-02T00:49:16.247884Z","shell.execute_reply":"2022-12-02T00:49:16.244677Z","shell.execute_reply.started":"2022-12-02T00:49:16.20876Z"},"papermill":{"duration":0.004797,"end_time":"2022-12-11T23:15:56.132313","exception":false,"start_time":"2022-12-11T23:15:56.127516","status":"completed"},"tags":[],"editable":false}},{"cell_type":"code","source":"# !unzip -q ../input/timm-with-dependencies/timm_all -d timm-with-dependencies\n# !pip install --no-index --find-links timm-with-dependencies timm\n# !pip install /kaggle/input/dicomsdl-offline-installer/dicomsdl-0.109.1-cp37-cp37m-manylinux_2_12_x86_64.manylinux2010_x86_64.whl\nimport os\n\nfrom fastai.vision.learner import *\nfrom fastai.data.all import *\nfrom fastai.vision.all import *\nfrom fastai.metrics import ActivationType, error_rate, accuracy\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom collections import defaultdict\nimport pandas as pd\nimport numpy as np\nfrom pdb import set_trace\n\nfrom matplotlib import pyplot as plt\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport seaborn as sns\nfrom pathlib import Path\nimport glob","metadata":{"papermill":{"duration":72.944243,"end_time":"2022-12-11T23:17:09.081513","exception":false,"start_time":"2022-12-11T23:15:56.13727","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:15:00.546081Z","iopub.execute_input":"2023-04-18T19:15:00.546789Z","iopub.status.idle":"2023-04-18T19:15:03.683145Z","shell.execute_reply.started":"2023-04-18T19:15:00.546748Z","shell.execute_reply":"2023-04-18T19:15:03.681717Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### The competition metric","metadata":{"editable":false}},{"cell_type":"markdown","source":"It is always a good idea to understand exactly the type of predictions our model will be required to deliver.\n\nThe metric that the organizers opted for here is the probabilistic F1 score:\n\n$pF_1 = 2\\frac{pPrecision \\cdot pRecall}{pPrecision+pRecall}$\n\nwith:\n\n$pPrecision = \\frac{pTP}{pTP+pFP}$\n\n$pRecall = \\frac{pTP}{pTP+pFN}$\n\nYou can find a Python implementation of the metric [here](https://www.kaggle.com/code/sohier/probabilistic-f-score)\n\nOur model should output the likelihood of cancer in the corresponding image.\n\nSo what are the labels we will train on?","metadata":{"editable":false}},{"cell_type":"markdown","source":"This is implements the probablistic F score described [here](https://aclanthology.org/2020.eval4nlp-1.9.pdf)","metadata":{"editable":false}},{"cell_type":"code","source":"num_splits = 4\n\nresize_to = (512, 512)\n\ndata_path = '/kaggle/input/rsna-breast-cancer-detection'\ntrain_image_dir = '/kaggle/input/rsna-mammography-images-as-pngs/images_as_pngs_cv2_512'\ntest_dicom_dir = '/kaggle/input/rsna-breast-cancer-detection/test_images'\nmodel_path = '/kaggle/input/rsna-trained-model-weights/tf_effv2_s_208_402/tf_effv2_s_208_402'\n#model_path = '/kaggle/input/rsna-trained-model-weights/res18_119_80/res18_119_80'\n# label_smoothing_weights = torch.tensor([1,10]).float()\n# if torch.cuda.is_available():\n#     label_smoothing_weights = label_smoothing_weights.cuda()","metadata":{"papermill":{"duration":2.925095,"end_time":"2022-12-11T23:17:12.01247","exception":false,"start_time":"2022-12-11T23:17:09.087375","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:15:03.685145Z","iopub.execute_input":"2023-04-18T19:15:03.685847Z","iopub.status.idle":"2023-04-18T19:15:03.692368Z","shell.execute_reply.started":"2023-04-18T19:15:03.685807Z","shell.execute_reply":"2023-04-18T19:15:03.691131Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Creating stratified splits for training","metadata":{"papermill":{"duration":0.005474,"end_time":"2022-12-11T23:17:12.023782","exception":false,"start_time":"2022-12-11T23:17:12.018308","status":"completed"},"tags":[],"editable":false}},{"cell_type":"code","source":"train_csv = pd.read_csv(f'{data_path}/train.csv')\npatient_id_any_cancer = train_csv.groupby('patient_id').cancer.max().reset_index()\nskf = StratifiedKFold(num_splits, shuffle=True, random_state=42)\nsplits = list(skf.split(patient_id_any_cancer.patient_id, patient_id_any_cancer.cancer))","metadata":{"papermill":{"duration":0.140411,"end_time":"2022-12-11T23:17:12.169705","exception":false,"start_time":"2022-12-11T23:17:12.029294","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:15:03.694379Z","iopub.execute_input":"2023-04-18T19:15:03.695261Z","iopub.status.idle":"2023-04-18T19:15:03.799749Z","shell.execute_reply.started":"2023-04-18T19:15:03.695219Z","shell.execute_reply":"2023-04-18T19:15:03.798439Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:15:03.803506Z","iopub.execute_input":"2023-04-18T19:15:03.804256Z","iopub.status.idle":"2023-04-18T19:15:03.831739Z","shell.execute_reply.started":"2023-04-18T19:15:03.804215Z","shell.execute_reply":"2023-04-18T19:15:03.830454Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patient_id_any_cancer.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:15:03.833386Z","iopub.execute_input":"2023-04-18T19:15:03.834117Z","iopub.status.idle":"2023-04-18T19:15:03.846886Z","shell.execute_reply.started":"2023-04-18T19:15:03.834062Z","shell.execute_reply":"2023-04-18T19:15:03.845669Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"splits","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:15:03.848542Z","iopub.execute_input":"2023-04-18T19:15:03.84929Z","iopub.status.idle":"2023-04-18T19:15:03.85964Z","shell.execute_reply.started":"2023-04-18T19:15:03.849247Z","shell.execute_reply":"2023-04-18T19:15:03.858045Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:15:03.8624Z","iopub.execute_input":"2023-04-18T19:15:03.863135Z","iopub.status.idle":"2023-04-18T19:15:03.886104Z","shell.execute_reply.started":"2023-04-18T19:15:03.863082Z","shell.execute_reply":"2023-04-18T19:15:03.884915Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining some functions","metadata":{"papermill":{"duration":0.005815,"end_time":"2022-12-11T23:17:12.182599","exception":false,"start_time":"2022-12-11T23:17:12.176784","status":"completed"},"tags":[],"editable":false}},{"cell_type":"code","source":"#https://www.kaggle.com/competitions/rsna-breast-cancer-detection/discussion/369267  \ndef pfbeta_torch(preds, labels, beta=1):\n    if preds.dim() != 2 or (preds.dim() == 2 and preds.shape[1] !=2): raise ValueError('Houston, we got a problem')\n    preds = preds[:, 1]\n    preds = preds.clip(0, 1)\n    y_true_count = labels.sum()\n    ctp = preds[labels==1].sum()\n    cfp = preds[labels==0].sum()\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / y_true_count\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall)\n        return result\n    else:\n        return 0.0\n\ndef optimize_preds(preds, labels=None, thresh=None, return_thresh=False, print_results=False):\n    preds = preds.clone()\n    if labels is not None: without_thresh = pfbeta_torch(preds, labels)\n    \n    if not thresh and labels is not None:\n        threshs = np.linspace(0, 1, 101)\n        f1s = [pfbeta_torch((preds > thr).float(), labels) for thr in threshs]\n        idx = np.argmax(f1s)\n        thresh, best_pfbeta = threshs[idx], f1s[idx]\n\n    preds = (preds > thresh).float()\n\n    if print_results:\n        print(f'without optimization: {without_thresh}')\n        pfbeta = pfbeta_torch(preds, labels)\n        print(f'with optimization: {pfbeta}')\n        print(f'best_thresh = {thresh}')\n    if return_thresh:\n        return thresh\n    return preds\n\n# https://www.kaggle.com/competitions/rsna-breast-cancer-detection/discussion/369886    \ndef pfbeta_torch_thresh(preds, labels):\n    optimized_preds = optimize_preds(preds, labels)\n    return pfbeta_torch(optimized_preds, labels)\n\nfn2label = {fn: cancer_or_not for fn, cancer_or_not in zip(train_csv['image_id'].astype('str'), train_csv['cancer'])}\n\ndef splitting_func(paths):\n    train = []\n    valid = []\n    for idx, path in enumerate(paths):\n        if int(path.parent.name) in patient_id_any_cancer.iloc[splits[SPLIT][0]].patient_id.values:\n            train.append(idx)\n        else:\n            valid.append(idx)\n    return train, valid\n\ndef label_func(path):\n    return fn2label[path.stem]\n\ndef get_items(image_dir_path):\n    items = []\n    for p in get_image_files(image_dir_path):\n        items.append(p)\n        if p.stem in fn2label and int(p.parent.name) in patient_id_any_cancer.iloc[splits[SPLIT][0]].patient_id.values:\n            if label_func(p) == 1:\n                for _ in range(5):\n                    items.append(p)\n    return items","metadata":{"_kg_hide-input":true,"papermill":{"duration":0.070613,"end_time":"2022-12-11T23:17:12.269689","exception":false,"start_time":"2022-12-11T23:17:12.199076","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:15:03.887889Z","iopub.execute_input":"2023-04-18T19:15:03.888631Z","iopub.status.idle":"2023-04-18T19:15:04.001767Z","shell.execute_reply.started":"2023-04-18T19:15:03.88858Z","shell.execute_reply":"2023-04-18T19:15:04.00046Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Wrapping getting data and getting a model into functions -- this way our logic for training will be cleaner to read.","metadata":{"papermill":{"duration":0.005327,"end_time":"2022-12-11T23:17:12.280715","exception":false,"start_time":"2022-12-11T23:17:12.275388","status":"completed"},"tags":[],"editable":false}},{"cell_type":"code","source":"from timm.models.layers.adaptive_avgmax_pool import SelectAdaptivePool2d\nfrom torch.nn import Flatten\n\ndef get_dataloaders():\n    train_image_path = train_image_dir\n\n    dblock = DataBlock(\n        blocks    = (ImageBlock, CategoryBlock),\n        get_items = get_items,\n        get_y = label_func,\n        splitter  = splitting_func,\n        batch_tfms=[Flip()],\n    )\n    dsets = dblock.datasets(train_image_path)\n    return dblock.dataloaders(train_image_path, batch_size=32)\n\ndef get_learner(arch=resnet18):\n    learner = vision_learner(\n        get_dataloaders(),\n        arch,\n        custom_head=nn.Sequential(SelectAdaptivePool2d(pool_type='avg', flatten=Flatten()), nn.Linear(1280, 2)),\n        #custom_head=nn.Sequential(SelectAdaptivePool2d(pool_type='avg', flatten=Flatten()), nn.Linear(512, 1)),\n        metrics=[\n            error_rate,\n            AccumMetric(pfbeta_torch, activation=ActivationType.Softmax, flatten=False),\n            AccumMetric(pfbeta_torch_thresh, activation=ActivationType.Softmax, flatten=False)\n        ],\n        loss_func=CrossEntropyLossFlat(weight=torch.tensor([1,50]).float()),\n        pretrained=True,\n        normalize=False\n    ).to_fp16()\n    return learner","metadata":{"papermill":{"duration":0.016358,"end_time":"2022-12-11T23:17:12.302564","exception":false,"start_time":"2022-12-11T23:17:12.286206","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:15:04.003653Z","iopub.execute_input":"2023-04-18T19:15:04.004349Z","iopub.status.idle":"2023-04-18T19:15:04.015736Z","shell.execute_reply.started":"2023-04-18T19:15:04.00431Z","shell.execute_reply":"2023-04-18T19:15:04.014499Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#for SPLIT in range(num_splits):\nSPLIT = 0\ndls = get_dataloaders()","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:15:04.017835Z","iopub.execute_input":"2023-04-18T19:15:04.018615Z","iopub.status.idle":"2023-04-18T19:17:24.241143Z","shell.execute_reply.started":"2023-04-18T19:15:04.018552Z","shell.execute_reply":"2023-04-18T19:17:24.239941Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls.show_batch()","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:17:24.24304Z","iopub.execute_input":"2023-04-18T19:17:24.243821Z","iopub.status.idle":"2023-04-18T19:17:26.564978Z","shell.execute_reply.started":"2023-04-18T19:17:24.243781Z","shell.execute_reply":"2023-04-18T19:17:26.563558Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Creating the learner and training","metadata":{"papermill":{"duration":0.005319,"end_time":"2022-12-11T23:17:12.313292","exception":false,"start_time":"2022-12-11T23:17:12.307973","status":"completed"},"tags":[],"editable":false}},{"cell_type":"code","source":"# This is a dependency that is needed for reading DICOM images\n\ntry:\n    import pylibjpeg\nexcept:\n    !rm -rf /root/.cache/torch/hub/checkpoints/\n    !mkdir -p /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}\n\n# copying the pretrained weights\n\nif not os.path.exists('/root/.cache/torch/hub/checkpoints/'):\n        os.makedirs('/root/.cache/torch/hub/checkpoints/')\n!cp '/kaggle/input/pretrained-model-weights-for-fastai/resnet18-f37072fd.pth' '/root/.cache/torch/hub/checkpoints/resnet18-f37072fd.pth'\n!cp '/kaggle/input/pretrained-model-weights-for-fastai/tf_efficientnetv2_s-eb54923e.pth' '/root/.cache/torch/hub/checkpoints/tf_efficientnetv2_s-eb54923e.pth'","metadata":{"papermill":{"duration":110.135069,"end_time":"2022-12-11T23:19:02.453784","exception":false,"start_time":"2022-12-11T23:17:12.318715","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:17:26.567334Z","iopub.execute_input":"2023-04-18T19:17:26.568366Z","iopub.status.idle":"2023-04-18T19:17:29.389824Z","shell.execute_reply.started":"2023-04-18T19:17:26.568313Z","shell.execute_reply":"2023-04-18T19:17:29.388234Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\npreds, labels = [], []\n\nSPLIT = 0 # our learner needs this to construct its dataloaders...\nlearn = get_learner('tf_efficientnetv2_s')\n#learn = get_learner('resnet18')","metadata":{"papermill":{"duration":124.064883,"end_time":"2022-12-11T23:21:06.525622","exception":false,"start_time":"2022-12-11T23:19:02.460739","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:17:29.399273Z","iopub.execute_input":"2023-04-18T19:17:29.401802Z","iopub.status.idle":"2023-04-18T19:19:47.510695Z","shell.execute_reply.started":"2023-04-18T19:17:29.401756Z","shell.execute_reply":"2023-04-18T19:19:47.50954Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn.summary()","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:19:47.515391Z","iopub.execute_input":"2023-04-18T19:19:47.517898Z","iopub.status.idle":"2023-04-18T19:19:48.368063Z","shell.execute_reply.started":"2023-04-18T19:19:47.517854Z","shell.execute_reply":"2023-04-18T19:19:48.366937Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Build a Classification Interpretation object from our learn model\n# it can show us where the model made the worse predictions:\ninterp = ClassificationInterpretation.from_learner(learn)","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:19:48.372682Z","iopub.execute_input":"2023-04-18T19:19:48.375218Z","iopub.status.idle":"2023-04-18T19:21:59.574769Z","shell.execute_reply.started":"2023-04-18T19:19:48.375173Z","shell.execute_reply":"2023-04-18T19:21:59.573545Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot the top ‘n’ classes where the classifier has least precision.\ninterp.plot_top_losses(9, figsize=(12,12))","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:21:59.576751Z","iopub.execute_input":"2023-04-18T19:21:59.577243Z","iopub.status.idle":"2023-04-18T19:22:01.108462Z","shell.execute_reply.started":"2023-04-18T19:21:59.577192Z","shell.execute_reply":"2023-04-18T19:22:01.107418Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"interp.plot_confusion_matrix(figsize=(6,6), dpi=60)","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:22:01.110456Z","iopub.execute_input":"2023-04-18T19:22:01.111151Z","iopub.status.idle":"2023-04-18T19:24:00.765294Z","shell.execute_reply.started":"2023-04-18T19:22:01.111111Z","shell.execute_reply":"2023-04-18T19:24:00.763714Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predicting on test<a id=\"section-two\">","metadata":{"papermill":{"duration":0.009508,"end_time":"2022-12-11T23:21:06.575219","exception":false,"start_time":"2022-12-11T23:21:06.565711","status":"completed"},"tags":[],"editable":false}},{"cell_type":"code","source":"pip install dicomsdl","metadata":{"execution":{"iopub.status.busy":"2023-04-18T19:24:00.772371Z","iopub.execute_input":"2023-04-18T19:24:00.776402Z","iopub.status.idle":"2023-04-18T19:24:10.263795Z","shell.execute_reply.started":"2023-04-18T19:24:00.776337Z","shell.execute_reply":"2023-04-18T19:24:10.262539Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport dicomsdl\n    \nfrom pathlib import Path\nimport multiprocessing as mp\nimport cv2\n\n!rm -rf test_resized_{resize_to[0]}\n\ndef dicom_file_to_ary(path):\n    dcm_file = dicomsdl.open(str(path))\n    data = dcm_file.pixelData()\n\n    data = (data - data.min()) / (data.max() - data.min())\n\n    if dcm_file.getPixelDataInfo()['PhotometricInterpretation'] == \"MONOCHROME1\":\n        data = 1 - data\n\n    data = cv2.resize(data, resize_to)\n    data = (data * 255).astype(np.uint8)\n    return data\n\ndirectories = list(Path(test_dicom_dir).iterdir())\n\ndef process_directory(directory_path):\n    parent_directory = str(directory_path).split('/')[-1]\n    !mkdir -p test_resized_{resize_to[0]}/{parent_directory}\n    for image_path in directory_path.iterdir():\n        processed_ary = dicom_file_to_ary(image_path)\n        cv2.imwrite(\n            f'test_resized_{resize_to[0]}/{parent_directory}/{image_path.stem}.png',\n            processed_ary\n        )\n\nwith mp.Pool(mp.cpu_count()) as p:\n    p.map(process_directory, directories)","metadata":{"_kg_hide-input":true,"papermill":{"duration":4.959966,"end_time":"2022-12-11T23:21:11.545524","exception":false,"start_time":"2022-12-11T23:21:06.585558","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:24:10.266458Z","iopub.execute_input":"2023-04-18T19:24:10.266955Z","iopub.status.idle":"2023-04-18T19:24:14.842789Z","shell.execute_reply.started":"2023-04-18T19:24:10.266907Z","shell.execute_reply":"2023-04-18T19:24:14.841304Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\npreds_all = []\n\ntest_dl = learn.dls.test_dl(get_image_files(f'test_resized_{resize_to[0]}'))\nfor SPLIT in range(num_splits):\n    print(SPLIT,(f'{model_path}/{SPLIT}'))\n    learn.load(f'{model_path}/{SPLIT}')\n    preds, _ = learn.get_preds(dl=test_dl)\n    preds_all.append(preds)","metadata":{"papermill":{"duration":19.402972,"end_time":"2022-12-11T23:21:30.955539","exception":false,"start_time":"2022-12-11T23:21:11.552567","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:24:14.844984Z","iopub.execute_input":"2023-04-18T19:24:14.845407Z","iopub.status.idle":"2023-04-18T19:24:18.394518Z","shell.execute_reply.started":"2023-04-18T19:24:14.845354Z","shell.execute_reply":"2023-04-18T19:24:18.393206Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"threshold = 0.4\npreds = torch.zeros_like(preds_all[0])\nfor pred in preds_all:\n    preds += pred\n\npreds /= num_splits\n\npreds = optimize_preds(preds, thresh=threshold)\nimage_ids = [path.stem for path in test_dl.items]\n\nimage_id2pred = defaultdict(lambda: 0)\nfor image_id, pred in zip(image_ids, preds[:, 1]):\n    image_id2pred[int(image_id)] = pred.item()","metadata":{"papermill":{"duration":0.019482,"end_time":"2022-12-11T23:21:30.982952","exception":false,"start_time":"2022-12-11T23:21:30.96347","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:24:18.396215Z","iopub.execute_input":"2023-04-18T19:24:18.397769Z","iopub.status.idle":"2023-04-18T19:24:18.405805Z","shell.execute_reply.started":"2023-04-18T19:24:18.397728Z","shell.execute_reply":"2023-04-18T19:24:18.404793Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-three\"></a>\n# Making a submission","metadata":{"papermill":{"duration":0.007116,"end_time":"2022-12-11T23:21:30.997643","exception":false,"start_time":"2022-12-11T23:21:30.990527","status":"completed"},"tags":[],"editable":false}},{"cell_type":"code","source":"test_csv = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/test.csv')\n\nprediction_ids = []\npreds = []\n\nfor _, row in test_csv.iterrows():\n    prediction_ids.append(row.prediction_id)\n    preds.append(image_id2pred[row.image_id])\n\nsubmission = pd.DataFrame(data={'prediction_id': prediction_ids, 'cancer': preds}).groupby('prediction_id').max().reset_index()\nsubmission.head()","metadata":{"papermill":{"duration":0.057453,"end_time":"2022-12-11T23:21:31.062316","exception":false,"start_time":"2022-12-11T23:21:31.004863","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:24:18.407582Z","iopub.execute_input":"2023-04-18T19:24:18.408187Z","iopub.status.idle":"2023-04-18T19:24:18.432309Z","shell.execute_reply.started":"2023-04-18T19:24:18.408148Z","shell.execute_reply":"2023-04-18T19:24:18.431096Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=False)","metadata":{"papermill":{"duration":0.036385,"end_time":"2022-12-11T23:21:31.116858","exception":false,"start_time":"2022-12-11T23:21:31.080473","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T19:24:18.434121Z","iopub.execute_input":"2023-04-18T19:24:18.434534Z","iopub.status.idle":"2023-04-18T19:24:18.442748Z","shell.execute_reply.started":"2023-04-18T19:24:18.434476Z","shell.execute_reply":"2023-04-18T19:24:18.441168Z"},"editable":false,"trusted":true},"execution_count":null,"outputs":[]}]}