{"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":"# Inference example with a Cellpose model: www.cellpose.org\nThe model is based on U-Net, however rather than training it directly on bitmask targets they first convert them to \"spatial flows\" representations and train on that. This makes segmentation of dense and touching cells more reliable. For details and additional tricks they use see the paper \"Cellpose: a generalist algorithm for cellular segmentation\".\n\nTo train it I used the script provided in the cellpose repo ie: `python -m cellpose --train ...` after I converted the dataset to the input format it expects.\n\nIn inference I just submit the masks as they were returned from the model - no postprocessing","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport cv2","metadata":{"execution":{"iopub.status.busy":"2021-12-29T06:56:18.474990Z","iopub.execute_input":"2021-12-29T06:56:18.475342Z","iopub.status.idle":"2021-12-29T06:56:18.482374Z","shell.execute_reply.started":"2021-12-29T06:56:18.475285Z","shell.execute_reply":"2021-12-29T06:56:18.481016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"RESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\n\nCELL_TYPES  = {0: 'shsy5y', 1: 'astro', 2: 'cort'}\n\nTEST_PATH = \"../input/sartorius-cell-instance-segmentation/test\"","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:02:11.930345Z","iopub.execute_input":"2021-12-29T07:02:11.931140Z","iopub.status.idle":"2021-12-29T07:02:11.937788Z","shell.execute_reply.started":"2021-12-29T07:02:11.931107Z","shell.execute_reply":"2021-12-29T07:02:11.936298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install ../input/offlinewhl/numpy-1.20.0-cp37-cp37m-manylinux2010_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2021-12-29T06:57:45.561072Z","iopub.execute_input":"2021-12-29T06:57:45.561370Z","iopub.status.idle":"2021-12-29T06:58:20.416706Z","shell.execute_reply.started":"2021-12-29T06:57:45.561334Z","shell.execute_reply":"2021-12-29T06:58:20.415601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install ../input/offlinewhl/natsort-7.1.1-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2021-12-29T06:58:20.419617Z","iopub.execute_input":"2021-12-29T06:58:20.420256Z","iopub.status.idle":"2021-12-29T06:58:50.711699Z","shell.execute_reply.started":"2021-12-29T06:58:20.420207Z","shell.execute_reply":"2021-12-29T06:58:50.710552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install ../input/offlinewhl/edt-2.0.2-cp37-cp37m-manylinux2014_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2021-12-29T06:58:50.713450Z","iopub.execute_input":"2021-12-29T06:58:50.713782Z","iopub.status.idle":"2021-12-29T06:59:20.784079Z","shell.execute_reply.started":"2021-12-29T06:58:50.713736Z","shell.execute_reply":"2021-12-29T06:59:20.782951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install ../input/offlinewhl/pytorch_ranger-0.1.1-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2021-12-29T06:59:20.787187Z","iopub.execute_input":"2021-12-29T06:59:20.787580Z","iopub.status.idle":"2021-12-29T06:59:49.880796Z","shell.execute_reply.started":"2021-12-29T06:59:20.787511Z","shell.execute_reply":"2021-12-29T06:59:49.879747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install ../input/offlinewhl/torch_optimizer-0.1.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2021-12-29T06:59:49.884083Z","iopub.execute_input":"2021-12-29T06:59:49.885303Z","iopub.status.idle":"2021-12-29T07:00:19.311472Z","shell.execute_reply.started":"2021-12-29T06:59:49.885233Z","shell.execute_reply":"2021-12-29T07:00:19.310364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install ../input/offlinewhl/fastremap-1.11.1-cp37-cp37m-manylinux2014_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:00:19.313563Z","iopub.execute_input":"2021-12-29T07:00:19.313963Z","iopub.status.idle":"2021-12-29T07:00:48.855660Z","shell.execute_reply.started":"2021-12-29T07:00:19.313912Z","shell.execute_reply":"2021-12-29T07:00:48.854596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install --no-index ../input/cellposewhl/cellpose-0.7.2-py3-none-any.whl --find-links=../input/cellposewhl/","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:00:48.859623Z","iopub.execute_input":"2021-12-29T07:00:48.859908Z","iopub.status.idle":"2021-12-29T07:00:58.147737Z","shell.execute_reply.started":"2021-12-29T07:00:48.859875Z","shell.execute_reply":"2021-12-29T07:00:58.146579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Classifier","metadata":{}},{"cell_type":"code","source":"CLASSIFIER_CHK = \"../input/sartorius-resnet-34-classifier-finetuned/resnet34-finetuned.bin\"","metadata":{"execution":{"iopub.status.busy":"2021-12-29T06:49:54.234245Z","iopub.execute_input":"2021-12-29T06:49:54.234989Z","iopub.status.idle":"2021-12-29T06:49:54.268012Z","shell.execute_reply.started":"2021-12-29T06:49:54.234885Z","shell.execute_reply":"2021-12-29T06:49:54.267091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classifier = torch.load(CLASSIFIER_CHK, map_location=torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu'))\nclassifier.to(torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu'))\nclassifier.eval();","metadata":{"execution":{"iopub.status.busy":"2021-12-29T06:52:07.604460Z","iopub.execute_input":"2021-12-29T06:52:07.605258Z","iopub.status.idle":"2021-12-29T06:52:13.336643Z","shell.execute_reply.started":"2021-12-29T06:52:07.605225Z","shell.execute_reply":"2021-12-29T06:52:13.335643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get the input of the classifier\n# The process overlaps a bit with the Mask R-CNN preprocessing\n# But they are different\ndef get_image_for_classifier(image_id):\n    image_path = os.path.join(TEST_PATH, image_id + '.png')\n    transforms = A.Compose([A.Resize(224, 224), \n                       A.Normalize(mean=RESNET_MEAN, std=RESNET_STD, p=1), \n                       ToTensorV2()])\n    image = transforms(image=cv2.imread(image_path))['image']\n    return image.unsqueeze(0).to(torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu'))\n\n# Assess the image_id cell_type with the classifier\ndef get_image_cell_type(classifier, image_id):\n    img = get_image_for_classifier(image_id)\n    with torch.no_grad():\n        logits = classifier(img)[0]\n        cell_type_idx = torch.argmax(logits).item()\n    return CELL_TYPES[cell_type_idx]","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:02:16.303415Z","iopub.execute_input":"2021-12-29T07:02:16.303756Z","iopub.status.idle":"2021-12-29T07:02:16.313477Z","shell.execute_reply.started":"2021-12-29T07:02:16.303711Z","shell.execute_reply":"2021-12-29T07:02:16.312081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#get_image_cell_type(classifier,\"01ae5a43a2ab\")","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:05:24.330712Z","iopub.execute_input":"2021-12-29T07:05:24.331433Z","iopub.status.idle":"2021-12-29T07:05:24.338243Z","shell.execute_reply.started":"2021-12-29T07:05:24.331394Z","shell.execute_reply":"2021-12-29T07:05:24.337283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### I'm running it in a separate process rather than as a regular notebook because I've faced issues with numpy version not updating after the `pip install` above. If you know how to fix this please let me know.","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom cellpose import models, io, plot\nfrom pathlib import Path\nimport pandas as pd\n\ndef rle_encode(img):\n    pixels = img.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\ntest_dir = Path('../input/sartorius-cell-instance-segmentation/test')\ntest_files = [fname for fname in test_dir.iterdir()]\n","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:08:55.501722Z","iopub.execute_input":"2021-12-29T07:08:55.502109Z","iopub.status.idle":"2021-12-29T07:08:55.512096Z","shell.execute_reply.started":"2021-12-29T07:08:55.502031Z","shell.execute_reply":"2021-12-29T07:08:55.510775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shy_model = models.CellposeModel(gpu=True, pretrained_model='../input/cellpose-sep-model/shy5_4/models/cellpose_residual_on_style_on_concatenation_off_train_2021_12_28_20_56_58.474692')\nastro_model = models.CellposeModel(gpu=True, pretrained_model='../input/cellpose-sep-model/astro_4/models/cellpose_residual_on_style_on_concatenation_off_train_2021_12_28_21_24_32.062928')\ncort_model = models.CellposeModel(gpu=True, pretrained_model='../input/cellpose-sep-model/cort_4/models/cellpose_residual_on_style_on_concatenation_off_train_2021_12_28_20_18_28.265027')\n\nids, masks = [],[]\nfor count,fn in enumerate(test_files):\n    img_id = os.listdir(TEST_PATH)[count][:-4]\n    \n   # CELL_TYPES  = {0: 'shsy5y', 1: 'astro', 2: 'cort'}\n    cell_type = get_image_cell_type(classifier,img_id)\n    \n    if cell_type == \"cort\":\n        preds, flows, _ = cort_model.eval(io.imread(str(fn)), diameter= None, channels=[0,0], augment=True, resample=True,flow_threshold=0.5)\n    elif cell_type == \"shsy5y\":\n        preds, flows, _ = shy_model.eval(io.imread(str(fn)), diameter= None, channels=[0,0], augment=True, resample=True,flow_threshold=0.5)\n    elif cell_type == \"astro\":\n        preds, flows, _ = astro_model.eval(io.imread(str(fn)), diameter= None, channels=[0,0], augment=True, resample=True,flow_threshold=0.5)\n\n    for i in range (1, preds.max() + 1):\n        ids.append(fn.stem)\n        masks.append(rle_encode(preds == i))\n        \npd.DataFrame({'id':ids, 'predicted':masks}).to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:53:21.276286Z","iopub.execute_input":"2021-12-29T07:53:21.276633Z","iopub.status.idle":"2021-12-29T07:53:26.660707Z","shell.execute_reply.started":"2021-12-29T07:53:21.276585Z","shell.execute_reply":"2021-12-29T07:53:26.659674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!python run.py","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:53:26.662871Z","iopub.execute_input":"2021-12-29T07:53:26.663240Z","iopub.status.idle":"2021-12-29T07:53:26.668680Z","shell.execute_reply.started":"2021-12-29T07:53:26.663207Z","shell.execute_reply":"2021-12-29T07:53:26.667460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\npd.read_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:53:26.670416Z","iopub.execute_input":"2021-12-29T07:53:26.671036Z","iopub.status.idle":"2021-12-29T07:53:26.696424Z","shell.execute_reply.started":"2021-12-29T07:53:26.670986Z","shell.execute_reply":"2021-12-29T07:53:26.695334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rles_to_mask(encs, shape):\n    \"\"\"\n    Decodes a rle.\n\n    Args:\n        encs (list of str): Rles for each class.\n        shape (tuple [2]): Mask size.\n\n    Returns:\n        np array [shape]: Mask.\n    \"\"\"\n    img = np.zeros(shape[0] * shape[1], dtype=np.uint)\n    if type(encs)==float:\n        return img\n    for m, enc in enumerate(encs):\n        if isinstance(enc, np.float) and np.isnan(enc):\n            continue\n        enc_split = enc.split()\n        for i in range(len(enc_split) // 2):\n            start = int(enc_split[2 * i]) - 1\n            length = int(enc_split[2 * i + 1])\n            img[start: start + length] = 1 + m\n    return img.reshape(shape)","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:53:26.699064Z","iopub.execute_input":"2021-12-29T07:53:26.699409Z","iopub.status.idle":"2021-12-29T07:53:26.708279Z","shell.execute_reply.started":"2021-12-29T07:53:26.699363Z","shell.execute_reply":"2021-12-29T07:53:26.707025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"width = 704\nheight = 520\nshape = [height,width]\n\ncellpose_predictions = pd.read_csv('submission.csv')\ncellpose_predictions = cellpose_predictions.groupby('id').predicted.agg(list).reset_index()","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:53:26.710446Z","iopub.execute_input":"2021-12-29T07:53:26.710886Z","iopub.status.idle":"2021-12-29T07:53:26.729670Z","shell.execute_reply.started":"2021-12-29T07:53:26.710841Z","shell.execute_reply":"2021-12-29T07:53:26.728529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cellpose_predictions","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:53:26.733074Z","iopub.execute_input":"2021-12-29T07:53:26.733864Z","iopub.status.idle":"2021-12-29T07:53:26.748057Z","shell.execute_reply.started":"2021-12-29T07:53:26.733813Z","shell.execute_reply":"2021-12-29T07:53:26.746855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i,row in cellpose_predictions.iterrows():\n    \n    print(row.id)\n    #gt_masks = rles_to_mask(row.annotation, shape).astype(np.uint16)\n    predicted_masks = rles_to_mask(row.predicted, shape).astype(np.uint16)\n    \n    #gt_masks = (gt_masks>0).astype(int)*(gt_masks%5)\n    predicted_masks = (predicted_masks>0).astype(int)*(predicted_masks%5)\n\n    _, axs = plt.subplots(1, 2, figsize=(36, 18))\n    axs = axs.flatten()\n    image = plt.imread(f\"../input/sartorius-cell-instance-segmentation/test/{row.id}.png\",)\n    axs[0].imshow(image)\n    axs[1].imshow(predicted_masks)\n    plt.show()\n    \n    if i==4: break","metadata":{"execution":{"iopub.status.busy":"2021-12-29T07:53:26.749833Z","iopub.execute_input":"2021-12-29T07:53:26.750959Z","iopub.status.idle":"2021-12-29T07:53:29.507724Z","shell.execute_reply.started":"2021-12-29T07:53:26.750857Z","shell.execute_reply":"2021-12-29T07:53:29.506796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}