{"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":"code","source":"!pip install \"../input/pycocotools/pycocotools-2.0-cp37-cp37m-linux_x86_64.whl\"\n!pip install \"../input/hpacellsegmentatorraman/HPA-Cell-Segmentation\"\n!pip install \"../input/hpapytorchzoozip/pytorch_zoo-master\"\n!pip install \"../input/localhpapackage/hpa-single-cell/\"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from copy import deepcopy\nimport os\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom cv2 import resize, INTER_NEAREST\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport torch\nfrom torch.nn import Conv2d, Sequential, ReLU, AdaptiveMaxPool2d, Flatten\nfrom tqdm.notebook import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### HPA Local Code","metadata":{}},{"cell_type":"code","source":"from hpa.data import N_CLASSES, CHANNEL_MEANS, CHANNEL_STDS\nfrom hpa.data.dataset import NEGATIVE_LABEL, load_channels\nfrom hpa.data.misc import parse_string_label, remove_empty_masks\nfrom hpa.data.transforms import ToCellMasks\nfrom hpa.infer.cells import get_cells\nfrom hpa.infer.label import *\nfrom hpa.model.bestfitting.densenet import DensenetClass\nfrom hpa.model.localizers import *\nfrom hpa.segment import HPACellSegmenter\nfrom hpa.utils.plot import *","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Inference","metadata":{}},{"cell_type":"code","source":"PROB_CUTOFF = 0.05\nMIN_AGREEMENT = 3","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMG_DIM = 1536\nDOWNSIZE_SCALE = 16\nFEATURE_MAP_DIM = int(IMG_DIM / DOWNSIZE_SCALE)\n\nFEATURE_ROI_METHOD = 'max_and_avg'\nPOSITION_ENCODING = True\nPOSITION_ENC_SHAPE = 8\nNUM_ENCODERS = 4\nEMB_DIM = 1024\nNUM_HEADS = 4\n\nif FEATURE_ROI_METHOD == 'max_and_avg':\n    cell_feature_dim = 2048\nelse:\n    cell_feature_dim = 1024\nif POSITION_ENCODING:\n    cell_feature_dim += POSITION_ENC_SHAPE * POSITION_ENC_SHAPE\nprint(f'Features extracted per cell = {cell_feature_dim}')\n\nROOT_DIR = '/kaggle/input/hpa-single-cell-image-classification'\nIMG_DIR = os.path.join(ROOT_DIR, 'test')\n\nDEVICE = 'cuda'\nMODEL_PATHS = [\n    '/kaggle/input/hparoimodels/roi15-model9.pth',\n    '/kaggle/input/hparoimodels/roi16-model9.pth',\n    '/kaggle/input/hparoimodels/roi17-model9.pth'\n]\n\nNUCLEI_PATH = '/kaggle/input/hpacellsegmentatormodelweights/dpn_unet_nuclei_v1.pth'\nCELL_PATH = '/kaggle/input/hpacellsegmentatormodelweights/dpn_unet_cell_3ch_v1.pth'","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv(os.path.join(ROOT_DIR, 'sample_submission.csv'))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_model(model_path):\n    densenet_model = DensenetClass(in_channels=4, dropout=True)\n    densenet_encoder = Sequential(densenet_model.conv1,\n                                  densenet_model.encoder2,\n                                  densenet_model.encoder3,\n                                  densenet_model.encoder4,\n                                  densenet_model.encoder5,\n                                  ReLU())\n    \n    feature_roi_pool = RoIPool(method=FEATURE_ROI_METHOD, \n                               positions=POSITION_ENCODING, \n                               tgt_shape=POSITION_ENC_SHAPE)\n    \n    upsample_fn = Upsample(scale_factor=2, mode='nearest')\n\n    model = CellTransformer(backbone=densenet_encoder,\n                            feature_roi=feature_roi_pool,\n                            num_encoders=NUM_ENCODERS,\n                            emb_dim=EMB_DIM,\n                            num_heads=NUM_HEADS,\n                            upsample=upsample_fn,\n                            cell_feature_dim=cell_feature_dim,\n                            device=DEVICE)\n    \n    model_state = torch.load(model_path, map_location=DEVICE)\n    model.load_state_dict(model_state)\n    model = model.to(DEVICE)\n    model = model.eval()\n    return model","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = [load_model(model_path) for model_path in MODEL_PATHS]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"segmenter = HPACellSegmenter(NUCLEI_PATH, CELL_PATH, device=DEVICE)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.tail(5)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def assign_cell_labels(cells, cell_probs, prob_cutoff):\n    for cell, probs in zip(cells, cell_probs):\n        cell_class_idx = np.where(probs > prob_cutoff)[0]\n        if len(cell_class_idx) == 0:\n            cell.add_prediction(NEGATIVE_LABEL, 0.5)\n        else:\n            for label_id in cell_class_idx:\n                cell.add_prediction(label_id, probs[label_id])\n    return cells","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def assign_cell_labels_ensemble(cells, ensemble_probs, prob_cutoff, min_agreement=2):\n    for cell, cell_probs in zip(cells, ensemble_probs):\n        assigned = False\n        for label_id, label_probs in enumerate(cell_probs):\n            prob_avg = label_probs.mean()\n            pred_idx, = np.where(label_probs > PROB_CUTOFF)\n            if len(pred_idx) >= min_agreement:\n                cell.add_prediction(label_id, prob_avg)\n                assigned = True\n        if not assigned:\n            cell.add_prediction(NEGATIVE_LABEL, 0.5)\n    return cells","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"normalize_fn = A.Normalize(mean=CHANNEL_MEANS, std=CHANNEL_STDS, max_pixel_value=255)\nresize_seg_fn = A.Resize(FEATURE_MAP_DIM, FEATURE_MAP_DIM, interpolation=INTER_NEAREST)\ncell_mask_fn = ToCellMasks()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\nnum_missed_cells = 0\nfor img_id, img_dim in tqdm(zip(sub_df['ID'], sub_df['ImageWidth']), total=len(sub_df)):\n\n    # load the image\n    channels = load_channels(img_id, IMG_DIR)\n    img_full = np.dstack([channels['red'], channels['green'], channels['blue'], channels['yellow']])\n    img_reduced = resize(img_full, (IMG_DIM, IMG_DIM))\n    img_shape = (img_dim, img_dim)\n    \n    # segment the cells\n    seg = segmenter(img_reduced[..., 0], img_reduced[..., 3], img_reduced[..., 2])\n    seg = resize(seg, img_shape, interpolation=INTER_NEAREST)\n    cells = get_cells(seg)\n    \n    # prep the image\n    x = normalize_fn(image=img_reduced)['image']\n    x = ToTensorV2()(image=x)['image']\n    x = x.float().to(DEVICE)\n    \n    # create the individual cell masks\n    subseg = resize_seg_fn(image=seg)['image']\n    cell_masks = cell_mask_fn(image=subseg)['image']\n    cell_masks = torch.from_numpy(cell_masks)\n    cell_masks = cell_masks.to(DEVICE)\n\n    # count the cells\n    num_cells = torch.LongTensor([len(cell_masks)])\n    num_cells = num_cells.to(DEVICE)\n    \n    # calculate the image level class probabilities\n    with torch.no_grad():\n        class_probs_ensemble = []\n        cell_probs_ensemble = []\n        for model in models:\n            logits, cell_logits = model(x.unsqueeze(0), cell_masks, num_cells, return_cells=True)\n\n            class_probs = torch.sigmoid(logits).cpu().numpy().squeeze()\n            class_probs_ensemble.append(class_probs)\n\n            cell_probs = torch.sigmoid(cell_logits).cpu().numpy()\n            cell_probs_ensemble.append(cell_probs)\n            \n    ensemble_probs = np.stack(cell_probs_ensemble).transpose((1, 2, 0))\n        \n    # identify the cells which get squashed from the segmentation resize and remove those cells\n    missed_cell_ids = set(np.unique(seg)).difference(set(np.unique(subseg)))\n    cells = [cell for cell in cells if cell.cell_id not in missed_cell_ids]\n    num_missed_cells += len(missed_cell_ids)\n    \n    # assign the cell labels\n    cells = assign_cell_labels_ensemble(cells, ensemble_probs, PROB_CUTOFF, min_agreement=MIN_AGREEMENT)\n    \n    # gather the prediction strings\n    pred_strings = [cell.get_prediction_string() for cell in cells]\n    pred_str = ' '.join(pred_strings)\n    predictions.append(pred_str)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(num_missed_cells)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df['PredictionString'] = predictions\nsub_df.to_csv('submission.csv', index=None)\nsub_df.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"get_percent_labeled_cells(cells)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tgt_img = Image.fromarray((img_reduced[..., 1]).astype(np.uint8))\nref_img = Image.fromarray((img_reduced[..., [0, 3, 2]]).astype(np.uint8))\nplot_example(ref_img, tgt_img, seg)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for class_probs in class_probs_ensemble:\n    tgt_class_idx = []\n    for i, p in enumerate(class_probs):\n        if p > PROB_CUTOFF:\n            tgt_class_idx.append(i)\n    ax = plot_predicted_probs(class_probs, tgt_class_idx)\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for cell_id, probs in enumerate(np.stack(cell_probs_ensemble).transpose((1, 2, 0))):\n    cell_class_idx = []\n    probs = probs.ravel()\n    for i, p in enumerate(probs):\n        if p > PROB_CUTOFF:\n            cell_class_idx.append(i)\n    ax = plot_predicted_probs(probs, cell_class_idx)\n\n    xticks = np.arange(1, 3 * 18 + 1, 3)\n    xtick_labels = range(18)\n    ax.set_xticks(xticks)\n    ax.set_xticklabels(xtick_labels)\n    \n    ax.set_title(f'Cell {cell_id + 1}')\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"overlay_cell_assignments(cells, seg)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}