{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Objective\n\nPost-processing steps may include thresholding the probability map to create binary objects and then locating the centroids of these objects.   This can be done using open-source connected component analysis libraries.  \n\nPreviously, I highlighted the 'connected-components-3d' library:  \nhttps://github.com/seung-lab/connected-components-3d/  \n\nThis library is efficient, easy to use, and fast.   \nHere, I introduce an alternative method using PyTorch with CUDA for enhanced performance.","metadata":{}},{"cell_type":"code","source":"try:\n    import cc3d\nexcept:\n    !pip install connected-components-3d ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-11-12T08:09:08.796591Z","iopub.execute_input":"2024-11-12T08:09:08.797433Z","iopub.status.idle":"2024-11-12T08:09:21.913150Z","shell.execute_reply.started":"2024-11-12T08:09:08.797387Z","shell.execute_reply":"2024-11-12T08:09:21.912188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cc3d\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nimport cv2\n\nimport matplotlib\nimport matplotlib.pyplot as plt\n\nfrom timeit import default_timer as timer\ndef time_to_str(t, mode='min'):\n\tif mode=='min':\n\t\tt  = int(t)/60\n\t\thr = t//60\n\t\tmin = t%60\n\t\treturn '%2d hr %02d min'%(hr,min)\n\n\telif mode=='sec':\n\t\tt   = int(t)\n\t\tmin = t//60\n\t\tsec = t%60\n\t\treturn '%2d min %02d sec'%(min,sec)\n\n\telse:\n\t\traise NotImplementedError","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-12T08:09:21.915081Z","iopub.execute_input":"2024-11-12T08:09:21.915406Z","iopub.status.idle":"2024-11-12T08:09:25.311426Z","shell.execute_reply.started":"2024-11-12T08:09:21.915371Z","shell.execute_reply":"2024-11-12T08:09:25.310468Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def np_find_connected_componet(probability, threshold):\n    num_particle_type, D, H, W = probability.shape\n    binary = probability > np.array(threshold).reshape(num_particle_type,1,1,1)\n    componet = np.zeros((num_particle_type, D, H, W), np.uint32)\n    for i in range(num_particle_type):\n        componet[i] = cc3d.connected_components(binary[i])\n    return componet\n\ndef np_find_centroid(componet):\n    centroid =[]\n    num_particle_type, D, H, W = componet.shape\n    for i in range(num_particle_type):\n        stats = cc3d.statistics(componet[i])\n        zyx=stats['centroids'][1:]\n        xyz = np.ascontiguousarray(zyx[:,::-1])\n        centroid.append(xyz)\n    return centroid","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-12T08:09:25.312712Z","iopub.execute_input":"2024-11-12T08:09:25.313233Z","iopub.status.idle":"2024-11-12T08:09:25.321822Z","shell.execute_reply.started":"2024-11-12T08:09:25.313186Z","shell.execute_reply":"2024-11-12T08:09:25.320822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#faster version using cuda\n\n#https://github.com/kornia/kornia/blob/9ccae8c297a00a35d811b5a6e4f468a1d54d17f4/kornia/contrib/connected_components.py#L7\n#https://stackoverflow.com/questions/46840707/efficiently-find-centroid-of-labelled-image-regions\n\ndef find_connected_componet(probability, threshold, max_radius = 10):\n    device = probability.device\n    probability = probability.detach()\n    num_particle_type, D, H, W = probability.shape\n    mask = probability > torch.tensor(threshold,device=device).reshape(num_particle_type,1,1,1)\n\n    # allocate the output tensors for labels\n    out = (torch.arange(D * H * W, device=device, dtype=torch.float32)+1).reshape(1,D, H, W)\n    out = out.repeat(num_particle_type,1,1,1)\n    out[~mask] = 0\n\n    out = out.reshape(num_particle_type,1,D, H, W)\n    mask = mask.reshape(num_particle_type,1,D, H, W)\n    for _ in range(max_radius):\n        out = F.max_pool3d(out, kernel_size=3, stride=1, padding=1)\n        out = torch.mul(out, mask)  # mask using element-wise multiplication\n    out = out.reshape(num_particle_type,D,H,W)\n    out = out.long()\n    componet=[]\n    for i in range(num_particle_type):\n        u, inverse = torch.unique(out[i], sorted=True, return_inverse=True)\n        componet.append(inverse)\n    componet = torch.stack(componet)\n    #plt.imshow(componet[1].data.cpu().numpy().max(0))\n    return componet\n    \ndef find_centroid(componet):\n    device = componet.device\n    num_particle_type, D, H, W = componet.shape\n    count = componet.flatten(1).max(-1)[0]+1\n    cumcount = torch.zeros(num_particle_type+1, dtype=torch.int32, device=device)\n    cumcount[1:] = torch.cumsum(count,0)\n    componet = componet+cumcount[:-1].reshape(num_particle_type,1,1,1)\n\n    # gridz, gridy, gridx = torch.meshgrid([\n    #     torch.arange(0,D,device=device),\n    #     torch.arange(0,H,device=device),\n    #     torch.arange(0,W,device=device),\n    # ],indexing='ij')\n\n    gridz = torch.arange(0, D, device=device).reshape(1,D,1,1).expand(num_particle_type,-1,H,W)\n    gridy = torch.arange(0, H, device=device).reshape(1,1,H,1).expand(num_particle_type,D,-1,W)\n    gridx = torch.arange(0, W, device=device).reshape(1,1,1,W).expand(num_particle_type,D,H,-1)\n    n  = torch.bincount(componet.flatten())\n    nx = torch.bincount(componet.flatten(),weights=gridx.flatten())\n    ny = torch.bincount(componet.flatten(),weights=gridy.flatten())\n    nz = torch.bincount(componet.flatten(),weights=gridz.flatten())\n\n    x=nx/n\n    y=ny/n\n    z=nz/n\n    xyz = torch.stack([x,y,z],1).float()\n    xyz = torch.split(xyz, count.tolist(), dim=0)\n    centroid = [xxyyzz[1:] for xxyyzz in xyz]\n    return centroid\n\n    gridz = gridz.unsqueeze(0).expand(num_particle_type,-1,-1,-1)\n    gridy = gridy.unsqueeze(0).expand(num_particle_type,-1,-1,-1)\n    gridx = gridx.unsqueeze(0).expand(num_particle_type,-1,-1,-1)\n\n    n  = torch.bincount(componet.flatten())\n    nx = torch.bincount(componet.flatten(),weights=gridx.flatten())\n    ny = torch.bincount(componet.flatten(),weights=gridy.flatten())\n    nz = torch.bincount(componet.flatten(),weights=gridz.flatten())\n\n    x=nx/n\n    y=ny/n\n    z=nz/n\n    xyz = torch.stack([x,y,z],1).float()\n    xyz = torch.split(xyz, count.tolist(), dim=0)\n    centroid = [xxyyzz[1:] for xxyyzz in xyz]\n    return centroid\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-12T08:09:25.323985Z","iopub.execute_input":"2024-11-12T08:09:25.324432Z","iopub.status.idle":"2024-11-12T08:09:25.346331Z","shell.execute_reply.started":"2024-11-12T08:09:25.324398Z","shell.execute_reply":"2024-11-12T08:09:25.345387Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#make dummy data\ndef fill_sphere(array, point, radius, color):\n    D,H,W = array.shape\n    x0, y0, z0 = point\n    for x in range(x0 - radius, x0 + radius + 1):\n        for y in range(y0 - radius, y0 + radius + 1):\n            for z in range(z0 - radius, z0 + radius + 1):\n                d = ((x0-x)**2 + (y0-y)**2 + (z0-z)**2)**0.5\n                #deb = radius - abs(x0 - x) - abs(y0 - y) - abs(z0 - z)\n                z = np.clip(z,0,D-1)\n                y = np.clip(y,0,H-1)\n                x = np.clip(x,0,W-1)\n                if d <=radius:\n                    array[z,y,x] = color\n\n\n###########################################################################\n\nD,H,W = 138,630,630\nlabel = np.zeros((D,H,W),dtype=np.int32)\nfor i in range(1,7):\n    num = np.random.randint(10,50)\n    for j in range(num):\n        radius=5\n        x = np.random.randint(2*radius,W-2*radius)\n        y = np.random.randint(2*radius,H-2*radius)\n        z = np.random.randint(2*radius,D-2*radius)\n        fill_sphere(label, (x,y,z), radius, i)\n\nprint('label', label.shape, label.max(), label.min())\nplt.imshow(label.max(0))\nplt.show()\n\nprobability=np.eye(7)[label]\nprobability = np.ascontiguousarray(probability.transpose(3,0,1,2,)).astype(np.float32)\n_7_,D,H,W = probability.shape\nthreshold=[-1,0.5,0.5,0.5,0.5,0.5,0.5]\n\nprint('probability', probability.shape, probability.max(), probability.min())\nprint('probability[1]')\nplt.imshow(probability[1].max(0))\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-12T08:09:25.347478Z","iopub.execute_input":"2024-11-12T08:09:25.348256Z","iopub.status.idle":"2024-11-12T08:09:37.517543Z","shell.execute_reply.started":"2024-11-12T08:09:25.348213Z","shell.execute_reply":"2024-11-12T08:09:37.516621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n#time np version\nnum_trial = 10\nstart_timer = timer()\nfor t in range(num_trial):\n    componet = np_find_connected_componet(probability[1:], threshold[1:])\nprint('np_find_connected_componet:', time_to_str(timer() - start_timer, 'sec'))\n\nstart_timer = timer()\nfor t in range(num_trial):\n    centroid = np_find_centroid(componet)\nprint('np_find_connected_componet:', time_to_str(timer() - start_timer, 'sec'))\n\n'''\nnp_find_connected_componet:  0 min 17 sec\nnp_find_connected_componet:  0 min 37 sec\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-12T08:09:37.518626Z","iopub.execute_input":"2024-11-12T08:09:37.518933Z","iopub.status.idle":"2024-11-12T08:10:32.121023Z","shell.execute_reply.started":"2024-11-12T08:09:37.518899Z","shell.execute_reply":"2024-11-12T08:10:32.120107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# time cuda version\n# https://www.speechmatics.com/company/articles-and-news/timing-operations-in-pytorch\n# https://auro-227.medium.com/timing-your-pytorch-code-fragments-e1a556e81f2\n\n#componet1 = torch.from_numpy(componet).cuda().int()\nprobability1 = torch.from_numpy(probability).cuda().half()\n# componet1 = find_connected_componet(probability1[1:], threshold[1:])\n# centroid1 = find_centroid(componet1)\n# centroid1 = [xyz.data.cpu().numpy() for xyz in centroid1]\n\n# warmup\nwith torch.no_grad():\n    for trial in range(3):\n        componet1 = find_connected_componet(probability1[1:], threshold[1:])\n        centroid1 = find_centroid(componet1)\n\n\nstart = torch.cuda.Event(enable_timing=True)\nend = torch.cuda.Event(enable_timing=True)\n\n\nstart.record() #-----------------------\nwith torch.no_grad():\n    for trial in range(num_trial):\n        componet1 = find_connected_componet(probability1[1:], threshold[1:], max_radius = 18)\nend.record() #-----------------------\ntorch.cuda.synchronize()\ntime_used = start.elapsed_time(end)\nprint('find_connected_componet:',time_to_str(time_used/1000, 'sec'))\ntorch. cuda. empty_cache() \n\n\nstart.record() #-----------------------\nwith torch.no_grad():\n    for trial in range(num_trial):\n        centroid1 = find_centroid(componet1)\nend.record() #-----------------------\ntorch.cuda.synchronize()\ntime_used = start.elapsed_time(end)\nprint('find_centroid:',time_to_str(time_used/1000, 'sec'))\n\n'''\nfind_connected_componet:  0 min 19 sec\nfind_centroid:  0 min 20 sec\n'''\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-12T08:17:26.839221Z","iopub.execute_input":"2024-11-12T08:17:26.840117Z","iopub.status.idle":"2024-11-12T08:18:21.880361Z","shell.execute_reply.started":"2024-11-12T08:17:26.840075Z","shell.execute_reply":"2024-11-12T08:18:21.879407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#verification\ntry:\n    centroid1 = [xyz.data.cpu().numpy() for xyz in centroid1]\nexcept:\n    pass\n\n\nprint('centroid')\n[print(xyz.shape) for xyz in centroid]\nprint('')\nprint('centroid1')\n[print(xyz.shape) for xyz in centroid1]\nprint('')\nprint('np.isclose')\n[print(np.all(np.isclose(xyz,xyz1))) for xyz,xyz1 in zip(centroid,centroid1)]\nprint('')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-12T08:18:32.448192Z","iopub.execute_input":"2024-11-12T08:18:32.448570Z","iopub.status.idle":"2024-11-12T08:18:32.456401Z","shell.execute_reply.started":"2024-11-12T08:18:32.448524Z","shell.execute_reply":"2024-11-12T08:18:32.455493Z"}},"outputs":[],"execution_count":null}]}