{"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":"# Inference for Mammography ⚕️ Breast Cancer with ⚡ Flash\n\nThis is follow-up of training baseline: https://www.kaggle.com/code/jirkaborovec/mammography-baseline-flash-effnet-augment","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"!mkdir -p /tmp/frozen_packages\n!cp /kaggle/input/pydicom-package/* /tmp/frozen_packages\n!cp /kaggle/input/mammography-baseline-flash-effnet-augment/frozen_packages/* /tmp/frozen_packages\n!ls /tmp/frozen_packages","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-07T08:05:50.219705Z","iopub.execute_input":"2022-12-07T08:05:50.220768Z","iopub.status.idle":"2022-12-07T08:06:11.266783Z","shell.execute_reply.started":"2022-12-07T08:05:50.220601Z","shell.execute_reply":"2022-12-07T08:06:11.26565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip uninstall -y torchtext\n# !pip install -q --upgrade torch torchvision\n!pip install -q \"lightning-flash[image]\" --find-links /tmp/frozen_packages --no-index\n!pip install -q -U timm --find-links /tmp/frozen_packages --no-index\n!pip install -qU \"python-gdcm\" pydicom pylibjpeg \"opencv-python-headless\" -f /tmp/frozen_packages --no-index\n\n! pip list | grep -e torch -e lightning\n! nvidia-smi -L","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-07T08:06:11.268967Z","iopub.execute_input":"2022-12-07T08:06:11.269261Z","iopub.status.idle":"2022-12-07T08:07:26.341878Z","shell.execute_reply.started":"2022-12-07T08:06:11.269228Z","shell.execute_reply":"2022-12-07T08:07:26.340707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%matplotlib inline\n\nimport os, glob\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nimport flash\nfrom flash.image import ImageClassificationData, ImageClassifier\n\nPATH_DATASET = \"/kaggle/input/rsna-breast-cancer-detection\"","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-07T08:07:26.34399Z","iopub.execute_input":"2022-12-07T08:07:26.344314Z","iopub.status.idle":"2022-12-07T08:07:39.561334Z","shell.execute_reply.started":"2022-12-07T08:07:26.344281Z","shell.execute_reply":"2022-12-07T08:07:39.560342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test = pd.read_csv(os.path.join(PATH_DATASET, \"test.csv\"))\ndisplay(df_test.head())\nprint(f'cases: {len(df_test)}')","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:07:39.564192Z","iopub.execute_input":"2022-12-07T08:07:39.564983Z","iopub.status.idle":"2022-12-07T08:07:39.594756Z","shell.execute_reply.started":"2022-12-07T08:07:39.564944Z","shell.execute_reply":"2022-12-07T08:07:39.592719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ls_image = glob.glob(os.path.join(PATH_DATASET, \"test_images\", \"*\", \"*.dcm\"))\nprint(f\"found images: {len(ls_image)}\")","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:07:39.596Z","iopub.execute_input":"2022-12-07T08:07:39.596978Z","iopub.status.idle":"2022-12-07T08:08:19.318443Z","shell.execute_reply.started":"2022-12-07T08:07:39.596939Z","shell.execute_reply":"2022-12-07T08:08:19.317393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Convert test images\n\nCheck the full dataset conversion: https://www.kaggle.com/code/jirkaborovec/mammography-convert-resize-dicom-png","metadata":{}},{"cell_type":"code","source":"import pydicom\nfrom PIL import Image\nfrom pydicom.pixel_data_handlers import apply_windowing\n\ndef convert_dicom(dicom_path, output_dir, img_size: int = 720):\n    dicom = pydicom.dcmread(dicom_path)\n    data = dicom.pixel_array\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = np.amax(data) - data\n    img = apply_windowing(data, dicom)\n    img = (img.astype(float) - img.min()) / (img.max() - img.min())\n    img = Image.fromarray((img * 255).astype(np.uint8))\n\n    img_name, _ = os.path.splitext(os.path.basename(dicom_path))\n    img_dir = os.path.basename(os.path.dirname(dicom_path))\n    png_path = os.path.join(output_dir, img_dir, f\"{img_name}.png\")\n    os.makedirs(os.path.dirname(png_path), exist_ok=True)\n    # plt.imsave(png_path, (img * 255).astype(np.uint8))\n\n    img.thumbnail((img_size, img_size))\n    img.save(png_path)","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:08:19.320127Z","iopub.execute_input":"2022-12-07T08:08:19.320835Z","iopub.status.idle":"2022-12-07T08:08:19.464681Z","shell.execute_reply.started":"2022-12-07T08:08:19.32079Z","shell.execute_reply":"2022-12-07T08:08:19.463734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from joblib import Parallel, delayed\nfrom tqdm.auto import tqdm\n\nPATH_CONVERT = \"/tmp/test_images\"\n\n_= Parallel(n_jobs=6)(\n    delayed(convert_dicom)(p_img, output_dir=PATH_CONVERT)\n    for p_img in tqdm(ls_image)\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:08:19.466237Z","iopub.execute_input":"2022-12-07T08:08:19.466594Z","iopub.status.idle":"2022-12-07T08:15:11.574981Z","shell.execute_reply.started":"2022-12-07T08:08:19.466547Z","shell.execute_reply":"2022-12-07T08:15:11.572206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 1. Create the DataModule\n\nNote that the validation augmentation and especially color normalization need to be the same as you used for training","metadata":{}},{"cell_type":"code","source":"import torch\nfrom dataclasses import dataclass\nfrom torchvision import transforms as T\nfrom typing import Tuple, Callable, Optional\nfrom flash.core.data.io.input import DataKeys\nfrom flash.core.data.io.input_transform import InputTransform\nfrom flash.core.data.transforms import ApplyToKeys\n\ndef _streaming(transforms):\n    return T.Compose([\n        ApplyToKeys(\n            DataKeys.INPUT,\n            T.Compose(transforms),\n        ),\n        ApplyToKeys(DataKeys.TARGET, torch.as_tensor),\n    ])\n\n@dataclass\nclass ImageClassifInputTransform(InputTransform):\n\n    image_size: Tuple[int, int] = (256, 256)\n    color_mean: float = 0.09092962741851807\n    color_std: float = 0.142587348818779\n\n    def per_sample_transform(self) -> Callable:\n        return _streaming([\n            T.ToTensor(),\n#             T.CenterCrop(size=(800, 800)),  # TODO\n            T.Resize(self.image_size),\n            T.Normalize(self.color_mean, self.color_std),\n        ])","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:15:19.790638Z","iopub.execute_input":"2022-12-07T08:15:19.791016Z","iopub.status.idle":"2022-12-07T08:15:19.801883Z","shell.execute_reply.started":"2022-12-07T08:15:19.790982Z","shell.execute_reply":"2022-12-07T08:15:19.800963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ls_imgs_png = glob.glob(os.path.join(PATH_CONVERT, \"*\", \"*.png\"))\nprint(len(ls_imgs_png))","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:15:19.803824Z","iopub.execute_input":"2022-12-07T08:15:19.80442Z","iopub.status.idle":"2022-12-07T08:15:19.824227Z","shell.execute_reply.started":"2022-12-07T08:15:19.804385Z","shell.execute_reply":"2022-12-07T08:15:19.823279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRANSFORM_PARAMS = {\n    \"image_size\": (456, 456),\n}\n\ndatamodule = ImageClassificationData.from_files(\n    predict_files=ls_imgs_png,\n    # size can be larger then for training as you do not use gradients\n    batch_size=40,\n    transform=ImageClassifInputTransform,\n    transform_kwargs=TRANSFORM_PARAMS,\n    num_workers=4,\n)\n\nprint(len(datamodule.predict_dataloader()))","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:15:19.82731Z","iopub.execute_input":"2022-12-07T08:15:19.828011Z","iopub.status.idle":"2022-12-07T08:15:19.845338Z","shell.execute_reply.started":"2022-12-07T08:15:19.827984Z","shell.execute_reply":"2022-12-07T08:15:19.844127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Note that these images are aready color normalized soe the collors would look diffrent / not natural\n\nHINT: if you want to see the images in human friendly from, drop color normalization","metadata":{}},{"cell_type":"code","source":"img_idx, spl_max, nb_rows = 0, 4, 2\nfig, axarr = plt.subplots(ncols=2, nrows=nb_rows, figsize=(8, 5 * nb_rows))\n\nfor batch in datamodule.predict_dataloader():\n    imgs = [np.rollaxis(im.numpy(), 0, 3) for im in batch['input']]\n    for i in range(len(imgs)):\n        if img_idx == spl_max:\n            break\n        print(imgs[i].min(), imgs[i].mean(), imgs[i].max())\n        axarr[img_idx // 2, img_idx % 2].imshow(imgs[i])\n        img_idx += 1\n    if img_idx >= spl_max:\n        break","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:15:19.848026Z","iopub.execute_input":"2022-12-07T08:15:19.848433Z","iopub.status.idle":"2022-12-07T08:15:30.822626Z","shell.execute_reply.started":"2022-12-07T08:15:19.848399Z","shell.execute_reply":"2022-12-07T08:15:30.821624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Create trainer\n\nJust for running inference, does not need much configuration...","metadata":{}},{"cell_type":"code","source":"import torch\nimport pytorch_lightning as pl\n\ntrainer = flash.Trainer(\n    gpus=int(torch.cuda.is_available()),\n    precision=16 if torch.cuda.is_available() else 32,\n)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-07T08:15:30.826697Z","iopub.execute_input":"2022-12-07T08:15:30.827035Z","iopub.status.idle":"2022-12-07T08:15:31.556357Z","shell.execute_reply.started":"2022-12-07T08:15:30.827002Z","shell.execute_reply":"2022-12-07T08:15:31.554984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Load the model","metadata":{}},{"cell_type":"code","source":"from torch import nn\nimport torch.nn.functional as F\n\nclass F1_Loss(nn.Module):\n    def __init__(self, epsilon=1e-7):\n        super().__init__()\n        self.epsilon = epsilon\n        \n    def forward(self, y_pred, y_true):\n        return None","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:15:31.560251Z","iopub.execute_input":"2022-12-07T08:15:31.568934Z","iopub.status.idle":"2022-12-07T08:15:31.578388Z","shell.execute_reply.started":"2022-12-07T08:15:31.568889Z","shell.execute_reply":"2022-12-07T08:15:31.57683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ImageClassifier.load_from_checkpoint(\n    \"/kaggle/input/mammography-baseline-flash-effnet-augment/image_classification_model.pt\",\n    pretrained=False,\n)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-07T08:15:31.582997Z","iopub.execute_input":"2022-12-07T08:15:31.585761Z","iopub.status.idle":"2022-12-07T08:15:36.739269Z","shell.execute_reply.started":"2022-12-07T08:15:31.585723Z","shell.execute_reply":"2022-12-07T08:15:36.73821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. Predictions","metadata":{}},{"cell_type":"code","source":"from itertools import chain\n\npredictions = trainer.predict(model, datamodule=datamodule, output=\"probabilities\")\npredictions = list(chain(*predictions))","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-07T08:15:36.740727Z","iopub.execute_input":"2022-12-07T08:15:36.741573Z","iopub.status.idle":"2022-12-07T08:15:52.935039Z","shell.execute_reply.started":"2022-12-07T08:15:36.741523Z","shell.execute_reply":"2022-12-07T08:15:52.933928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = []\nfor p_img, pred in zip(ls_imgs_png, predictions):\n    # print(pred)\n    pid = os.path.basename(os.path.dirname(p_img))\n    pim, _ = os.path.splitext(os.path.basename(p_img))\n    preds.append({\n        \"patient_id\": int(pid),\n        \"image_id\": int(pim),\n        \"cancer\": pred[1],\n    })\n\ndf_preds = pd.DataFrame(preds)\ndisplay(df_preds.head())","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:15:52.938878Z","iopub.execute_input":"2022-12-07T08:15:52.939196Z","iopub.status.idle":"2022-12-07T08:15:52.957056Z","shell.execute_reply.started":"2022-12-07T08:15:52.939162Z","shell.execute_reply":"2022-12-07T08:15:52.956135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Fuse predictions for submission\n\nsee the dummy submission in https://www.kaggle.com/code/jirkaborovec/mammography-eda-loading-dicom","metadata":{}},{"cell_type":"code","source":"df_preds = df_test.merge(df_preds, on=[\"patient_id\", \"image_id\"])\ndisplay(df_preds.head())","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:15:52.958422Z","iopub.execute_input":"2022-12-07T08:15:52.959111Z","iopub.status.idle":"2022-12-07T08:15:53.005853Z","shell.execute_reply.started":"2022-12-07T08:15:52.959073Z","shell.execute_reply":"2022-12-07T08:15:53.005022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"use max cancer then average as single potive view is for the whole subject","metadata":{}},{"cell_type":"code","source":"df_preds = df_preds.groupby('prediction_id').max()\ndf_preds[\"cancer\"].to_csv(\"submission.csv\")\n\n! head submission.csv","metadata":{"execution":{"iopub.status.busy":"2022-12-07T08:15:53.007857Z","iopub.execute_input":"2022-12-07T08:15:53.008503Z","iopub.status.idle":"2022-12-07T08:15:54.181145Z","shell.execute_reply.started":"2022-12-07T08:15:53.008468Z","shell.execute_reply":"2022-12-07T08:15:54.179833Z"},"trusted":true},"execution_count":null,"outputs":[]}]}