{"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":"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\n\nfrom fastai.vision.learner import *\nfrom fastai.data.all import *\nfrom fastai.vision.all import *\nfrom fastai.metrics import ActivationType\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom collections import defaultdict\nimport pandas as pd\nimport numpy as np\nfrom pdb import set_trace","metadata":{"execution":{"iopub.status.busy":"2023-01-23T07:01:46.333299Z","iopub.execute_input":"2023-01-23T07:01:46.333763Z","iopub.status.idle":"2023-01-23T07:03:03.921497Z","shell.execute_reply.started":"2023-01-23T07:01:46.333673Z","shell.execute_reply":"2023-01-23T07:03:03.920384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_EPOCHS = 4\nNUM_SPLITS = 4\n\nRESIZE_TO = (1024, 1024)\n\nDATA_PATH = '/kaggle/input/rsna-breast-cancer-detection'\nTRAIN_IMAGE_DIR = '/kaggle/input/rsna-mammography-images-as-pngs/images_as_pngs_cv2_vl_asp_1024'\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\nlabel_smoothing_weights = torch.tensor([1,10]).float()\nif torch.cuda.is_available():\n    label_smoothing_weights = label_smoothing_weights.cuda()","metadata":{"execution":{"iopub.status.busy":"2023-01-23T07:08:26.072912Z","iopub.execute_input":"2023-01-23T07:08:26.07336Z","iopub.status.idle":"2023-01-23T07:08:26.082835Z","shell.execute_reply.started":"2023-01-23T07:08:26.073319Z","shell.execute_reply":"2023-01-23T07:08:26.081829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-23T07:03:06.906743Z","iopub.execute_input":"2023-01-23T07:03:06.907112Z","iopub.status.idle":"2023-01-23T07:03:07.055274Z","shell.execute_reply.started":"2023-01-23T07:03:06.907078Z","shell.execute_reply":"2023-01-23T07:03:07.054352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\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\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\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,"execution":{"iopub.status.busy":"2023-01-23T07:03:07.057792Z","iopub.execute_input":"2023-01-23T07:03:07.058148Z","iopub.status.idle":"2023-01-23T07:03:07.122266Z","shell.execute_reply.started":"2023-01-23T07:03:07.058116Z","shell.execute_reply":"2023-01-23T07:03:07.121452Z"},"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":{}},{"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        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":{"execution":{"iopub.status.busy":"2023-01-23T07:03:07.123493Z","iopub.execute_input":"2023-01-23T07:03:07.124399Z","iopub.status.idle":"2023-01-23T07:03:07.132953Z","shell.execute_reply.started":"2023-01-23T07:03:07.124365Z","shell.execute_reply":"2023-01-23T07:03:07.132051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This is a dependency that is needed for reading DICOM images\n#creating learner\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":{"execution":{"iopub.status.busy":"2023-01-23T07:03:07.135167Z","iopub.execute_input":"2023-01-23T07:03:07.136444Z","iopub.status.idle":"2023-01-23T07:05:23.033419Z","shell.execute_reply.started":"2023-01-23T07:03:07.136416Z","shell.execute_reply":"2023-01-23T07:05:23.031912Z"},"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\n# instead of training, to conserve pipeline time, I am uploading models trained locally\n# uncomment the lines below for training\n  \n# for SPLIT in range(NUM_SPLITS):\n#     learn = get_learner()\n#     learn.unfreeze()\n#     learn.fit_one_cycle(NUM_EPOCHS, 1e-4, pct_start=0.1)\n#     learn.save(f'{MODEL_PATH}/{SPLIT}')\n        \n#     output = learn.get_preds()\n#     preds.append(output[0])\n#     labels.append(output[1])","metadata":{"execution":{"iopub.status.busy":"2023-01-23T07:08:36.04809Z","iopub.execute_input":"2023-01-23T07:08:36.048474Z","iopub.status.idle":"2023-01-23T07:11:17.43499Z","shell.execute_reply.started":"2023-01-23T07:08:36.048443Z","shell.execute_reply":"2023-01-23T07:11:17.433359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# threshold = optimize_preds(torch.cat(preds), torch.cat(labels), return_thresh=True, print_results=True)\nthreshold = 0.402","metadata":{"execution":{"iopub.status.busy":"2023-01-23T07:11:23.199128Z","iopub.execute_input":"2023-01-23T07:11:23.199496Z","iopub.status.idle":"2023-01-23T07:11:23.204433Z","shell.execute_reply.started":"2023-01-23T07:11:23.199466Z","shell.execute_reply":"2023-01-23T07:11:23.203098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import pydicom\n# from pydicom.pixel_data_handlers.util import apply_voi_lut\n#predicting on test\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,"execution":{"iopub.status.busy":"2023-01-23T07:11:30.204335Z","iopub.execute_input":"2023-01-23T07:11:30.204831Z","iopub.status.idle":"2023-01-23T07:11:34.79779Z","shell.execute_reply.started":"2023-01-23T07:11:30.204798Z","shell.execute_reply":"2023-01-23T07:11:34.796411Z"},"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    learn.load(f'{MODEL_PATH}/{SPLIT}')\n    preds, _ = learn.get_preds(dl=test_dl)\n    preds_all.append(preds)","metadata":{"execution":{"iopub.status.busy":"2023-01-23T07:11:44.070218Z","iopub.execute_input":"2023-01-23T07:11:44.07124Z","iopub.status.idle":"2023-01-23T07:12:05.304952Z","shell.execute_reply.started":"2023-01-23T07:11:44.071174Z","shell.execute_reply":"2023-01-23T07:12:05.303595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = torch.zeros_like(preds_all[0])\nfor pred in preds_all:\n    preds += pred\n\npreds /= NUM_SPLITS\n\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":{"execution":{"iopub.status.busy":"2023-01-23T07:12:10.497133Z","iopub.execute_input":"2023-01-23T07:12:10.497887Z","iopub.status.idle":"2023-01-23T07:12:10.515643Z","shell.execute_reply.started":"2023-01-23T07:12:10.497824Z","shell.execute_reply":"2023-01-23T07:12:10.514487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-23T07:12:14.7984Z","iopub.execute_input":"2023-01-23T07:12:14.798979Z","iopub.status.idle":"2023-01-23T07:12:14.840778Z","shell.execute_reply.started":"2023-01-23T07:12:14.798943Z","shell.execute_reply":"2023-01-23T07:12:14.839669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-01-23T07:05:29.940975Z","iopub.status.idle":"2023-01-23T07:05:29.943466Z","shell.execute_reply.started":"2023-01-23T07:05:29.943182Z","shell.execute_reply":"2023-01-23T07:05:29.943206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}