{"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":"# Baseline of Mammography Breast Cancer Detection with ⚡ Flash\n\nThis is follow-up of EDA in https://www.kaggle.com/code/jirkaborovec/mammography-eda-loading-dicom","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"!mkdir -p frozen_packages\n!pip download -q \"lightning-flash[image]\" timm --dest frozen_packages --prefer-binary\n\n!rm frozen_packages/torch-*\n!ls -l frozen_packages/ | grep -e torch -e lightning","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-05T00:57:24.222408Z","iopub.execute_input":"2022-12-05T00:57:24.223634Z","iopub.status.idle":"2022-12-05T01:00:32.924654Z","shell.execute_reply.started":"2022-12-05T00:57:24.223501Z","shell.execute_reply":"2022-12-05T01:00:32.923315Z"},"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 frozen_packages/\n!pip install -q -U timm --find-links frozen_packages/\n\n! pip list | grep -e torch -e lightning\n! nvidia-smi -L","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-05T01:00:32.929101Z","iopub.execute_input":"2022-12-05T01:00:32.929404Z","iopub.status.idle":"2022-12-05T01:01:13.923655Z","shell.execute_reply.started":"2022-12-05T01:00:32.929371Z","shell.execute_reply":"2022-12-05T01:01:13.922487Z"},"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\"\nPATH_CONVERTED = \"/kaggle/input/mammography-breast-cancer-detection-png\"","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-05T01:01:13.925303Z","iopub.execute_input":"2022-12-05T01:01:13.925659Z","iopub.status.idle":"2022-12-05T01:01:26.764452Z","shell.execute_reply.started":"2022-12-05T01:01:13.925613Z","shell.execute_reply":"2022-12-05T01:01:26.763444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(os.path.join(PATH_DATASET, \"train.csv\"))\ndisplay(df_train.head())\nprint(f'cases: {len(df_train)}')","metadata":{"execution":{"iopub.status.busy":"2022-12-05T01:01:26.767276Z","iopub.execute_input":"2022-12-05T01:01:26.768132Z","iopub.status.idle":"2022-12-05T01:01:26.894054Z","shell.execute_reply.started":"2022-12-05T01:01:26.768092Z","shell.execute_reply":"2022-12-05T01:01:26.892977Z"},"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\")\nls_images = glob.glob(os.path.join(PATH_CONVERTED, \"train_images\", \"*\", \"*.png\"))\nnp.random.shuffle(ls_images)\nclr_mean_std = Parallel(n_jobs=os.cpu_count())(delayed(_color_means)(fn) for fn in tqdm(ls_images[::10000]))","metadata":{"execution":{"iopub.status.busy":"2022-12-05T01:01:26.89547Z","iopub.execute_input":"2022-12-05T01:01:26.895939Z","iopub.status.idle":"2022-12-05T01:02:06.942043Z","shell.execute_reply.started":"2022-12-05T01:01:26.895902Z","shell.execute_reply":"2022-12-05T01:02:06.940994Z"},"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\"])[0]\nimg_color_std = list(img_color_std.T[\"mean\"])[0]\nprint(img_color_mean, img_color_std)","metadata":{"execution":{"iopub.status.busy":"2022-12-05T01:02:06.943675Z","iopub.execute_input":"2022-12-05T01:02:06.944043Z","iopub.status.idle":"2022-12-05T01:02:06.99231Z","shell.execute_reply.started":"2022-12-05T01:02:06.944001Z","shell.execute_reply":"2022-12-05T01:02:06.99116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Adjust dataset","metadata":{}},{"cell_type":"code","source":"df_train['img_path'] = df_train.apply(\n    lambda r: os.path.join(str(r['patient_id']), f\"{r['image_id']}.png\"), axis=1)\n\ndisplay(df_train.head())\nprint(f\"total samples: {len(df_train)}\")\nmissing = df_train['img_path'].apply(\n    lambda n: not os.path.isfile(os.path.join(PATH_CONVERTED, \"train_images\", n)))\ndf_train = df_train[~missing]\nprint(f\"validated samples: {len(df_train)}\")","metadata":{"execution":{"iopub.status.busy":"2022-12-05T01:02:06.993963Z","iopub.execute_input":"2022-12-05T01:02:06.99432Z","iopub.status.idle":"2022-12-05T01:02:28.075696Z","shell.execute_reply.started":"2022-12-05T01:02:06.994285Z","shell.execute_reply":"2022-12-05T01:02:28.074422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Pre-scale dataset shall improve training time as the major scaling is done only once not in each epoch","metadata":{}},{"cell_type":"code","source":"from PIL import Image\nfrom joblib import Parallel, delayed\nfrom tqdm.auto import tqdm\n\nPATH_RESIZED = \"/tmp/train_images\"\n\nls_image = glob.glob(os.path.join(PATH_CONVERTED, \"train_images\", \"*\", \"*.png\"))\nprint(f\"found images: {len(ls_image)}\")\n\ndef resize_image(img_path, output_dir, img_size: int = 512):\n    img_name, _ = os.path.splitext(os.path.basename(img_path))\n    img_dir = os.path.basename(os.path.dirname(img_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\n    img = Image.open(img_path)\n    img.thumbnail((img_size, img_size))\n    img.save(png_path)\n\n\n_= Parallel(n_jobs=4)(\n    delayed(resize_image)(p_img, output_dir=PATH_RESIZED)\n    for p_img in tqdm(ls_image)\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-05T01:39:30.497846Z","iopub.execute_input":"2022-12-05T01:39:30.49823Z","iopub.status.idle":"2022-12-05T01:39:32.340428Z","shell.execute_reply.started":"2022-12-05T01:39:30.498196Z","shell.execute_reply":"2022-12-05T01:39:32.338034Z"},"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 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 = img_color_mean\n    color_std: float = img_color_std\n\n    def train_per_sample_transform(self) -> Callable:\n        return _streaming([\n#             T.RandomInvert(),\n            T.ToTensor(),\n            T.Lambda(lambda x: (x * 255).to(torch.uint8)),\n            T.RandomPosterize(bits=7, p=0.2),\n            T.Lambda(lambda x: x.to(torch.float32) / 255),\n#             T.RandomCrop(size=(800, 800), pad_if_needed=True),  # TODO\n            T.Resize(self.image_size),\n            T.GaussianBlur(kernel_size=3, sigma=(0.5, 4)),\n            T.RandomAffine(degrees=15, translate=(0.10, 0.10)),\n            T.Normalize(self.color_mean, self.color_std),\n            T.RandomHorizontalFlip(),\n            T.RandomVerticalFlip(),\n        ])\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-05T01:38:22.312529Z","iopub.execute_input":"2022-12-05T01:38:22.313614Z","iopub.status.idle":"2022-12-05T01:38:22.328742Z","shell.execute_reply.started":"2022-12-05T01:38:22.313555Z","shell.execute_reply":"2022-12-05T01:38:22.327434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRANSFORM_PARAMS = {\n    \"image_size\": (456, 456),\n}\n\ndatamodule = ImageClassificationData.from_data_frame(\n    \"img_path\",\n    \"cancer\",\n    train_data_frame=df_train,\n    train_images_root=PATH_RESIZED,\n#     train_images_root=os.path.join(PATH_CONVERTED, \"train_images\"),\n    val_split=0.1,\n    batch_size=18,\n    transform=ImageClassifInputTransform,\n    transform_kwargs=TRANSFORM_PARAMS,\n    num_workers=2,\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-05T01:39:40.157973Z","iopub.execute_input":"2022-12-05T01:39:40.158343Z","iopub.status.idle":"2022-12-05T01:40:09.591713Z","shell.execute_reply.started":"2022-12-05T01:39:40.158309Z","shell.execute_reply":"2022-12-05T01:40:09.590565Z"},"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":"fig, axarr = plt.subplots(ncols=2, nrows=4, figsize=(8, 16))\n\nspl_max = 4\nlb_counts = [0, 0]\nfor batch in datamodule.train_dataloader():\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] >= spl_max:\n            continue\n        print(imgs[i].min(), imgs[i].mean(), imgs[i].max())\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) == spl_max:\n        break","metadata":{"execution":{"iopub.status.busy":"2022-12-05T01:40:09.593736Z","iopub.execute_input":"2022-12-05T01:40:09.594202Z","iopub.status.idle":"2022-12-05T01:40:18.498363Z","shell.execute_reply.started":"2022-12-05T01:40:09.594163Z","shell.execute_reply":"2022-12-05T01:40:18.497213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Build the task\n\nsee TIMM gallery to chose your model: https://github.com/rwightman/pytorch-image-models/blob/main/results/results-imagenet.csv","metadata":{}},{"cell_type":"code","source":"from torch import nn\nimport torch.nn.functional as F\n\nclass F1_Loss(nn.Module):\n    '''Calculate F1 score. Can work with gpu tensors\n    \n    The original implmentation is written by Michal Haltuf on Kaggle.\n    \n    see: https://gist.github.com/SuperShinyEyes/dcc68a08ff8b615442e3bc6a9b55a354\n    '''\n    def __init__(self, epsilon=1e-7):\n        super().__init__()\n        self.epsilon = epsilon\n        \n    def forward(self, y_pred, y_true):\n        assert y_pred.ndim == 2\n        assert y_true.ndim == 1\n        y_true = F.one_hot(y_true, 2).to(torch.float32)\n        y_pred = F.softmax(y_pred, dim=1)\n        \n        tp = (y_true * y_pred).sum(dim=0).to(torch.float32)\n        tn = ((1 - y_true) * (1 - y_pred)).sum(dim=0).to(torch.float32)\n        fp = ((1 - y_true) * y_pred).sum(dim=0).to(torch.float32)\n        fn = (y_true * (1 - y_pred)).sum(dim=0).to(torch.float32)\n\n        precision = tp / (tp + fp + self.epsilon)\n        recall = tp / (tp + fn + self.epsilon)\n\n        f1 = 2* (precision*recall) / (precision + recall + self.epsilon)\n        f1 = f1.clamp(min=self.epsilon, max=1-self.epsilon)\n        return 1 - f1.mean()","metadata":{"execution":{"iopub.status.busy":"2022-12-05T01:40:18.500277Z","iopub.execute_input":"2022-12-05T01:40:18.500841Z","iopub.status.idle":"2022-12-05T01:40:18.514155Z","shell.execute_reply.started":"2022-12-05T01:40:18.500781Z","shell.execute_reply":"2022-12-05T01:40:18.512988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from pprint import pprint\n# pprint(vars(datamodule))\nfrom torchmetrics import F1Score\n\n# optimizer_config = {\n#     \"name\": \"AdamW\",\n#     \"lr\": 3e-4,\n#     \"warmup_prop\": 0.1,\n#     \"betas\": (0.9, 0.999),\n#     \"max_grad_norm\": 10.,\n# }\n\nmodel = ImageClassifier(\n    backbone=\"tf_efficientnet_b5_ns\",\n    metrics=F1Score(num_classes=2),\n    pretrained=True,\n    num_classes=2,\n    optimizer=(\"adamw\", {\n#         \"warmup_prop\": 0.1,\n        \"betas\": (0.9, 0.999),\n#         \"max_grad_norm\": 10.,\n    }),\n    loss_fn=F1_Loss(),\n    learning_rate=0.005,\n    lr_scheduler=(\"StepLR\", {\"step_size\": 15_000}),\n)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-05T01:42:33.804665Z","iopub.execute_input":"2022-12-05T01:42:33.805449Z","iopub.status.idle":"2022-12-05T01:42:34.759281Z","shell.execute_reply.started":"2022-12-05T01:42:33.805401Z","shell.execute_reply":"2022-12-05T01:42:34.758223Z"},"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=5,\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=20,\n#     limit_train_batches=500,\n#     limit_val_batches=100,\n)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-05T01:42:34.761679Z","iopub.execute_input":"2022-12-05T01:42:34.762304Z","iopub.status.idle":"2022-12-05T01:42:35.6001Z","shell.execute_reply.started":"2022-12-05T01:42:34.762266Z","shell.execute_reply":"2022-12-05T01:42:35.598801Z"},"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-12-05T01:42:35.603089Z","iopub.execute_input":"2022-12-05T01:42:35.603778Z","iopub.status.idle":"2022-12-05T01:54:08.029041Z","shell.execute_reply.started":"2022-12-05T01:42:35.603729Z","shell.execute_reply":"2022-12-05T01:54:08.027887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"###  Visualize progress","metadata":{}},{"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)\n# plt.gca().set_yscale('log')\nplt.grid()","metadata":{"execution":{"iopub.status.busy":"2022-12-05T01:54:08.033288Z","iopub.execute_input":"2022-12-05T01:54:08.03499Z","iopub.status.idle":"2022-12-05T01:54:08.09597Z","shell.execute_reply.started":"2022-12-05T01:54:08.03494Z","shell.execute_reply":"2022-12-05T01:54:08.094569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference 🔥!","metadata":{}},{"cell_type":"code","source":"!pip install -qU \"python-gdcm\" pydicom pylibjpeg \"opencv-python-headless\"","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-05T01:54:08.097025Z","iopub.status.idle":"2022-12-05T01:54:08.098904Z","shell.execute_reply.started":"2022-12-05T01:54:08.098637Z","shell.execute_reply":"2022-12-05T01:54:08.098663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PATH_CONVERT = \"/tmp/test_images\"\n\nls_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-05T01:54:08.100388Z","iopub.status.idle":"2022-12-05T01:54:08.101195Z","shell.execute_reply.started":"2022-12-05T01:54:08.100937Z","shell.execute_reply":"2022-12-05T01:54:08.100963Z"},"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 = 1024):\n    dicom = pydicom.dcmread(dicom_path)\n    img = apply_windowing(dicom.pixel_array, 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-05T01:54:08.102624Z","iopub.status.idle":"2022-12-05T01:54:08.103399Z","shell.execute_reply.started":"2022-12-05T01:54:08.10313Z","shell.execute_reply":"2022-12-05T01:54:08.103154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from joblib import Parallel, delayed\n\n_= Parallel(n_jobs=5)(\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-05T01:54:08.104813Z","iopub.status.idle":"2022-12-05T01:54:08.105624Z","shell.execute_reply.started":"2022-12-05T01:54:08.10534Z","shell.execute_reply":"2022-12-05T01:54:08.105365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Predictions","metadata":{}},{"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-05T01:54:08.10706Z","iopub.status.idle":"2022-12-05T01:54:08.107871Z","shell.execute_reply.started":"2022-12-05T01:54:08.107612Z","shell.execute_reply":"2022-12-05T01:54:08.107636Z"},"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    transform=ImageClassifInputTransform,\n    transform_kwargs=TRANSFORM_PARAMS,\n)\npredictions = trainer.predict(model, datamodule=dm, output=\"probabilities\")\npredictions = list(chain(*predictions))","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-05T01:54:08.109254Z","iopub.status.idle":"2022-12-05T01:54:08.110031Z","shell.execute_reply.started":"2022-12-05T01:54:08.109775Z","shell.execute_reply":"2022-12-05T01:54:08.109799Z"},"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\": pid,\n        \"image_id\": pim,\n        \"cancer\": pred[1],\n    })\n\ndf_preds = pd.DataFrame(preds)\ndisplay(df_preds.head())","metadata":{"execution":{"iopub.status.busy":"2022-12-05T01:54:08.11144Z","iopub.status.idle":"2022-12-05T01:54:08.112253Z","shell.execute_reply.started":"2022-12-05T01:54:08.111998Z","shell.execute_reply":"2022-12-05T01:54:08.112023Z"},"trusted":true},"execution_count":null,"outputs":[]}]}