{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":6844262,"sourceType":"datasetVersion","datasetId":3934666},{"sourceId":6979608,"sourceType":"datasetVersion","datasetId":4009169},{"sourceId":7092658,"sourceType":"datasetVersion","datasetId":4087402},{"sourceId":7736460,"sourceType":"datasetVersion","datasetId":4521196},{"sourceId":7120688,"sourceType":"datasetVersion","datasetId":4107067},{"sourceId":7738559,"sourceType":"datasetVersion","datasetId":4521217,"isSourceIdPinned":true}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Lunit + ABMIL Inference for Ovarian Cancer","metadata":{}},{"cell_type":"markdown","source":"This is an adapted version of my competion inference notebook, replacing CLAM with ABMIL, as the open-source code of CLAM was licensed under GPL-3, but with a clause excluding commercial use.\n\nThe featuere extraction, model training and inference code are located in the dataset \"[abmil-for-ovarian-cancer](https://www.kaggle.com/datasets/dantee/abmil-for-ovarian-cancer-3rd-place)\" later referenced as ABMIL.\n\n1. Feature extraction via [Lunit-Dioo](https://github.com/lunit-io/benchmark-ssl-pathology) in the function ABMIL.extract_png_features . \n2. Training code of the [ABMIL](https://github.com/AMLab-Amsterdam/AttentionDeepMIL) model in ABMIL.main. This notebook does not run the training. It rather uses the weights from my local training run, located in the dataset .\n3. Inference code of the [ABMIL](https://github.com/AMLab-Amsterdam/AttentionDeepMIL) model is the ABMIL.eval.eval_utils and uses the model weights in the dataset [traind-abmil-for-ovarian-cancer](https://www.kaggle.com/datasets/dantee/trained--abmil-for-ovarian-cancer). There are 5 models each trained on 80% of my training data as described in the [3rd place model summary](https://www.kaggle.com/competitions/UBC-OCEAN/discussion/465527).\n   ","metadata":{}},{"cell_type":"code","source":"%load_ext autoreload\n%autoreload 2","metadata":{"execution":{"iopub.status.busy":"2024-03-01T11:59:06.149364Z","iopub.execute_input":"2024-03-01T11:59:06.149896Z","iopub.status.idle":"2024-03-01T11:59:06.180115Z","shell.execute_reply.started":"2024-03-01T11:59:06.149869Z","shell.execute_reply":"2024-03-01T11:59:06.179126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nif os.path.isdir('/kaggle/input'):\n    !cp /kaggle/input/library-fastkaggle/fastkaggle-0.0.8-py3-none-any.whl /kaggle/working \n    !pip install  /kaggle/working/fastkaggle-0.0.8-py3-none-any.whl -q\nfrom fastkaggle import setup_comp, iskaggle\n!pip install /kaggle/input/library-histomicstk/wheelhouse/histomicstk-1.3.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl --no-index --find-links /kaggle/input/library-histomicstk/wheelhouse -q\n!pip install /kaggle/input/xformers-wheel/xformers/xformers-0.0.22.post7+cu118-cp310-cp310-manylinux2014_x86_64.whl --no-index --find-links /kaggle/input/xformers-wheel/xformers","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-03-01T11:59:06.181748Z","iopub.execute_input":"2024-03-01T11:59:06.182714Z","iopub.status.idle":"2024-03-01T12:01:58.257648Z","shell.execute_reply.started":"2024-03-01T11:59:06.18269Z","shell.execute_reply":"2024-03-01T12:01:58.256671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"comp_slug = 'UBC-OCEAN'\ncomp_path = setup_comp(comp_slug)\ncomp_path","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-03-01T12:01:58.259276Z","iopub.execute_input":"2024-03-01T12:01:58.259958Z","iopub.status.idle":"2024-03-01T12:01:58.284168Z","shell.execute_reply.started":"2024-03-01T12:01:58.259915Z","shell.execute_reply":"2024-03-01T12:01:58.283209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nimport re\nimport json\nfrom pathlib import Path\nimport time\n# import psutil\nimport gc\nimport ctypes\nimport pandas as pd\nimport numpy as np\nimport matplotlib as mpl\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn.functional as F","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-03-01T12:01:58.286367Z","iopub.execute_input":"2024-03-01T12:01:58.286635Z","iopub.status.idle":"2024-03-01T12:02:00.753464Z","shell.execute_reply.started":"2024-03-01T12:01:58.286613Z","shell.execute_reply":"2024-03-01T12:02:00.752603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ['VIPS_CONCURRENCY'] = '4'\nos.environ['VIPS_DISC_THRESHOLD'] = '15gb'","metadata":{"execution":{"iopub.status.busy":"2024-03-01T12:02:00.754721Z","iopub.execute_input":"2024-03-01T12:02:00.755161Z","iopub.status.idle":"2024-03-01T12:02:00.789511Z","shell.execute_reply.started":"2024-03-01T12:02:00.755133Z","shell.execute_reply":"2024-03-01T12:02:00.788495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Make Juptyter allow multi-processing in the data loader\n# also seems to save memory\nfrom multiprocessing import set_start_method\ntry:\n    set_start_method('spawn')\nexcept RuntimeError as ex:\n    print(ex)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T12:02:00.79087Z","iopub.execute_input":"2024-03-01T12:02:00.791149Z","iopub.status.idle":"2024-03-01T12:02:01.030708Z","shell.execute_reply.started":"2024-03-01T12:02:00.791126Z","shell.execute_reply":"2024-03-01T12:02:01.028347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torch.set_num_threads(4)\n\nif iskaggle:\n    abmil_path = '/kaggle/input/abmil-for-ovarian-cancer-3rd-place'    \nelse:\n    abmil_path = '../abmil'\nsys.path.append(abmil_path)\n\nfrom ABMIL.extract_png_features import extract_png_features\nfrom ABMIL.datasets.dataset_generic import Generic_WSI_Classification_Dataset\nfrom ABMIL.utils.eval_utils import initiate_model","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-03-01T12:06:32.915672Z","iopub.execute_input":"2024-03-01T12:06:32.916458Z","iopub.status.idle":"2024-03-01T12:06:40.789958Z","shell.execute_reply.started":"2024-03-01T12:06:32.916423Z","shell.execute_reply":"2024-03-01T12:06:40.78912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ckpt_folder = Path('/kaggle/input/trained--abmil-for-ovarian-cancer')\nwith open(ckpt_folder/'settings.json', 'r') as file:\n    settings = json.load(file)\nprint(settings)\n\nmodel_path = '/kaggle/input/lunit-dino-weights/dino_vit_small_patch16_ep200.torch'\nmodel = settings['FEATURE_EXTRACT_MODEL']\nuse_fp16 = settings['USE_FP16_FOR_FEATURE_EXTRACTION']\ntile_size = settings['TILE_SIZE']\nmodel_size = settings['MODEL_SIZE']\ntma_megapixel_threshold = settings['TMA_MEGAPIXEL_THRESHOLD']\n    \n# Kaggle Submission\nif len(os.listdir(comp_path/'test_images')) != 1:\n    is_submission = True\n    test_meta = pd.read_csv(comp_path/'test.csv')\n    train_test = 'test'\n    test_folder = comp_path/f'test_images'\n# Kaggle Test Run\nelse:\n    is_submission = False\n    train_test = 'train'\n    test_meta = pd.read_csv(comp_path/f'{train_test}.csv')\n    # test_meta = test_meta[test_meta['image_id'].isin([39728, 39872, 39880, 29084, 44232, 34247, 42125, 5264])]\n    # test_meta = test_meta[test_meta['image_id'].isin([45630, 36678, 8713, 14424])] # four largest images\n    test_meta = test_meta.sample(5)\n    test_folder = comp_path/f'{train_test}_images'\nproject_root = Path('/tmp')\n\nfile_sizes = []\nformatted_sizes = []\nfor image_id in test_meta['image_id']:\n    size = os.path.getsize(test_folder/f'{image_id}.png')\n    file_sizes.append(size)\n    formatted_sizes.append(f\"{size / 1024**3:.2f} GB\")\ntest_meta['file_gb'] = np.array(file_sizes) / 1024 **3\ntest_meta['formatted_file_size'] = formatted_sizes\ntest_meta['n_mega_pixels'] = (test_meta['image_width'] * test_meta['image_height'] // 1e6).astype(int)\nn_classes = 6","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-03-01T12:14:22.813738Z","iopub.execute_input":"2024-03-01T12:14:22.814118Z","iopub.status.idle":"2024-03-01T12:14:22.927702Z","shell.execute_reply.started":"2024-03-01T12:14:22.814089Z","shell.execute_reply":"2024-03-01T12:14:22.926819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_meta.shape[0]","metadata":{"execution":{"iopub.status.busy":"2024-03-01T12:14:25.668527Z","iopub.execute_input":"2024-03-01T12:14:25.668874Z","iopub.status.idle":"2024-03-01T12:14:25.743321Z","shell.execute_reply.started":"2024-03-01T12:14:25.668849Z","shell.execute_reply":"2024-03-01T12:14:25.742375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 512\nstart = time.time()\nextract_png_features(test_meta,\n                     comp_path,\n                     project_root,\n                     train_or_test=train_test,\n                     tma_megapixel_threshold=tma_megapixel_threshold,\n                     model=model,\n                     model_path=model_path,\n                     use_fp16=use_fp16,\n                     num_workers=4,\n                     prefetch_factor=2,\n                     tile_size=tile_size,\n                     batch_size=batch_size,\n                     print_every_batches=5,\n                     tissue_threshold=0.05,\n                     gc_after_batch=False,\n                     print_memory=False,\n                     skip_existing=False)\nprint(f'Finished in {(time.time()-start) // 60:.0f} min {(time.time()-start) % 60:.1f} s')","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-03-01T12:14:28.243349Z","iopub.execute_input":"2024-03-01T12:14:28.244061Z","iopub.status.idle":"2024-03-01T12:17:11.915972Z","shell.execute_reply.started":"2024-03-01T12:14:28.244027Z","shell.execute_reply":"2024-03-01T12:17:11.914983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load Model Checkpoints","metadata":{"tags":[]}},{"cell_type":"code","source":"labels = pd.read_csv(ckpt_folder/'label_mapping.csv', header=None)\nlabels.columns=['label', 'idx']\nlabels","metadata":{"execution":{"iopub.status.busy":"2024-03-01T12:17:11.917742Z","iopub.execute_input":"2024-03-01T12:17:11.918069Z","iopub.status.idle":"2024-03-01T12:17:12.015279Z","shell.execute_reply.started":"2024-03-01T12:17:11.918043Z","shell.execute_reply":"2024-03-01T12:17:12.01431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"drop_out=0.7\nuse_inst_predictions=False\n\nre_checkpoint = re.compile(r's_\\d+_checkpoint.pt')\nckpt_files = [path for path in os.listdir(ckpt_folder) if re_checkpoint.match(path)]\nlabel_dict = {row['label']: row['idx'] for i, row in labels.iterrows()}\nmodels = [initiate_model(label_dict, \n                         'abmil', \n                         os.path.join(ckpt_folder, ckpt_file),\n                         model_size=model_size,\n                         drop_out=drop_out,\n                         feature_dim=384,\n                         use_inst_predictions=use_inst_predictions) for ckpt_file in ckpt_files]","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-03-01T12:17:32.849415Z","iopub.execute_input":"2024-03-01T12:17:32.850112Z","iopub.status.idle":"2024-03-01T12:17:33.055392Z","shell.execute_reply.started":"2024-03-01T12:17:32.850078Z","shell.execute_reply":"2024-03-01T12:17:33.054598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_meta","metadata":{"execution":{"iopub.status.busy":"2024-03-01T12:17:37.613415Z","iopub.execute_input":"2024-03-01T12:17:37.614091Z","iopub.status.idle":"2024-03-01T12:17:37.698132Z","shell.execute_reply.started":"2024-03-01T12:17:37.614057Z","shell.execute_reply":"2024-03-01T12:17:37.696904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bag_weight = 0.7\n\ndefault_prediction = 'HGSC'\npredictions = {}\n\nnot_found_count = 0\nfor i, row in test_meta.iterrows():\n    image_id = row['image_id']\n    is_tma = (row['n_mega_pixels'] <= 50)\n    feature_path = project_root/'features'/f'{image_id}.pt'\n    try:\n        features = torch.load(feature_path, map_location='cuda')\n    except FileNotFoundError:\n        not_found_count += 1\n        continue\n\n    probs = []\n    for model in models:\n        model.eval()\n        with torch.no_grad():\n            result = model(features, bag_weight, is_tma)[0]\n        probs.append(F.softmax(result, dim=1).cpu().numpy())\n    model_avg = np.array(probs).mean(axis=0)\n\n    predictions[image_id] = labels['label'].iloc[model_avg.argmax()]","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-03-01T12:17:38.421799Z","iopub.execute_input":"2024-03-01T12:17:38.422222Z","iopub.status.idle":"2024-03-01T12:17:38.598637Z","shell.execute_reply.started":"2024-03-01T12:17:38.422189Z","shell.execute_reply":"2024-03-01T12:17:38.597735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = pd.DataFrame(predictions.items(), columns=['image_id', 'label'])\nif not_found_count > 0:\n    print(f'Could not find {not_found_count} feature files.')","metadata":{"execution":{"iopub.status.busy":"2024-03-01T12:17:42.070061Z","iopub.execute_input":"2024-03-01T12:17:42.070443Z","iopub.status.idle":"2024-03-01T12:17:42.145184Z","shell.execute_reply.started":"2024-03-01T12:17:42.070404Z","shell.execute_reply":"2024-03-01T12:17:42.144031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"Replacing {predictions['label'].isnull().sum()} null labels\")\npredictions['label'].fillna(default_prediction, inplace=True)\nprint(f\"Replacing {(~(predictions['label'].isin(labels['label']))).sum()} invalid labels\")\npredictions[~predictions['label'].isin(labels['label'])] = default_prediction\nprint(f\"Fill default prediction for {(~test_meta['image_id'].isin(predictions['image_id'])).sum()} misisng image_ids.\")\nmissing_ids = test_meta[~test_meta['image_id'].isin(predictions['image_id'])][['image_id']]\nmissing_ids['label'] = default_prediction\npredictions = pd.concat([predictions, missing_ids])\npredictions.sort_values('image_id').to_csv('/kaggle/working/submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T12:17:42.998191Z","iopub.execute_input":"2024-03-01T12:17:42.998576Z","iopub.status.idle":"2024-03-01T12:17:43.105125Z","shell.execute_reply.started":"2024-03-01T12:17:42.998545Z","shell.execute_reply":"2024-03-01T12:17:43.103845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cat submission.csv","metadata":{"execution":{"iopub.status.busy":"2024-03-01T12:17:44.663161Z","iopub.execute_input":"2024-03-01T12:17:44.663588Z","iopub.status.idle":"2024-03-01T12:17:45.745194Z","shell.execute_reply.started":"2024-03-01T12:17:44.663555Z","shell.execute_reply":"2024-03-01T12:17:45.744184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# while \"submission.csv\" not in os.listdir(\"/kaggle/working\"):\n#     predictions.to_csv('/kaggle/working/submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T12:17:12.628555Z","iopub.status.idle":"2024-03-01T12:17:12.628915Z","shell.execute_reply.started":"2024-03-01T12:17:12.628723Z","shell.execute_reply":"2024-03-01T12:17:12.628736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}