{"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":"# CLAM\n\nNOTE: Some of the descriptions or images are cited from: https://github.com/mahmoodlab/CLAM\n\n![img](https://github.com/mahmoodlab/CLAM/raw/master/docs/CLAM2.jpg)\n\n## TL;DR:\n\n+ CLAM is a high-throughput and interpretable method for data efficient whole slide image (WSI) classification using slide-level labels without any ROI extraction or patch-level annotations, and is capable of handling multi-class subtyping problems. Tested on three different WSI datasets, trained models adapt to independent test cohorts of WSI resections and biopsies as well as smartphone microscopy images (photomicrographs).\n+ paper: https://arxiv.org/abs/2004.09666\n\n## How to apply CLAM on the STRIP AI dataset ?\n\n+ I prepared four notebooks for pre-process, train and inference:\n\n### pre-process\n\n+ (1) image generation: https://www.kaggle.com/code/fx6300/clam-strip-ai-image-generation\n+ (2) feature extraction: https://www.kaggle.com/code/fx6300/clam-strip-ai-feature-extraction\n\n### train\n\n+ (3) train: https://www.kaggle.com/code/fx6300/clam-strip-ai-train\n\n### inference\n\n+ (4) inference: https://www.kaggle.com/code/fx6300/clam-strip-ai-inference\n\n## How to visualize the attention generated by CLAM ?\n\n+ I prepared an example:\n  + <b>&gt; THIS NOTEBOOK &lt;</b>: https://www.kaggle.com/fx6300/clam-strip-ai-attention-heatmap\n\n## NOTE\n\n+ The source code from CLAM (https://github.com/mahmoodlab/CLAM) is licensed under GPLv3 and available for non-commercial academic purposes.","metadata":{}},{"cell_type":"code","source":"import gc\nimport cv2\nimport numpy as np \nimport pandas as pd \nimport torch\nfrom torch import nn\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import Dataset, DataLoader\nimport h5py\nfrom sklearn.model_selection import train_test_split\nfrom tqdm.notebook import tqdm\ngc.enable()","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":3.063995,"end_time":"2022-07-08T14:24:41.045696","exception":false,"start_time":"2022-07-08T14:24:37.981701","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-28T14:00:20.942452Z","iopub.execute_input":"2022-09-28T14:00:20.942847Z","iopub.status.idle":"2022-09-28T14:00:21.475482Z","shell.execute_reply.started":"2022-09-28T14:00:20.942815Z","shell.execute_reply":"2022-09-28T14:00:21.474501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(\"../input/mayo-clinic-strip-ai/train.csv\")\n_, test_df = train_test_split(train_df, test_size=0.1, random_state=42, stratify = train_df.label)\ndirs = [\"../input/mayo-clinic-strip-ai/train/\", \"../input/mayo-clinic-strip-ai/test/\"]\ntest_df","metadata":{"papermill":{"duration":0.02504,"end_time":"2022-07-08T14:24:41.073811","exception":false,"start_time":"2022-07-08T14:24:41.048771","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-28T14:00:21.478821Z","iopub.execute_input":"2022-09-28T14:00:21.479508Z","iopub.status.idle":"2022-09-28T14:00:21.506558Z","shell.execute_reply.started":"2022-09-28T14:00:21.47947Z","shell.execute_reply":"2022-09-28T14:00:21.50564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_heatmap(attention, image_id):\n    plt.imshow(cv2.cvtColor(np.array(cv2.imread(f\"../input/4096-tiles-v4/train-4096-tiles-v4/4096-tiles-v4/{image_id}.jpg\")), cv2.COLOR_BGR2RGB))\n    hm = plt.imshow(np.flip(attention, axis=0), extent=[0, 4096, 0, 4096], cmap='Reds',interpolation=\"spline16\", alpha = 0.5)\n    plt.colorbar(hm)\n    plt.title(image_id)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-28T14:00:21.508113Z","iopub.execute_input":"2022-09-28T14:00:21.508474Z","iopub.status.idle":"2022-09-28T14:00:21.515996Z","shell.execute_reply.started":"2022-09-28T14:00:21.508438Z","shell.execute_reply":"2022-09-28T14:00:21.514853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(model, dataloader):\n    model.cuda()\n    model.eval()\n    outputs = []\n    attentions = []\n    s = nn.Softmax(dim=1)\n    patient_ids = []\n    image_ids = []\n    for item in tqdm(dataloader, leave=False):\n        patient_id = item[2][0]\n        image_id = item[3][0]\n        patient_ids.append(patient_id)\n        image_ids.append(image_id)\n        try:\n            images = item[0][0].cuda().float()  \n            _, output, _, _, attention = model(images)\n            outputs.append(s(output.cpu()[:,:2])[0].detach().numpy())\n            attentions.append(attention[0].view(8, 8).cpu().detach().numpy())\n            del output, images\n        except Exception as e:\n            print(e)\n            outputs.append(s(torch.tensor([[1, 1]]).float())[0].detach().numpy())\n            attentions.append(torch.ones(8, 8).detach().cpu().numpy())\n        gc.collect()\n        torch.cuda.empty_cache()\n    return np.array(outputs), patient_ids, image_ids, attentions","metadata":{"execution":{"iopub.status.busy":"2022-09-28T14:00:21.519667Z","iopub.execute_input":"2022-09-28T14:00:21.52015Z","iopub.status.idle":"2022-09-28T14:00:21.530232Z","shell.execute_reply.started":"2022-09-28T14:00:21.52012Z","shell.execute_reply":"2022-09-28T14:00:21.529105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FeatureDataset(Dataset):\n    def __init__(self, df, data_dir, **kwargs):\n        self.df = df\n        self.data_dir = data_dir\n    def __len__(self):\n        return len(self.df)\n        \n    def __getitem__(self, index):\n        image_id = self.df.iloc[index].image_id\n        patient_id = self.df.iloc[index].patient_id\n        label = {\"CE\":0,\"LAA\":1}[self.df.iloc[index].label]\n        full_path = f\"{self.data_dir}/{self.df.iloc[index].image_id}.h5\"\n        with h5py.File(full_path,'r') as hdf5_file:\n            features = torch.stack([torch.tensor(hdf5_file[str(i)]) for i in range(64)]).view(64, 1024)\n            coords = torch.tensor([i for i in range(64)]).view(64)\n        return features, coords, patient_id, image_id, label","metadata":{"execution":{"iopub.status.busy":"2022-09-28T14:00:21.532082Z","iopub.execute_input":"2022-09-28T14:00:21.532476Z","iopub.status.idle":"2022-09-28T14:00:21.542675Z","shell.execute_reply.started":"2022-09-28T14:00:21.532438Z","shell.execute_reply":"2022-09-28T14:00:21.541666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_paths = [\n    '../input/clambaseline/model_fold1.pth',\n]\n\nprob = pd.DataFrame()\nexamples = test_df[:10]\nfor model_id, model_path in enumerate(model_paths):\n    torch.cuda.empty_cache()\n    model = torch.jit.load(model_path)\n    torch.cuda.empty_cache()\n    \n    batch_size = 1\n    test_loader = DataLoader(\n        FeatureDataset(examples, \"../input/my-features-v4/my_features-v4\"), \n        batch_size=batch_size, \n        shuffle=False, \n        num_workers=1\n    )\n    anss, ids, image_ids, attentions = predict(model, test_loader)\n    for i in range(10):\n        show_heatmap(attentions[i], image_ids[i])","metadata":{"execution":{"iopub.status.busy":"2022-09-28T14:00:21.544431Z","iopub.execute_input":"2022-09-28T14:00:21.544808Z","iopub.status.idle":"2022-09-28T14:02:01.058594Z","shell.execute_reply.started":"2022-09-28T14:00:21.544773Z","shell.execute_reply":"2022-09-28T14:02:01.057579Z"},"trusted":true},"execution_count":null,"outputs":[]}]}