{"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":"# Stroke Blood Clot Origin🔬: classification baseline with ⚡Flash","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y torchtext\n# !pip install -q --upgrade torch torchvision\n!mkdir -p frozen_packages\n!cp ../input/starter-flash-semantic-segmentation/frozen_packages/* frozen_packages/\n!pip install -q \"lightning-flash[image]\" \"torchmetrics<0.8\" --no-index --find-links frozen_packages/\n!pip install -q -U timm --no-index --find-links frozen_packages/\n!rm -rf frozen_packages\n\n! pip list | grep -e torch -e lightning\n! nvidia-smi -L","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-23T21:16:52.738934Z","iopub.execute_input":"2022-07-23T21:16:52.739391Z","iopub.status.idle":"2022-07-23T21:18:29.030965Z","shell.execute_reply.started":"2022-07-23T21:16:52.739293Z","shell.execute_reply":"2022-07-23T21:18:29.029356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os, glob\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nDATASET_FOLDER = \"/kaggle/input/mayo-clinic-strip-ai/\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-23T21:18:29.034182Z","iopub.execute_input":"2022-07-23T21:18:29.035249Z","iopub.status.idle":"2022-07-23T21:18:29.042016Z","shell.execute_reply.started":"2022-07-23T21:18:29.03521Z","shell.execute_reply":"2022-07-23T21:18:29.040581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path_csv = os.path.join(DATASET_FOLDER, \"train.csv\")\ndf_train = pd.read_csv(path_csv)\ndisplay(df_train.head())","metadata":{"execution":{"iopub.status.busy":"2022-07-23T21:18:29.044123Z","iopub.execute_input":"2022-07-23T21:18:29.044605Z","iopub.status.idle":"2022-07-23T21:18:29.086047Z","shell.execute_reply.started":"2022-07-23T21:18:29.044555Z","shell.execute_reply":"2022-07-23T21:18:29.084905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Converting test images\n\nthe image conversion is using https://www.kaggle.com/code/jirkaborovec/bloodclots-classif-eda-load-crop-images","metadata":{}},{"cell_type":"code","source":"from PIL import Image\n\nImage.MAX_IMAGE_PIXELS = 25_000_000_000\n\ndef prune_image_rows_cols(im, mask, thr=0.990):\n    # delete empty columns\n    for l in reversed(range(im.shape[1])):\n        if (np.sum(mask[:, l]) / float(mask.shape[0])) > thr:\n            im = np.delete(im, l, 1)\n    # delete empty rows\n    for l in reversed(range(im.shape[0])):\n        if (np.sum(mask[l, :]) / float(mask.shape[1])) > thr:\n            im = np.delete(im, l, 0)\n    return im\n\n\ndef mask_median(im, val=255):\n    masks = [None] * 3\n    for c in range(3):\n        masks[c] = im[..., c] >= np.median(im[:, :, c]) - 5\n    mask = np.logical_and(*masks)\n    im[mask, :] = val\n    return im, mask\n\n\ndef image_load_scale_norm(img_path, prune_thr=0.990, bg_val=255):\n    img = Image.open(img_path)\n    if (img.width * img.height) > 3_000_000_000:\n        print(img.width, img.height)\n        return None\n    scale = min(img.height / 2e3, img.width / 2e3)\n    tmp_size = int(img.width / scale), int(img.height / scale)\n    img.thumbnail(tmp_size, resample=Image.Resampling.BILINEAR, reducing_gap=scale)\n    im, mask = mask_median(np.array(img), val=bg_val)\n    im = prune_image_rows_cols(im, mask, thr=prune_thr)\n    img = Image.fromarray(im)\n    scale = min(img.height / 1e3, img.width / 1e3)\n    if scale > 1:\n        img = img.resize((int(img.width / scale), int(img.height / scale)), Image.ANTIALIAS)\n    return img","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-23T21:18:29.088513Z","iopub.execute_input":"2022-07-23T21:18:29.089072Z","iopub.status.idle":"2022-07-23T21:18:29.105029Z","shell.execute_reply.started":"2022-07-23T21:18:29.08904Z","shell.execute_reply":"2022-07-23T21:18:29.103627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nfrom tqdm.auto import tqdm\n\nls_imgs_tif = glob.glob(os.path.join(DATASET_FOLDER, \"test\", \"*.tif\"))\nnames = [os.path.splitext(os.path.basename(p))[0] for p in ls_imgs_tif]\npatient_ids = set([n.split(\"_\")[0] for n in names])\n\n! mkdir -p /kaggle/temp/images\n\nfor img_path in tqdm(ls_imgs_tif):\n    name, _ = os.path.splitext(os.path.basename(img_path))\n    img = image_load_scale_norm(img_path)\n    if not img:\n        print(f\"missing: {name}\")\n        continue\n    img.save(os.path.join(\"/kaggle/temp/images\", f\"{name}.png\"))\n    del img\n    gc.collect()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-23T21:18:29.106979Z","iopub.execute_input":"2022-07-23T21:18:29.107554Z","iopub.status.idle":"2022-07-23T21:21:07.501691Z","shell.execute_reply.started":"2022-07-23T21:18:29.10751Z","shell.execute_reply":"2022-07-23T21:21:07.500444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loading model...","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_transform import InputTransform\n\n@dataclass\nclass ImageClassifInputTransform(InputTransform):\n\n    image_size: Tuple[int, int] = (256, 256)\n    # Default from ImageNet\n    color_mean: Tuple[float, float, float] = (0.947, 0.881, 0.863)\n    color_std: Tuple[float, float, float] = (0.093, 0.201, 0.245)\n\n    def input_per_sample_transform(self) -> Callable:\n        return T.Compose([\n            T.ToTensor(),\n            T.CenterCrop(size=(800, 800)),\n            T.Resize(self.image_size),\n            T.Normalize(self.color_mean, self.color_std),\n        ])\n\n    def target_per_sample_transform(self) -> Callable:\n        return torch.as_tensor","metadata":{"execution":{"iopub.status.busy":"2022-07-23T21:21:07.503349Z","iopub.execute_input":"2022-07-23T21:21:07.503711Z","iopub.status.idle":"2022-07-23T21:21:19.844323Z","shell.execute_reply.started":"2022-07-23T21:21:07.503661Z","shell.execute_reply":"2022-07-23T21:21:19.842932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport pytorch_lightning as pl\n\nimport flash\nfrom flash.image import ImageClassificationData, ImageClassifier\n\ntrainer = flash.Trainer(gpus=torch.cuda.device_count())\n\nmodel = ImageClassifier.load_from_checkpoint(\n    \"../input/bloodclots-classif-baseline-flash-effnet-aug/image_classification_model.pt\",\n    pretrained=False,\n)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-23T21:21:19.84618Z","iopub.execute_input":"2022-07-23T21:21:19.847325Z","iopub.status.idle":"2022-07-23T21:21:29.889309Z","shell.execute_reply.started":"2022-07-23T21:21:19.847283Z","shell.execute_reply":"2022-07-23T21:21:29.88792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRANSFORM_PARAMS = {\n    \"image_size\": (528, 528),\n    # \"mean\": (0.485, 0.456, 0.406),\n    # \"std\": (0.229, 0.224, 0.225),\n}","metadata":{"execution":{"iopub.status.busy":"2022-07-23T21:21:29.891191Z","iopub.execute_input":"2022-07-23T21:21:29.891994Z","iopub.status.idle":"2022-07-23T21:21:29.898451Z","shell.execute_reply.started":"2022-07-23T21:21:29.891942Z","shell.execute_reply":"2022-07-23T21:21:29.896989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference 🔥!","metadata":{}},{"cell_type":"code","source":"ls_imgs_png = glob.glob(os.path.join(\"/kaggle/temp/images\", \"*.png\"))\n\nfig, axes = plt.subplots(nrows=3, figsize=(8, 12))\nfor i, img_path in enumerate(ls_imgs_png[:3]):\n    img = plt.imread(img_path)\n    if img.shape[0] > img.shape[1]:\n        img = np.rollaxis(img, 1, 0)\n    axes[i].imshow(img)","metadata":{"execution":{"iopub.status.busy":"2022-07-23T21:21:29.900256Z","iopub.execute_input":"2022-07-23T21:21:29.901567Z","iopub.status.idle":"2022-07-23T21:21:31.641525Z","shell.execute_reply.started":"2022-07-23T21:21:29.901514Z","shell.execute_reply":"2022-07-23T21:21:31.640152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from itertools import chain\n\ndm = ImageClassificationData.from_files(\n    predict_files=ls_imgs_png,\n    batch_size=3,\n    predict_transform=ImageClassifInputTransform,\n    transform_kwargs=TRANSFORM_PARAMS,\n)\nprint(model.labels)\npredictions = trainer.predict(model, datamodule=dm, output=\"probabilities\")\npredictions = [dict(zip(model.labels, pred)) for pred in chain(*predictions)]\ndf_pred = pd.DataFrame(predictions)\ndisplay(df_pred.head())","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-23T21:21:31.64622Z","iopub.execute_input":"2022-07-23T21:21:31.646732Z","iopub.status.idle":"2022-07-23T21:21:37.635683Z","shell.execute_reply.started":"2022-07-23T21:21:31.646665Z","shell.execute_reply":"2022-07-23T21:21:37.634361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"names = [os.path.splitext(os.path.basename(p))[0] for p in ls_imgs_png]\ndf_pred[\"patient_id\"] = [n.split(\"_\")[0] for n in names]","metadata":{"execution":{"iopub.status.busy":"2022-07-23T21:21:37.637501Z","iopub.execute_input":"2022-07-23T21:21:37.638227Z","iopub.status.idle":"2022-07-23T21:21:37.649949Z","shell.execute_reply.started":"2022-07-23T21:21:37.638163Z","shell.execute_reply":"2022-07-23T21:21:37.648933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Fill missing cases","metadata":{}},{"cell_type":"code","source":"# fill mmissing if skipped for too large image size\nmissed = [{\"CE\": 0.5, \"LAA\": 0.5, \"patient_id\": pid}\n          for pid in patient_ids if pid not in df_pred[\"patient_id\"].values]\ndf_pred = df_pred.append(pd.DataFrame(missed), ignore_index=True)\ndisplay(df_pred.head())","metadata":{"execution":{"iopub.status.busy":"2022-07-23T21:21:37.651411Z","iopub.execute_input":"2022-07-23T21:21:37.65249Z","iopub.status.idle":"2022-07-23T21:21:37.679213Z","shell.execute_reply.started":"2022-07-23T21:21:37.652453Z","shell.execute_reply":"2022-07-23T21:21:37.677775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Aggregate multiple cases per `parent_id`","metadata":{}},{"cell_type":"code","source":"# in case there are mo samples use mean value\ndf_pred = df_pred.groupby(\"patient_id\").mean()\ndisplay(df_pred.head())","metadata":{"execution":{"iopub.status.busy":"2022-07-23T21:21:37.680817Z","iopub.execute_input":"2022-07-23T21:21:37.681164Z","iopub.status.idle":"2022-07-23T21:21:37.702491Z","shell.execute_reply.started":"2022-07-23T21:21:37.681133Z","shell.execute_reply":"2022-07-23T21:21:37.70158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Export submission","metadata":{}},{"cell_type":"code","source":"df_pred[[\"CE\", \"LAA\"]].round(6).to_csv(\"submission.csv\")\n\n!head submission.csv","metadata":{"execution":{"iopub.status.busy":"2022-07-23T21:21:37.703456Z","iopub.execute_input":"2022-07-23T21:21:37.703818Z","iopub.status.idle":"2022-07-23T21:21:38.509887Z","shell.execute_reply.started":"2022-07-23T21:21:37.703786Z","shell.execute_reply":"2022-07-23T21:21:38.508628Z"},"trusted":true},"execution_count":null,"outputs":[]}]}