{"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#!pip install -q 'kaggle-image-segmentation' --no-index --find-links 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-23T17:55:29.696329Z","iopub.execute_input":"2022-07-23T17:55:29.697033Z","iopub.status.idle":"2022-07-23T17:56:53.019687Z","shell.execute_reply.started":"2022-07-23T17:55:29.696949Z","shell.execute_reply":"2022-07-23T17:56:53.018504Z"},"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\nimport flash\nfrom flash.image import ImageClassificationData, ImageClassifier\n\nDATASET_FOLDER = \"/kaggle/input/mayo-clinic-strip-ai/\"\nDATASET_SMALL_FOLDER = \"/kaggle/input/stroke-blood-clot-origin-1k-scale-bg-crop\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-23T17:56:53.022493Z","iopub.execute_input":"2022-07-23T17:56:53.023655Z","iopub.status.idle":"2022-07-23T17:57:03.951406Z","shell.execute_reply.started":"2022-07-23T17:56:53.023611Z","shell.execute_reply":"2022-07-23T17:57:03.950361Z"},"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-23T17:57:03.952817Z","iopub.execute_input":"2022-07-23T17:57:03.953976Z","iopub.status.idle":"2022-07-23T17:57:03.984808Z","shell.execute_reply.started":"2022-07-23T17:57:03.953939Z","shell.execute_reply":"2022-07-23T17:57:03.983792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Color 🦩 normalizations","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom tqdm.auto import tqdm\nfrom joblib import Parallel, delayed\n\ndef _color_means(img_path):\n    img = plt.imread(img_path)\n    if np.max(img) > 1.5:\n        img = img / 255.0\n    clr_mean = np.mean(img) if img.ndim == 2 else {i: np.mean(img[..., i]) for i in range(3)}\n    clr_std = np.std(img) if img.ndim == 2 else {i: np.std(img[..., i]) for i in range(3)}\n    return clr_mean, clr_std\n\n# os.path.join(DATASET_SMALL_FOLDER, \"train_images\")\nimages = glob.glob(os.path.join(DATASET_SMALL_FOLDER, \"train_images\", \"*.png\"))\nclr_mean_std = Parallel(n_jobs=os.cpu_count())(delayed(_color_means)(fn) for fn in tqdm(images[::10]))","metadata":{"execution":{"iopub.status.busy":"2022-07-23T17:57:03.988626Z","iopub.execute_input":"2022-07-23T17:57:03.988897Z","iopub.status.idle":"2022-07-23T17:57:11.241322Z","shell.execute_reply.started":"2022-07-23T17:57:03.988872Z","shell.execute_reply":"2022-07-23T17:57:11.240233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_color_mean = pd.DataFrame([c[0] for c in clr_mean_std]).describe()\ndisplay(img_color_mean.T)\nimg_color_std = pd.DataFrame([c[1] for c in clr_mean_std]).describe()\ndisplay(img_color_std.T)\n\nimg_color_mean = list(img_color_mean.T[\"mean\"])\nimg_color_std = list(img_color_std.T[\"mean\"])\nprint(img_color_mean, img_color_std)","metadata":{"execution":{"iopub.status.busy":"2022-07-23T17:57:11.242757Z","iopub.execute_input":"2022-07-23T17:57:11.243149Z","iopub.status.idle":"2022-07-23T17:57:11.298153Z","shell.execute_reply.started":"2022-07-23T17:57:11.243114Z","shell.execute_reply":"2022-07-23T17:57:11.297194Z"},"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) > 1_500_000_000:  # todo: for train images it was fine 4_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.LANCZOS)\n    return img","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-23T17:57:11.29961Z","iopub.execute_input":"2022-07-23T17:57:11.300198Z","iopub.status.idle":"2022-07-23T17:57:11.313628Z","shell.execute_reply.started":"2022-07-23T17:57:11.300162Z","shell.execute_reply":"2022-07-23T17:57:11.31269Z"},"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])\nprint(patient_ids)\n\n! mkdir test_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(\"test_images\", f\"{name}.png\"))\n    del img\n    gc.collect()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-23T17:57:11.31579Z","iopub.execute_input":"2022-07-23T17:57:11.316695Z","iopub.status.idle":"2022-07-23T17:59:29.452325Z","shell.execute_reply.started":"2022-07-23T17:57:11.316648Z","shell.execute_reply":"2022-07-23T17:59:29.451303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 1. Create the DataModule","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 train_input_per_sample_transform(self) -> Callable:\n        return T.Compose([\n            T.ToTensor(),\n            # T.Lambda(lambda x: (x * 255).to(torch.uint8)),\n            # T.RandomPosterize(bits=7, p=0.2),\n            # T.RandomEqualize(),\n            # T.Lambda(lambda x: x.to(torch.float32) / 255),\n            T.RandomCrop(size=(800, 800), pad_if_needed=True),\n            T.GaussianBlur(kernel_size=5, sigma=(0.5, 4)),\n            T.Resize(self.image_size),\n            T.RandomHorizontalFlip(),\n            T.RandomVerticalFlip(),\n            T.RandomAffine(degrees=30, translate=(0.15, 0.15)),\n            T.Normalize(self.color_mean, self.color_std),\n        ])\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-23T18:05:46.090574Z","iopub.execute_input":"2022-07-23T18:05:46.090982Z","iopub.status.idle":"2022-07-23T18:05:46.107057Z","shell.execute_reply.started":"2022-07-23T18:05:46.090949Z","shell.execute_reply":"2022-07-23T18:05:46.105903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train['img_name'] = df_train['image_id'].apply(lambda n: f\"{n}.png\")\ndisplay(df_train.head())\nprint(len(df_train))\nmissing = df_train['img_name'].apply(lambda n: not os.path.isfile(os.path.join(DATASET_SMALL_FOLDER, \"train_images\", n)))\ndf_train = df_train[~missing]\nprint(len(df_train))","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-23T18:05:46.109462Z","iopub.execute_input":"2022-07-23T18:05:46.109958Z","iopub.status.idle":"2022-07-23T18:05:46.494614Z","shell.execute_reply.started":"2022-07-23T18:05:46.109916Z","shell.execute_reply":"2022-07-23T18:05:46.493442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRANSFORM_PARAMS = {\n    \"image_size\": (528, 528),\n#     \"color_mean\": (0.947, 0.881, 0.863),\n#     \"color_std\": (0.093, 0.201, 0.245),\n}\n\ndatamodule = ImageClassificationData.from_data_frame(\n    \"img_name\",\n    \"label\",\n    train_data_frame=df_train,\n    train_images_root=os.path.join(DATASET_SMALL_FOLDER, \"train_images\"),\n    val_split=0.1,\n    batch_size=8,\n    train_transform=ImageClassifInputTransform,\n    val_transform=ImageClassifInputTransform,\n    transform_kwargs=TRANSFORM_PARAMS,\n    num_workers=2,\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-23T18:05:46.496937Z","iopub.execute_input":"2022-07-23T18:05:46.497306Z","iopub.status.idle":"2022-07-23T18:05:46.531116Z","shell.execute_reply.started":"2022-07-23T18:05:46.497271Z","shell.execute_reply":"2022-07-23T18:05:46.530206Z"},"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","metadata":{}},{"cell_type":"code","source":"fig, axarr = plt.subplots(ncols=2, nrows=4, figsize=(8, 16))\n\nlb_counts = [0, 0]\nfor batch in datamodule.train_dataloader():\n    print(batch.keys())\n    imgs = [np.rollaxis(im.numpy(), 0, 3) for im in batch['input']]\n    lbs = batch['target']\n    for i in range(len(imgs)):\n        lb = lbs[i]\n        j = lb_counts[lb]\n        if lb_counts[lb] >= 4:\n            continue\n        axarr[j, lb].imshow(imgs[i])\n        axarr[j, lb].set_title(f\"label: {lb}\")\n        lb_counts[lb] += 1\n    if min(lb_counts) == 4:\n        break","metadata":{"execution":{"iopub.status.busy":"2022-07-23T18:05:46.532588Z","iopub.execute_input":"2022-07-23T18:05:46.532952Z","iopub.status.idle":"2022-07-23T18:05:53.38464Z","shell.execute_reply.started":"2022-07-23T18:05:46.532917Z","shell.execute_reply":"2022-07-23T18:05:53.383631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Build the task","metadata":{}},{"cell_type":"code","source":"print(ImageClassifier.available_backbones())","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-23T18:05:53.387481Z","iopub.execute_input":"2022-07-23T18:05:53.388375Z","iopub.status.idle":"2022-07-23T18:05:53.397118Z","shell.execute_reply.started":"2022-07-23T18:05:53.388335Z","shell.execute_reply":"2022-07-23T18:05:53.396137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ImageClassifier(\n    backbone=\"efficientnet_b6\",\n    pretrained=True,\n    labels=datamodule.labels,\n    multi_label=datamodule.multi_label,\n    optimizer=\"Adamax\",\n    learning_rate=0.005,\n    lr_scheduler=(\"StepLR\", {\"step_size\": 150}),\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-23T18:05:53.398551Z","iopub.execute_input":"2022-07-23T18:05:53.399597Z","iopub.status.idle":"2022-07-23T18:06:04.665242Z","shell.execute_reply.started":"2022-07-23T18:05:53.399558Z","shell.execute_reply":"2022-07-23T18:06:04.664258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Create the trainer and finetune the model","metadata":{}},{"cell_type":"code","source":"import torch\nimport pytorch_lightning as pl\n\ntrainer = flash.Trainer(\n    max_epochs=10,\n    logger=pl.loggers.CSVLogger(save_dir='logs/'),\n    gpus=torch.cuda.device_count(),\n    precision=16 if torch.cuda.is_available() else 32,\n    accumulate_grad_batches=12,\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-23T18:06:04.666558Z","iopub.execute_input":"2022-07-23T18:06:04.666919Z","iopub.status.idle":"2022-07-23T18:06:04.676895Z","shell.execute_reply.started":"2022-07-23T18:06:04.666884Z","shell.execute_reply":"2022-07-23T18:06:04.675765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.finetune(model, datamodule=datamodule, strategy=\"no_freeze\")\n\n# Save the model!\ntrainer.save_checkpoint(\"image_classification_model.pt\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-23T18:06:04.678429Z","iopub.execute_input":"2022-07-23T18:06:04.678898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sn\n\nmetrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\ndel metrics[\"step\"]\nmetrics.set_index(\"epoch\", inplace=True)\n# display(metrics.dropna(axis=1, how=\"all\").head())\ng = sn.relplot(data=metrics, kind=\"line\")\nplt.gcf().set_size_inches(12, 4)\nplt.gca().set_yscale('log')\nplt.grid()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference 🔥!","metadata":{}},{"cell_type":"code","source":"ls_imgs_png = glob.glob(os.path.join(\"test_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":{"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,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Finish submission","metadata":{}},{"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]\n\n# fill mmissing if skipped for too large image size\nmissed = [{\"CE\": 0.6, \"LAA\": 0.4, \"patient_id\": pid} 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":{"trusted":true},"execution_count":null,"outputs":[]},{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_pred[[\"CE\", \"LAA\"]].round(6).to_csv(\"submission.csv\")\n\n!head submission.csv","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}