In [1]:
import numpy as np
import glob
import matplotlib.pyplot as plt
import cv2
from skimage import measure, segmentation
from vis_utils import load_volume, VolumeVisualizer, ColorMapVisualizer
from scipy.ndimage import zoom
from skimage.morphology import skeletonize, skeletonize_3d
from skimage import filters, morphology

from scipy.ndimage.filters import convolve, correlate
from scipy import signal

In [2]:
source_dir = './data/'
files = list(sorted(glob.glob(source_dir + '/*/*.raw')))
list(enumerate(files))

[(0, './data/P12/P12_60um_1333x443x864.raw')]

In [13]:
%%time
volume = load_volume(files[0], scale=0.5)
visualizer = VolumeVisualizer(volume, binary=False).visualize()

CPU times: user 7.93 s, sys: 3.88 s, total: 11.8 s
Wall time: 24.7 s


In [23]:
np.unique(volume)

array([  0,   1,   2,   3,   4,   5,   6,   7,   8,   9,  10,  11,  12,
        13,  14,  15,  16,  17,  18,  19,  20,  21,  22,  23,  24,  25,
        26,  27,  28,  29,  30,  31,  32,  33,  34,  35,  36,  37,  38,
        39,  40,  41,  42,  43,  44,  45,  46,  47,  48,  49,  50,  51,
        52,  53,  54,  55,  56,  57,  58,  59,  60,  61,  62,  63,  64,
        65,  66,  67,  68,  69,  70,  71,  72,  73,  74,  75,  76,  77,
        78,  79,  80,  81,  82,  83,  84,  85,  86,  87,  88,  89,  90,
        91,  92,  93,  94,  95,  96,  97,  98,  99, 100, 101, 102, 103,
       104, 105, 106, 107, 108, 109, 110, 111, 112, 113, 114, 115, 116,
       117, 118, 119, 120, 121, 122, 123, 124, 125, 126, 127, 128, 129,
       130, 131, 132, 133, 134, 135, 136, 137, 138, 139, 140, 141, 142,
       143, 144, 145, 146, 147, 148, 149, 150, 151, 152, 153, 154, 155,
       156, 157, 158, 159, 160, 161, 162, 163, 164, 165, 166, 167, 168,
       169, 170, 171, 172, 173, 174, 175, 176, 177, 178, 179, 18

## simple threshold segmentation

In [4]:
threshold = 70
mask_raw = volume > threshold

# VolumeVisualizer(mask_raw).visualize()

In [5]:
def get_main_regions(binary_mask, min_size=10_000, connectivity=3):
    labeled = measure.label(binary_mask, connectivity=connectivity)
    region_props = measure.regionprops(labeled)
    
    main_regions_masks = []
    
    for props in region_props:
        if props.area >= min_size:
            main_regions_masks.append((props.filled_image, props.bbox))
            
    return main_regions_masks

def merge_masks(masks, img_shape):
    result_mask = np.zeros(img_shape, dtype=np.uint8)
    for mask, bbox in masks:
        min1, min2, min3, max1, max2, max3 = bbox
        result_mask[min1:max1, min2:max2, min3:max3] += mask.astype(np.uint8)
        
    return result_mask

In [6]:
%%time
main_regions_masks = get_main_regions(mask_raw, min_size=5_000, connectivity=1)

CPU times: user 8.79 s, sys: 304 ms, total: 9.1 s
Wall time: 9.1 s


In [7]:
mask = main_regions_masks[0][0]
# VolumeVisualizer(mask).visualize()

## utility functions

In [24]:
def spherical_kernel(outer_radius, thickness=1, filled=True):    
    outer_sphere = morphology.ball(radius=outer_radius)
    if filled:
        return outer_sphere
    
    inner_radius = outer_radius - thickness
    inner_sphere = morphology.ball(radius=inner_radius)
    
    begin = outer_radius - inner_radius
    end = begin + inner_sphere.shape[0]
    outer_sphere[begin:end, begin:end, begin:end] -= inner_sphere
    return outer_sphere

def convolve_with_ball(mask, ball_radius, dtype=np.uint16):
    kernel = spherical_kernel(ball_radius, filled=True)
    return signal.convolve(mask.astype(dtype), kernel.astype(dtype), mode='same')

def get_arterial_regions(conv_img, lower_hyst_fraction, upper_hyst_fraction):
    lower_hyst_value = lower_hyst_fraction * conv_img.max()
    upper_hyst_value = upper_hyst_fraction * conv_img.max()
    return filters.apply_hysteresis_threshold(conv_img, lower_hyst_value, upper_hyst_value)

def reconstruct_from_skeleton(skeleton, ball_radius):    
    mask = np.zeros(skeleton.shape, dtype=np.uint8)
    mask = np.pad(mask, ball_radius)
    
    kernel = spherical_kernel(ball_radius, filled=True)
    central_points = np.argwhere(skeleton == 1)
    
    for central_point in central_points:
        start_corner = tuple(central_point)
        end_corner = tuple(central_point + 2*ball_radius + 1)
        
        start1, start2, start3 = start_corner
        end1, end2, end3 = end_corner
        
        mask_slice = mask[start1:end1, start2:end2, start3:end3]
        mask_slice[:] = np.logical_or(mask_slice, kernel)
                
    return mask[ball_radius:-ball_radius, ball_radius:-ball_radius, ball_radius:-ball_radius]

# high level functions

def get_tree_core(tree_mask, kernel_radius, max_fraction):
    convolved_mask = convolve_with_ball(tree_mask, kernel_radius)
    core_voxels = convolved_mask > max_fraction * convolved_mask.max()
    core_skeleton = skeletonize_3d(core_voxels.astype(np.uint8))
    core_reconstruction = reconstruct_from_skeleton(core_skeleton, kernel_radius)
    
    return core_reconstruction

def expand_tree_reconstruction(tree_mask, reconstruction, kernel_radius, max_fraction):
    convolved_mask = convolve_with_ball(tree_mask, kernel_radius)
    
    kernel_vol = spherical_kernel(kernel_radius).sum()
    threshold_value = int(max_fraction * kernel_vol)
    
    # set current reconstruction to infinity
    convolved_mask_with_huge_core = convolved_mask + reconstruction * (kernel_vol + 2)
        
        
    expanded_rec = filters.apply_hysteresis_threshold(convolved_mask_with_huge_core, threshold_value, kernel_vol + 5)
    expansion = expanded_rec - reconstruction
    
    convolved_mask_with_huge_expansion = convolved_mask + expansion * (kernel_vol + 2)
    expanded_expansion = filters.apply_hysteresis_threshold(convolved_mask_with_huge_expansion, threshold_value, kernel_vol + 5)
    
    ee_skeleton = skeletonize_3d(expanded_expansion.astype(np.uint8))
    ee_reconstruction = reconstruct_from_skeleton(ee_skeleton, kernel_radius)
    
    return ee_reconstruction, ee_skeleton

In [9]:
%%time
core_rec = get_tree_core(mask, 15, 0.95)
# VolumeVisualizer(np.logical_or(core_rec, mask)).visualize()
VolumeVisualizer(core_rec).visualize()

CPU times: user 8.91 s, sys: 2.17 s, total: 11.1 s
Wall time: 12.7 s


In [10]:
%%time
ee = expand_tree_reconstruction(mask, core_rec, kernel_radius=10, max_fraction=0.8)
VolumeVisualizer(ee).visualize()

CPU times: user 11.6 s, sys: 2.22 s, total: 13.8 s
Wall time: 15.9 s


In [25]:
%%time

total_skel = np.zeros(mask.shape)
rec = get_tree_core(mask, 15, 0.95)
total_rec = rec.copy().astype(np.uint8)
new_rec = rec.copy()
print('core is nice')

for i, kernel_radius in enumerate([10, 8, 7, 6, 5, 4, 3, 2, 1]):
    new_rec, new_ee_skel = expand_tree_reconstruction(mask, new_rec, kernel_radius=kernel_radius, max_fraction=0.5)
    rec = np.logical_or(rec, new_rec).astype(np.uint8)
    
    total_skel = np.logical_or(total_skel, new_ee_skel).astype(np.uint8)
    
    just_expansion = new_rec.copy()
    just_expansion[total_rec > 0] = 0
    total_rec += just_expansion * (i + 2)
    
    print('iter for', kernel_radius, 'ended successfully XD')
    

core is nice
iter for 10 ended successfully XD
iter for 8 ended successfully XD
iter for 7 ended successfully XD
iter for 6 ended successfully XD
iter for 5 ended successfully XD
iter for 4 ended successfully XD
iter for 3 ended successfully XD
iter for 2 ended successfully XD
iter for 1 ended successfully XD
CPU times: user 1min 51s, sys: 19.7 s, total: 2min 10s
Wall time: 2min 11s


In [14]:
VolumeVisualizer(rec).visualize()

In [15]:
ColorMapVisualizer(total_rec).visualize(interactive=False)

In [18]:
ColorMapVisualizer(mask + total_rec).visualize(interactive=True)

In [30]:
VolumeVisualizer(total_skel).visualize()

In [17]:
reconstruction_skeleton = skeletonize_3d(rec)
VolumeVisualizer(reconstruction_skeleton).visualize()

In [None]:
mask_with_no_skeleton = mask.copy()
mask_with_no_skeleton[reconstruction_skeleton == 1] = 0

VolumeVisualizer(reconstruction_skeleton.astype(np.uint8) * 2 + mask_with_no_skeleton, binary=False).visualize(primary_color=(1,1,1))

In [17]:
VolumeVisualizer(reconstruction_skeleton).visualize()

In [18]:
mask_skel = skeletonize_3d(mask.astype(np.uint8))

mask_with_no_skeleton = mask.copy()
mask_with_no_skeleton[mask_skel == 1] = 0

VolumeVisualizer(mask_skel.astype(np.uint8) * 2 + mask_with_no_skeleton, binary=False).visualize(primary_color=(1,1,1))