From 084dfb657ea3bb727214bce75786dc3c303b43d7 Mon Sep 17 00:00:00 2001 From: vqdang Date: Fri, 29 Jan 2021 18:01:56 +0000 Subject: [PATCH 01/19] UPD: refactor to optimize for speed --- metrics/stats_utils.py | 28 ++++++++++------------------ 1 file changed, 10 insertions(+), 18 deletions(-) diff --git a/metrics/stats_utils.py b/metrics/stats_utils.py index bef1a827..729cd168 100644 --- a/metrics/stats_utils.py +++ b/metrics/stats_utils.py @@ -1,9 +1,10 @@ import warnings -import numpy as np -from scipy.optimize import linear_sum_assignment import cv2 import matplotlib.pyplot as plt +import numpy as np +import scipy +from scipy.optimize import linear_sum_assignment # --------------------------Optimised for Speed @@ -407,32 +408,23 @@ def pair_coordinates(setA, setB, radius): """ # * Euclidean distance as the cost matrix - setA_tile = np.expand_dims(setA, axis=1) - setB_tile = np.expand_dims(setB, axis=0) - setA_tile = np.repeat(setA_tile, setB.shape[0], axis=1) - setB_tile = np.repeat(setB_tile, setA.shape[0], axis=0) - pair_distance = (setA_tile - setB_tile) ** 2 - # set A is row, and set B is paired against set A - pair_distance = np.sqrt(np.sum(pair_distance, axis=-1)) + pair_distance = scipy.spatial.distance.cdist(setA, setB, metric='euclidean') + print(setA.shape, setB.shape) # * Munkres pairing with scipy library # the algorithm return (row indices, matched column indices) - # if there is multiple same cost in a row, index of first occurence + # if there is multiple same cost in a row, index of first occurence # is return, thus the unique pairing is ensured indicesA, paired_indicesB = linear_sum_assignment(pair_distance) - # extract the paired cost and remove instances + # extract the paired cost and remove instances # outside of designated radius pair_cost = pair_distance[indicesA, paired_indicesB] pairedA = indicesA[pair_cost <= radius] pairedB = paired_indicesB[pair_cost <= radius] - unpairedA = [idx for idx in range(setA.shape[0]) if idx not in list(pairedA)] - unpairedB = [idx for idx in range(setB.shape[0]) if idx not in list(pairedB)] - - pairing = np.array(list(zip(pairedA, pairedB))) - unpairedA = np.array(unpairedA, dtype=np.int64) - unpairedB = np.array(unpairedB, dtype=np.int64) - + pairing = np.concatenate([pairedA[:,None], pairedB[:,None]], axis=-1) + unpairedA = np.delete(np.arange(setA.shape[0]), pairedA) + unpairedB = np.delete(np.arange(setB.shape[0]), pairedB) return pairing, unpairedA, unpairedB From 0b037a24d142dbc6bebdbf290dda56daa6f270a1 Mon Sep 17 00:00:00 2001 From: vqdang Date: Tue, 2 Mar 2021 23:07:53 +0000 Subject: [PATCH 02/19] UPD: add test code for tile pos --- infer/super_wsi.py | 424 +++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 424 insertions(+) create mode 100644 infer/super_wsi.py diff --git a/infer/super_wsi.py b/infer/super_wsi.py new file mode 100644 index 00000000..e6ca9cb6 --- /dev/null +++ b/infer/super_wsi.py @@ -0,0 +1,424 @@ +import multiprocessing as mp +from concurrent.futures import FIRST_EXCEPTION, ProcessPoolExecutor, as_completed, wait +from multiprocessing import Lock, Pool + +mp.set_start_method("spawn", True) # ! must be at top for VScode debugging + +import argparse +import glob +import json +import logging +import math +import os +import pathlib +import re +import shutil +import sys +import time +from functools import reduce +from importlib import import_module + +import cv2 +import numpy as np +import psutil +import scipy.io as sio +import torch +import torch.utils.data as data +import tqdm +from docopt import docopt + +# from dataloader.infer_loader import SerializeArray, SerializeFileList +# from misc.utils import ( +# cropping_center, +# get_bounding_box, +# log_debug, +# log_info, +# rm_n_mkdir, +# ) +# from misc.wsi_handler import get_file_handler +# from . import base + + +#### +def _remove_inst(inst_map, remove_id_list): + """Remove instances with id in remove_id_list. + + Args: + inst_map: map of instances + remove_id_list: list of ids to remove from inst_map + """ + for inst_id in remove_id_list: + inst_map[inst_map == inst_id] = 0 + return inst_map + + +#### +def _get_patch_info(img_shape, input_size, output_size): + """Get top left coordinate information of patches from original image. + + Args: + img_shape: input image shape + input_size: patch input shape + output_size: patch output shape + + """ + def flat_mesh_grid_coord(y, x): + y, x = np.meshgrid(y, x) + return np.stack([y.flatten(), x.flatten()], axis=-1) + + in_out_diff = input_size - output_size + nr_step = np.floor((img_shape - in_out_diff) / output_size) + 1 + last_output_coord = (in_out_diff // 2) + (nr_step) * output_size + # generating subpatches index from orginal + output_tl_y_list = np.arange( + in_out_diff[0] // 2, last_output_coord[0], output_size[0], dtype=np.int32 + ) + output_tl_x_list = np.arange( + in_out_diff[1] // 2, last_output_coord[1], output_size[1], dtype=np.int32 + ) + output_tl = flat_mesh_grid_coord(output_tl_y_list, output_tl_x_list) + output_br = output_tl + output_size + + input_tl = output_tl - in_out_diff // 2 + input_br = input_tl + input_size + # exclude any patch where the input exceed the range of image, + # can comment this out if do padding in reading + sel = np.any(input_br > img_shape, axis=-1) + + info_list = np.stack( + [ + np.stack([ input_tl[~sel], input_br[~sel]], axis=1), + np.stack([output_tl[~sel], output_br[~sel]], axis=1), + ], axis=1) + print(info_list.shape) + return info_list + +# info_list = _get_patch_info( +# np.array([70, 100]), +# np.array([40, 40]), +# np.array([20, 20])) +# print(info_list[:,1,0]) + +# info_list = _get_patch_info( +# np.array([100, 100]), +# np.array([80, 80]), +# np.array([60, 60])) +# print(info_list[:,0,1]) +# exit() + +def _get_tile_info(img_shape, input_size, output_size, margin_size, unit_size): + """Get top left coordinate information of patches from original image. + + Args: + img_shape: input image shape + input_size: patch input shape + output_size: patch output shape + + """ + # ! ouput tile size must be multiple of unit + assert np.sum(output_size % unit_size) == 0 + assert np.sum((margin_size*2) % unit_size) == 0 + + def flat_mesh_grid_coord(y, x): + y, x = np.meshgrid(y, x) + return np.stack([y.flatten(), x.flatten()], axis=-1) + + in_out_diff = input_size - output_size + nr_step = np.ceil((img_shape - in_out_diff) / output_size) + last_output_coord = (in_out_diff // 2) + (nr_step) * output_size + + assert np.sum(output_size % unit_size) == 0 + nr_unit_step = np.floor((img_shape - in_out_diff) / unit_size) + last_unit_output_coord = (in_out_diff // 2) + (nr_unit_step) * unit_size + + # generating subpatches index from orginal + def get_top_left_1d(axis): + o_tl_list = np.arange( + in_out_diff[axis] // 2, + last_output_coord[axis], + output_size[axis], dtype=np.int32 + ) + o_br_list = o_tl_list + output_size[axis] + o_br_list[-1] = last_unit_output_coord[axis] + # in default behavior, last pos >= last multiple of unit + # hence may cause duplication, do a check and remove if necessary + if o_br_list[-1] == o_br_list[-2]: + o_br_list = o_br_list[:-1] + o_tl_list = o_tl_list[:-1] + return o_tl_list, o_br_list + + output_tl_y_list, output_br_y_list = get_top_left_1d(axis=0) + output_tl_x_list, output_br_x_list = get_top_left_1d(axis=1) + + output_tl = flat_mesh_grid_coord(output_tl_y_list, output_tl_x_list) + output_br = flat_mesh_grid_coord(output_br_y_list, output_br_x_list) + + def get_info_stack(output_tl, output_br): + input_tl = output_tl - (in_out_diff // 2) + input_br = output_br + (in_out_diff // 2) + + info_list = np.stack( + [ + np.stack([ input_tl, input_br], axis=1), + np.stack([output_tl, output_br], axis=1), + ], axis=1) + return info_list + info_list = get_info_stack(output_tl, output_br) + + # get the fix grid tile info + # sel position not on the image boundary + sel = (output_tl[:,0] == np.min(output_tl[:,0])) + y_fix_output_tl = output_tl - np.array([margin_size[0], 0])[None,:] + y_fix_output_br = np.stack([output_tl[:,0], output_br[:,1]], axis=-1) + y_fix_output_br = y_fix_output_br + np.array([margin_size[0], 0])[None,:] + y_info_list = get_info_stack(y_fix_output_tl[~sel], y_fix_output_br[~sel]) + print(y_info_list[:,1]) + + sel = (output_tl[:,1] == np.min(output_tl[:,1])) + x_fix_output_tl = output_tl - np.array([0, margin_size[1]])[None,:] + x_fix_output_br = np.stack([output_tl[:,1], output_br[:,0]], axis=-1) + x_fix_output_br = x_fix_output_br + np.array([0, margin_size[1]])[None,:] + x_info_list = get_info_stack(x_fix_output_tl[~sel], x_fix_output_br[~sel]) + print(x_info_list[:,1]) + + return info_list + +info_list = _get_tile_info( + np.array([170, 100]), # [130, 100], [140, 100] + np.array([80, 80]), + np.array([60, 60]), + np.array([20, 20]), + np.array([20, 20]),) +# print(info_list[:,0]) +exit() + +#### +def _get_tile_patch_info( + img_shape, tile_shape, patch_input_shape, patch_output_shape +): + """Get chunk patch info. Here, chunk refers to tiles used during inference. + + Args: + img_shape: input image shape + tile_input_shape: shape of tiles used for processing + patch_input_shape: input patch shape + patch_output_shape: output patch shape + + """ + def flat_mesh_grid_coord(y, x): + y, x = np.meshgrid(y, x) + return np.stack([y.flatten(), x.flatten()], axis=-1) + + patch_diff_shape = patch_input_shape - patch_output_shape + patch_info_list = _get_patch_info(img_shape, + patch_input_shape, patch_output_shape, + drop_out_of_range=True) + + round_to_multiple = lambda x, y: np.floor(x / y) * y + # derive tile output placement as consecutive tiling with step size of 0 + # and tile output will have shape of multiple of patch_output_shape (round down) + tile_output_shape = int(tile_shape / patch_output_shape) * patch_output_shape + tile_input_shape = tile_output_shape + patch_diff_shape + tile_info_list = _get_patch_info(img_shape, + patch_input_shape, patch_output_shape, + drop_out_of_range=False) + tile_i_list = tile_info_list[:,0] + tile_o_list = tile_info_list[:,1] + + return + +#### +class InferManager(base.InferManager): + def __run_model(self, patch_top_left_list, pbar_desc): + # TODO: the cost of creating dataloader may not be cheap ? + dataset = SerializeArray( + "%s/cache_chunk.npy" % self.cache_path, + patch_top_left_list, + self.patch_input_shape, + ) + + dataloader = data.DataLoader( + dataset, + num_workers=self.nr_inference_workers, + batch_size=self.batch_size, + drop_last=False, + ) + + pbar = tqdm.tqdm( + desc=pbar_desc, + leave=True, + total=int(len(dataloader)), + ncols=80, + ascii=True, + position=0, + ) + + # run inference on input patches + accumulated_patch_output = [] + for batch_idx, batch_data in enumerate(dataloader): + sample_data_list, sample_info_list = batch_data + sample_output_list = self.run_step(sample_data_list) + sample_info_list = sample_info_list.numpy() + curr_batch_size = sample_output_list.shape[0] + sample_output_list = np.split(sample_output_list, curr_batch_size, axis=0) + sample_info_list = np.split(sample_info_list, curr_batch_size, axis=0) + sample_output_list = list(zip(sample_info_list, sample_output_list)) + accumulated_patch_output.extend(sample_output_list) + pbar.update() + pbar.close() + return accumulated_patch_output + + def __select_valid_patches(self, patch_info_list, has_output_info=True): + """Select valid patches from the list of input patch information. + + Args: + patch_info_list: patch input coordinate information + has_output_info: whether output information is given + + """ + down_sample_ratio = self.wsi_mask.shape[0] / self.wsi_proc_shape[0] + selected_indices = [] + for idx in range(patch_info_list.shape[0]): + patch_info = patch_info_list[idx] + patch_info = np.squeeze(patch_info) + # get the box at corresponding mag of the mask + if has_output_info: + output_bbox = patch_info[1] * down_sample_ratio + else: + output_bbox = patch_info * down_sample_ratio + output_bbox = np.rint(output_bbox).astype(np.int64) + # coord of the output of the patch (i.e center regions) + output_roi = self.wsi_mask[ + output_bbox[0][0] : output_bbox[1][0], + output_bbox[0][1] : output_bbox[1][1], + ] + if np.sum(output_roi) > 0: + selected_indices.append(idx) + sub_patch_info_list = patch_info_list[selected_indices] + return sub_patch_info_list + + def _parse_args(self, run_args): + """Parse command line arguments and set as instance variables.""" + for variable, value in run_args.items(): + self.__setattr__(variable, value) + # to tuple + make_shape_array = lambda x : np.array([x, x]).astype(np.int64) + self.tile_shape = make_shape_array(self.tile_shape) + self.patch_input_shape = make_shape_array(self.patch_input_shape) + self.patch_output_shape = make_shape_array(self.patch_output_shape) + return + + def _get_wsi_mask(self, wsi_handler, mask_path): + if mask_path is not None and os.path.isfile(mask_path): + wsi_mask = cv2.imread(mask_path) + wsi_mask = cv2.cvtColor(self.wsi_mask, cv2.COLOR_BGR2GRAY) + wsi_mask[wsi_mask > 0] = 1 + else: + log_info( + "WARNING: No mask found, generating mask via thresholding at 1.25x!" + ) + from skimage import morphology + + # simple method to extract tissue regions using intensity thresholding and morphological operations + def simple_get_mask(): + scaled_wsi_mag = 1.25 # ! hard coded + wsi_thumb_rgb = wsi_handler.get_full_img(read_mag=scaled_wsi_mag) + gray = cv2.cvtColor(wsi_thumb_rgb, cv2.COLOR_RGB2GRAY) + _, mask = cv2.threshold(gray, 0, 255, cv2.THRESH_OTSU) + mask = morphology.remove_small_objects( + mask == 0, min_size=16 * 16, connectivity=2 + ) + mask = morphology.remove_small_holes(mask, area_threshold=128 * 128) + mask = morphology.binary_dilation(mask, morphology.disk(16)) + return mask + + wsi_mask = np.array(simple_get_mask() > 0, dtype=np.uint8) + return wsi_mask + + def process_single_file(self, wsi_path, mask_path, output_dir): + """Process a single whole-slide image and save the results. + + Args: + wsi_path: path to input whole-slide image + msk_path: path to input mask. If not supplied, mask will be automatically generated. + output_dir: path where output will be saved + + """ + path_obj = pathlib.Path(wsi_path) + wsi_ext = path_obj.suffix + wsi_name = path_obj.stem + + # TODO: expose read mpp mode + self.wsi_handler = get_file_handler(wsi_path, backend=wsi_ext) + self.wsi_proc_shape = self.wsi_handler.get_dimensions(self.proc_mag) + # ! cache here is for legacy and to deal with esoteric internal wsi format + self.wsi_handler.prepare_reading( + read_mag=self.proc_mag, cache_path="%s/src_wsi.npy" % self.cache_path + ) + self.wsi_proc_shape = np.array(self.wsi_proc_shape[::-1]) # to Y, X + + self.wsi_mask = self._get_wsi_mask(self.wsi_handler, mask_path) + if np.sum(self.wsi_mask) == 0: + log_info("Skip due to empty mask!") + return + if self.save_mask: + cv2.imwrite("%s/mask/%s.png" % (output_dir, wsi_name), + self.wsi_mask * 255) + if self.save_thumb: + wsi_thumb_rgb = self.wsi_handler.get_full_img(read_mag=1.25) + cv2.imwrite( + "%s/thumb/%s.png" % (output_dir, wsi_name), + cv2.cvtColor(wsi_thumb_rgb, cv2.COLOR_RGB2BGR), + ) + + # * raw prediction + chunk_info_list, patch_info_list = _get_chunk_patch_info( + self.wsi_proc_shape, + chunk_input_shape, + patch_input_shape, + patch_output_shape, + ) + + return + + def process_wsi_list(self, run_args): + """Process a list of whole-slide images. + + Args: + run_args: arguments as defined in run_infer.py + + """ + self._parse_args(run_args) + + if not os.path.exists(self.cache_path): + rm_n_mkdir(self.cache_path) + + if not os.path.exists(self.output_dir + "/json/"): + rm_n_mkdir(self.output_dir + "/json/") + if self.save_thumb: + if not os.path.exists(self.output_dir + "/thumb/"): + rm_n_mkdir(self.output_dir + "/thumb/") + if self.save_mask: + if not os.path.exists(self.output_dir + "/mask/"): + rm_n_mkdir(self.output_dir + "/mask/") + + wsi_path_list = glob.glob(self.input_dir + "/*") + wsi_path_list.sort() # ensure ordering + for wsi_path in wsi_path_list[:]: + wsi_base_name = pathlib.Path(wsi_path).stem + msk_path = "%s/%s.png" % (self.input_mask_dir, wsi_base_name) + if self.save_thumb or self.save_mask: + output_file = "%s/json/%s.json" % (self.output_dir, wsi_base_name) + else: + output_file = "%s/%s.json" % (self.output_dir, wsi_base_name) + if os.path.exists(output_file): + log_info("Skip: %s" % wsi_base_name) + continue + try: + log_info("Process: %s" % wsi_base_name) + self.process_single_file(wsi_path, msk_path, self.output_dir) + log_info("Finish") + except: + logging.exception("Crash") + rm_n_mkdir(self.cache_path) # clean up all cache + return From 74ec8d838c1c05c33a70be57f52d06aff3dd3191 Mon Sep 17 00:00:00 2001 From: vqdang Date: Wed, 3 Mar 2021 00:43:08 +0000 Subject: [PATCH 03/19] UPD: fix x and add removal flag for boundary --- infer/super_wsi.py | 57 ++++++++++++++++++++++++++++++++++++++-------- 1 file changed, 48 insertions(+), 9 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index e6ca9cb6..24b02ac4 100644 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -165,21 +165,60 @@ def get_info_stack(output_tl, output_br): return info_list info_list = get_info_stack(output_tl, output_br) + br_most = np.max(output_br, axis=0) # get the fix grid tile info - # sel position not on the image boundary - sel = (output_tl[:,0] == np.min(output_tl[:,0])) y_fix_output_tl = output_tl - np.array([margin_size[0], 0])[None,:] y_fix_output_br = np.stack([output_tl[:,0], output_br[:,1]], axis=-1) y_fix_output_br = y_fix_output_br + np.array([margin_size[0], 0])[None,:] + # bound reassignment + y_fix_output_br[y_fix_output_br[:,0] > br_most[0], 0] = br_most[0] + y_fix_output_br[y_fix_output_br[:,1] > br_most[1], 1] = br_most[1] + # sel position not on the image boundary + sel = (output_tl[:,0] == np.min(output_tl[:,0])) y_info_list = get_info_stack(y_fix_output_tl[~sel], y_fix_output_br[~sel]) - print(y_info_list[:,1]) - - sel = (output_tl[:,1] == np.min(output_tl[:,1])) - x_fix_output_tl = output_tl - np.array([0, margin_size[1]])[None,:] - x_fix_output_br = np.stack([output_tl[:,1], output_br[:,0]], axis=-1) - x_fix_output_br = x_fix_output_br + np.array([0, margin_size[1]])[None,:] + # print(y_info_list[...,::-1][:,1],'\n') + + # flag horizontal ambiguous region for y (left margin, right margin) + # |----|------------|----| + # |\\\\| |\\\\| + # |----|------------|----| + # ambiguous ambiguous (margin size) + removal_flag = np.zeros((y_info_list.shape[0], 4,)) # left, right, top, bot + removal_flag[:,[0,1]] = 1 + # exclude the left most boundary + removal_flag[(y_info_list[:,1,0,1] == np.min(output_tl[:,1])),0] = 0 + # exclude the right most boundary + removal_flag[(y_info_list[:,1,1,1] == np.max(output_br[:,1])),1] = 0 + print(removal_flag) + print(y_info_list[...,::-1][:,1]) + + x_fix_output_br = output_br + np.array([0, margin_size[1]])[None,:] + x_fix_output_tl = np.stack([output_tl[:,0], output_br[:,1]], axis=-1) + x_fix_output_tl = x_fix_output_tl - np.array([0, margin_size[1]])[None,:] + # bound reassignment + x_fix_output_br[x_fix_output_br[:,0] > br_most[0], 0] = br_most[0] + x_fix_output_br[x_fix_output_br[:,1] > br_most[1], 1] = br_most[1] + # sel position not on the image boundary + sel = (output_br[:,1] == np.max(output_br[:,1])) x_info_list = get_info_stack(x_fix_output_tl[~sel], x_fix_output_br[~sel]) - print(x_info_list[:,1]) + # print(x_info_list[...,::-1][:,1],'\n') + # flag vertical ambiguous region for x (top margin, bottom margin) + # |----| + # |\\\\| ambiguous + # |----| + # | | + # | | + # |----| + # |\\\\| ambiguous + # |----| + removal_flag = np.zeros((x_info_list.shape[0], 4,)) # left, right, top, bot + removal_flag[:,[2,3]] = 1 + # exclude the left most boundary + removal_flag[(x_info_list[:,1,0,0] == np.min(output_tl[:,0])),2] = 0 + # exclude the right most boundary + removal_flag[(x_info_list[:,1,1,0] == np.max(output_br[:,0])),3] = 0 + print(removal_flag) + print(x_info_list[...,::-1][:,1]) return info_list From 201e63cd8f940036929b34192860507265569995 Mon Sep 17 00:00:00 2001 From: vqdang Date: Wed, 3 Mar 2021 01:04:43 +0000 Subject: [PATCH 04/19] UPD: add removal flag for normal tiling --- infer/super_wsi.py | 22 ++++++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index 24b02ac4..6acabdfe 100644 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -164,7 +164,28 @@ def get_info_stack(output_tl, output_br): ], axis=1) return info_list info_list = get_info_stack(output_tl, output_br) + # flag surrounding ambiguous (left margin, right margin) + # |----|------------|----| + # |\\\\\\\\\\\\\\\\\\\\\\| + # |\\\\ \\\\| + # |\\\\ \\\\| + # |\\\\ \\\\| + # |\\\\\\\\\\\\\\\\\\\\\\| + # |----|------------|----| + removal_flag = np.full((info_list.shape[0], 4,), 1) # left, right, top, bot + # exclude those contain left most boundary + removal_flag[(info_list[:,1,0,1] == np.min(output_tl[:,1])),0] = 0 + # exclude those contain right most boundary + removal_flag[(info_list[:,1,1,1] == np.max(output_br[:,1])),1] = 0 + # exclude those contain top most boundary + removal_flag[(info_list[:,1,0,0] == np.min(output_tl[:,0])),2] = 0 + # exclude those contain bot most boundary + removal_flag[(info_list[:,1,1,0] == np.max(output_br[:,0])),3] = 0 + print(removal_flag) + print(info_list[...,::-1][:,1]) + exit() + # * ------------------------------- br_most = np.max(output_br, axis=0) # get the fix grid tile info y_fix_output_tl = output_tl - np.array([margin_size[0], 0])[None,:] @@ -220,6 +241,7 @@ def get_info_stack(output_tl, output_br): print(removal_flag) print(x_info_list[...,::-1][:,1]) + # * define the tile cross section return info_list info_list = _get_tile_info( From c664a5b93b694ce3381d15cff8c1a0ea2459ecf9 Mon Sep 17 00:00:00 2001 From: vqdang Date: Wed, 3 Mar 2021 20:34:31 +0000 Subject: [PATCH 05/19] UPD: test persistent worker --- infer/super_wsi.py | 394 ++++++++++++++++++++++++++++++--------------- 1 file changed, 265 insertions(+), 129 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index 6acabdfe..e59fd983 100644 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -1,6 +1,5 @@ import multiprocessing as mp from concurrent.futures import FIRST_EXCEPTION, ProcessPoolExecutor, as_completed, wait -from multiprocessing import Lock, Pool mp.set_start_method("spawn", True) # ! must be at top for VScode debugging @@ -23,22 +22,45 @@ import psutil import scipy.io as sio import torch +import torch.multiprocessing as torch_mp import torch.utils.data as data import tqdm from docopt import docopt -# from dataloader.infer_loader import SerializeArray, SerializeFileList -# from misc.utils import ( -# cropping_center, -# get_bounding_box, -# log_debug, -# log_info, -# rm_n_mkdir, -# ) -# from misc.wsi_handler import get_file_handler -# from . import base +from misc.utils import ( + cropping_center, + get_bounding_box, + log_debug, + log_info, + rm_n_mkdir, +) +from misc.wsi_handler import get_file_handler +from . import base +#### +class SerializeArray(data.Dataset): + def __init__(self, mp_shared_space, preproc=None): + super().__init__() + self.mp_shared_space = mp_shared_space + + self.preproc = preproc + return + + def __len__(self): + return len(self.mp_shared_space.patch_info_list) + + def __getitem__(self, idx): + patch_info = self.mp_shared_space.patch_info_list[idx] + patch_data = self.mp_shared_space.tile_img[ + patch_info[0,0] : patch_info[1,0], + patch_info[0,1] : patch_info[1,1], + ] + if self.preproc is not None: + patch_dat = patch_data.copy() + patch_data = self.preproc(patch_data) + return patch_data, patch_info + #### def _remove_inst(inst_map, remove_id_list): """Remove instances with id in remove_id_list. @@ -51,7 +73,6 @@ def _remove_inst(inst_map, remove_id_list): inst_map[inst_map == inst_id] = 0 return inst_map - #### def _get_patch_info(img_shape, input_size, output_size): """Get top left coordinate information of patches from original image. @@ -90,7 +111,7 @@ def flat_mesh_grid_coord(y, x): np.stack([ input_tl[~sel], input_br[~sel]], axis=1), np.stack([output_tl[~sel], output_br[~sel]], axis=1), ], axis=1) - print(info_list.shape) + # print(info_list.shape) return info_list # info_list = _get_patch_info( @@ -163,7 +184,8 @@ def get_info_stack(output_tl, output_br): np.stack([output_tl, output_br], axis=1), ], axis=1) return info_list - info_list = get_info_stack(output_tl, output_br) + info_list = get_info_stack(output_tl, output_br).astype(np.int64) + # flag surrounding ambiguous (left margin, right margin) # |----|------------|----| # |\\\\\\\\\\\\\\\\\\\\\\| @@ -181,17 +203,19 @@ def get_info_stack(output_tl, output_br): removal_flag[(info_list[:,1,0,0] == np.min(output_tl[:,0])),2] = 0 # exclude those contain bot most boundary removal_flag[(info_list[:,1,1,0] == np.max(output_br[:,0])),3] = 0 - print(removal_flag) - print(info_list[...,::-1][:,1]) - exit() + # print(removal_flag) + # print(info_list[...,::-1][:,1]) + # exit() - # * ------------------------------- br_most = np.max(output_br, axis=0) + tl_most = np.min(output_tl, axis=0) + # * ------------------------------- # get the fix grid tile info y_fix_output_tl = output_tl - np.array([margin_size[0], 0])[None,:] y_fix_output_br = np.stack([output_tl[:,0], output_br[:,1]], axis=-1) y_fix_output_br = y_fix_output_br + np.array([margin_size[0], 0])[None,:] # bound reassignment + # ? do we need to do bound check for tl ? (extreme case of 1 tile of size < margin size ?) y_fix_output_br[y_fix_output_br[:,0] > br_most[0], 0] = br_most[0] y_fix_output_br[y_fix_output_br[:,1] > br_most[1], 1] = br_most[1] # sel position not on the image boundary @@ -210,8 +234,8 @@ def get_info_stack(output_tl, output_br): removal_flag[(y_info_list[:,1,0,1] == np.min(output_tl[:,1])),0] = 0 # exclude the right most boundary removal_flag[(y_info_list[:,1,1,1] == np.max(output_br[:,1])),1] = 0 - print(removal_flag) - print(y_info_list[...,::-1][:,1]) + # print(removal_flag) + # print(y_info_list[...,::-1][:,1]) x_fix_output_br = output_br + np.array([0, margin_size[1]])[None,:] x_fix_output_tl = np.stack([output_tl[:,0], output_br[:,1]], axis=-1) @@ -238,98 +262,147 @@ def get_info_stack(output_tl, output_br): removal_flag[(x_info_list[:,1,0,0] == np.min(output_tl[:,0])),2] = 0 # exclude the right most boundary removal_flag[(x_info_list[:,1,1,0] == np.max(output_br[:,0])),3] = 0 - print(removal_flag) - print(x_info_list[...,::-1][:,1]) + # print(removal_flag) + # print(x_info_list[...,::-1][:,1]) # * define the tile cross section + sel = np.any(output_br == br_most, axis=-1) + xsect = output_br[~sel] + xsect_tl = xsect - margin_size * 2 + xsect_br = xsect + margin_size * 2 + # do the bound check to ensure range stay within + xsect_br[xsect_br[:,0] > br_most[0], 0] = br_most[0] + xsect_br[xsect_br[:,1] > br_most[1], 1] = br_most[1] + xsect_tl[xsect_tl[:,0] < tl_most[0], 0] = tl_most[0] + xsect_tl[xsect_tl[:,1] < tl_most[1], 1] = tl_most[1] + xsect_info = get_info_stack(xsect_tl, xsect_br) + return info_list -info_list = _get_tile_info( - np.array([170, 100]), # [130, 100], [140, 100] - np.array([80, 80]), - np.array([60, 60]), - np.array([20, 20]), - np.array([20, 20]),) +# info_list = _get_tile_info( +# np.array([170, 100]), # [130, 100], [140, 100] +# np.array([80, 80]), +# np.array([60, 60]), +# np.array([20, 20]), +# np.array([20, 20]),) # print(info_list[:,0]) -exit() - -#### -def _get_tile_patch_info( - img_shape, tile_shape, patch_input_shape, patch_output_shape -): - """Get chunk patch info. Here, chunk refers to tiles used during inference. - - Args: - img_shape: input image shape - tile_input_shape: shape of tiles used for processing - patch_input_shape: input patch shape - patch_output_shape: output patch shape - - """ - def flat_mesh_grid_coord(y, x): - y, x = np.meshgrid(y, x) - return np.stack([y.flatten(), x.flatten()], axis=-1) - - patch_diff_shape = patch_input_shape - patch_output_shape - patch_info_list = _get_patch_info(img_shape, - patch_input_shape, patch_output_shape, - drop_out_of_range=True) - - round_to_multiple = lambda x, y: np.floor(x / y) * y - # derive tile output placement as consecutive tiling with step size of 0 - # and tile output will have shape of multiple of patch_output_shape (round down) - tile_output_shape = int(tile_shape / patch_output_shape) * patch_output_shape - tile_input_shape = tile_output_shape + patch_diff_shape - tile_info_list = _get_patch_info(img_shape, - patch_input_shape, patch_output_shape, - drop_out_of_range=False) - tile_i_list = tile_info_list[:,0] - tile_o_list = tile_info_list[:,1] - - return +# exit() #### -class InferManager(base.InferManager): - def __run_model(self, patch_top_left_list, pbar_desc): - # TODO: the cost of creating dataloader may not be cheap ? - dataset = SerializeArray( - "%s/cache_chunk.npy" % self.cache_path, - patch_top_left_list, - self.patch_input_shape, - ) - - dataloader = data.DataLoader( - dataset, - num_workers=self.nr_inference_workers, - batch_size=self.batch_size, - drop_last=False, - ) - - pbar = tqdm.tqdm( - desc=pbar_desc, - leave=True, - total=int(len(dataloader)), - ncols=80, - ascii=True, - position=0, - ) +def run_model( + mp_queue_output, + tile_info_list, + wsi_path, wsi_ext, wsi_proc_mag,): + + # wsi_handler = get_file_handler(wsi_path, backend=wsi_ext) + # # ! cache here is for legacy and to deal with esoteric internal wsi format + # wsi_handler.prepare_reading( + # read_mag=self.proc_mag, cache_path="%s/src_wsi.npy" % self.cache_path + # ) + + # using shared memory namespace so all the loader workers use same + # underlying image data, also allowing persistent worker and fast data switching + mp_manager = torch_mp.Manager() + mp_shared_space = mp_manager.Namespace() + + ds = SerializeArray(mp_shared_space) + loader = data.DataLoader(ds, + num_workers=32, + batch_size=16, + drop_last=False, + persistent_workers=True, + ) + + for i in range (1000000): + patch_info_list = np.array([[[0, 0], [180, 180]]]*256, dtype=np.int32) + tile_img = np.full([5000, 5000, 3], i) + mp_shared_space.tile_img = torch.from_numpy(tile_img).share_memory_() + mp_shared_space.patch_info_list = torch.from_numpy(patch_info_list).share_memory_() + + # change the data without changing the loader + # so that the loader multiproc stays alives + for batch_idx, batch_data in enumerate(loader): + sample_data_list, sample_info_list = batch_data + print(sample_data_list[:,0,0,0]) + break + + while len(mp_queue_input > 0): + tile_info, patch_info_list = mp_queue_input.pop() + tile_tl, tile_br = tile_info + tile_img = wsi_handler.read_region( + tile_tl[::-1], (tile_br - tile_tl)[::-1] + ) + # change the data without changing the loader + # so that the loader multiproc stays alives - # run inference on input patches accumulated_patch_output = [] - for batch_idx, batch_data in enumerate(dataloader): + for batch_idx, batch_data in enumerate(loader): sample_data_list, sample_info_list = batch_data - sample_output_list = self.run_step(sample_data_list) + sample_output_list = run_step(sample_data_list) sample_info_list = sample_info_list.numpy() curr_batch_size = sample_output_list.shape[0] sample_output_list = np.split(sample_output_list, curr_batch_size, axis=0) sample_info_list = np.split(sample_info_list, curr_batch_size, axis=0) sample_output_list = list(zip(sample_info_list, sample_output_list)) accumulated_patch_output.extend(sample_output_list) - pbar.update() + mp_queue_output.append(accumulated_patch_output) + run inference on input patches + return + +def mp_dispatcher(data_list, nr_worker=0, show_pbar=True): + """ + data_list is alist of [[func, arg1, arg2, etc.]] + Resutls are alway sorted according to source position + """ + if nr_worker > 0: + proc_pool = ProcessPoolExecutor(nr_worker) + + result_list = [] + future_list = [] + + if show_pbar: + pbar = tqdm(total=len(data_list), ascii=True, position=0) + for run_idx, dat in enumerate(data_list): + func = dat[0] + args = dat[1:] + if nr_worker > 0: + future = proc_pool.submit(func, run_idx, *args) + future_list.append(future) + else: + # ! assume 1st return is alwasy run_id + result = func(run_idx, *args) + result_list.append(result) + if show_pbar: + pbar.update() + if nr_worker > 0: + for future in as_completed(future_list): + if future.exception() is not None: + logging.info(future.exception()) + else: + result = future.result() + result_list.append(result) + if show_pbar: + pbar.update() + if show_pbar: pbar.close() - return accumulated_patch_output + + result_list = sorted(result_list, key=lambda k : k[0]) + result_list = [v[1:] for v in result_list] + return result_list + +#### +def check_valid(run_idx, info, wsi_mask): + output_bbox = np.rint(info[1]).astype(np.int64) + output_roi = wsi_mask[ + output_bbox[0][0] : output_bbox[1][0], + output_bbox[0][1] : output_bbox[1][1], + ] + return run_idx, (torch.sum(output_roi) > 0).item() - def __select_valid_patches(self, patch_info_list, has_output_info=True): +#### +class InferManager(base.InferManager): + + def __select_valid_patches(self, patch_info_list): """Select valid patches from the list of input patch information. Args: @@ -337,26 +410,23 @@ def __select_valid_patches(self, patch_info_list, has_output_info=True): has_output_info: whether output information is given """ + down_sample_ratio = self.wsi_mask.shape[0] / self.wsi_proc_shape[0] - selected_indices = [] - for idx in range(patch_info_list.shape[0]): - patch_info = patch_info_list[idx] - patch_info = np.squeeze(patch_info) - # get the box at corresponding mag of the mask - if has_output_info: - output_bbox = patch_info[1] * down_sample_ratio - else: - output_bbox = patch_info * down_sample_ratio - output_bbox = np.rint(output_bbox).astype(np.int64) - # coord of the output of the patch (i.e center regions) - output_roi = self.wsi_mask[ - output_bbox[0][0] : output_bbox[1][0], - output_bbox[0][1] : output_bbox[1][1], - ] - if np.sum(output_roi) > 0: - selected_indices.append(idx) - sub_patch_info_list = patch_info_list[selected_indices] - return sub_patch_info_list + torch_mask = torch.from_numpy(self.wsi_mask).share_memory_() + run_list = [[check_valid, info * down_sample_ratio, torch_mask] for info in patch_info_list] + # somehow multiproc is slower than single thread + valid_indices = mp_dispatcher(run_list, nr_worker=0, show_pbar=False) + valid_indices = np.concatenate(valid_indices, axis=0) + return patch_info_list[valid_indices] + + def __select_patches_in_tile(self, tile_info, patch_info_list): + # checking basing on the output alignment + tile_tl, tile_br = tile_info[1] + patch_tl_list = patch_info_list[:,1,0] + patch_br_list = patch_info_list[:,1,1] + sel = (patch_tl_list[:,0] == tile_tl[0]) | (patch_tl_list[:,1] == tile_tl[1]) + sel |= (patch_br_list[:,0] == tile_br[0]) | (patch_br_list[:,1] == tile_br[1]) + return patch_info_list[sel] def _parse_args(self, run_args): """Parse command line arguments and set as instance variables.""" @@ -365,6 +435,7 @@ def _parse_args(self, run_args): # to tuple make_shape_array = lambda x : np.array([x, x]).astype(np.int64) self.tile_shape = make_shape_array(self.tile_shape) + self.ambiguous_size = make_shape_array(self.ambiguous_size) self.patch_input_shape = make_shape_array(self.patch_input_shape) self.patch_output_shape = make_shape_array(self.patch_output_shape) return @@ -432,14 +503,76 @@ def process_single_file(self, wsi_path, mask_path, output_dir): cv2.cvtColor(wsi_thumb_rgb, cv2.COLOR_RGB2BGR), ) - # * raw prediction - chunk_info_list, patch_info_list = _get_chunk_patch_info( - self.wsi_proc_shape, - chunk_input_shape, - patch_input_shape, - patch_output_shape, + # * retrieve patch and tile placement + patch_info_list = _get_patch_info( + self.wsi_proc_shape, self.patch_input_shape, self.patch_output_shape, + ) + patch_diff_shape = self.patch_input_shape - self.patch_output_shape + # derive tile output placement as consecutive tiling with step size of 0 + # and tile output will have shape of multiple of patch_output_shape (round down) + tile_output_shape = np.floor(self.tile_shape / self.patch_output_shape) * self.patch_output_shape + tile_input_shape = tile_output_shape + patch_diff_shape + tile_info_list = _get_tile_info( + self.wsi_proc_shape, tile_input_shape, tile_output_shape, + self.ambiguous_size, self.patch_output_shape ) + # * Async Inference + # * launch a seperate process to do forward and store the future, + # * the loader_workers may not need, to be spawned by the forward processs + # * then while polling for forward result, launch separate process for + # * doing the postproc + # + # / forward \ (loop) + # main ------------- main----------------main + # \ postproc (loop)/ + + # future_list = collections.deque() + # mp_pool = ProcessPoolExecutor(self.nr_post_proc_workers + 1) + + patch_info_list = self.__select_valid_patches(patch_info_list) + tile_info_list = self.__select_valid_patches(tile_info_list) + + mp_manager = mp.Manager() + # contain at most 5 tile ouput before polling + mp_forward_o_queue = mp_manager.Queue(maxsize=5) + + import collections + forward_info_list = collections.deque() + + nr_tile = tile_info_list.shape[0] + for tile_info in tile_info_list: + # retrieve valid patch within tile + patch_in_tile_info_list = self.__select_patches_in_tile(tile_info, patch_info_list) + # shifting patch coord from wrt wsi to wrt to tile + tile_input_info = tile_info[0] + patch_input_info_list = patch_in_tile_info_list[:,0] + patch_input_info_list -= tile_input_info + forward_info_list.append([tile_input_info, patch_input_info_list]) + + forward_process = mp.Process(target=run_model, + args=(mp_forward_o_queue, forward_info_list, + wsi_path, wsi_ext, self.proc_mag)) + + forward_process.start() + forward_process.join() + + # do double queue polling + # while len(future_list) > 0: + # while not future_list[0].done(): + # future_list.rotate() + # proc_future = future_list.popleft() + # if proc_future.exception() is not None: + # print(proc_future.exception()) + # proc_run_id, proc_result = proc_future.result() + # if proc_run_id == 'forward': + # proc_future = postproc_mp_pool.submit() + # future_list.append(proc_future) + # elif proc_run_id == 'postproc': + # proc_future.result() + # else: + # assert False + return def process_wsi_list(self, run_args): @@ -465,21 +598,24 @@ def process_wsi_list(self, run_args): wsi_path_list = glob.glob(self.input_dir + "/*") wsi_path_list.sort() # ensure ordering - for wsi_path in wsi_path_list[:]: + for wsi_path in wsi_path_list[::-1]: wsi_base_name = pathlib.Path(wsi_path).stem msk_path = "%s/%s.png" % (self.input_mask_dir, wsi_base_name) if self.save_thumb or self.save_mask: output_file = "%s/json/%s.json" % (self.output_dir, wsi_base_name) else: output_file = "%s/%s.json" % (self.output_dir, wsi_base_name) - if os.path.exists(output_file): - log_info("Skip: %s" % wsi_base_name) - continue - try: - log_info("Process: %s" % wsi_base_name) - self.process_single_file(wsi_path, msk_path, self.output_dir) - log_info("Finish") - except: - logging.exception("Crash") + + # if os.path.exists(output_file): + # log_info("Skip: %s" % wsi_base_name) + # continue + # try: + # log_info("Process: %s" % wsi_base_name) + # self.process_single_file(wsi_path, msk_path, self.output_dir) + # log_info("Finish") + # except: + # logging.exception("Crash") + self.process_single_file(wsi_path, msk_path, self.output_dir) + break rm_n_mkdir(self.cache_path) # clean up all cache return From f2b20e8dc51d6c76cad66364b9de30688eaa966b Mon Sep 17 00:00:00 2001 From: vqdang Date: Thu, 4 Mar 2021 02:53:42 +0000 Subject: [PATCH 06/19] UPD: fleshed out async pipeline --- infer/super_wsi.py | 252 +++++++++++++++++++++++---------------------- 1 file changed, 128 insertions(+), 124 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index e59fd983..b2d3d483 100644 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -40,10 +40,16 @@ #### class SerializeArray(data.Dataset): + """ + `mp_shared_space` must be from torch.multiprocessing, for example + + mp_manager = torch_mp.Manager() + mp_shared_space = mp_manager.Namespace() + mp_shared_space.image = torch.from_numpy(image) + """ def __init__(self, mp_shared_space, preproc=None): super().__init__() self.mp_shared_space = mp_shared_space - self.preproc = preproc return @@ -52,10 +58,8 @@ def __len__(self): def __getitem__(self, idx): patch_info = self.mp_shared_space.patch_info_list[idx] - patch_data = self.mp_shared_space.tile_img[ - patch_info[0,0] : patch_info[1,0], - patch_info[0,1] : patch_info[1,1], - ] + tl, br = patch_info[0] # retrieve input placement, [1] is output + patch_data = self.mp_shared_space.tile_img[tl[0] : br[0], tl[1] : br[1]] if self.preproc is not None: patch_dat = patch_data.copy() patch_data = self.preproc(patch_data) @@ -290,114 +294,87 @@ def get_info_stack(output_tl, output_br): #### def run_model( - mp_queue_output, + forward_output_queue, tile_info_list, - wsi_path, wsi_ext, wsi_proc_mag,): + wsi_path, wsi_ext, wsi_proc_mag, wsi_cache_path, + run_step, model, loader_kwargs): - # wsi_handler = get_file_handler(wsi_path, backend=wsi_ext) - # # ! cache here is for legacy and to deal with esoteric internal wsi format - # wsi_handler.prepare_reading( - # read_mag=self.proc_mag, cache_path="%s/src_wsi.npy" % self.cache_path - # ) + wsi_handler = get_file_handler(wsi_path, backend=wsi_ext) + # ! cache here is for legacy and to deal with esoteric internal wsi format + wsi_handler.prepare_reading(read_mag=wsi_proc_mag, cache_path=wsi_cache_path) # using shared memory namespace so all the loader workers use same - # underlying image data, also allowing persistent worker and fast data switching + # underlying image data, also allow persistent worker and fast data switching mp_manager = torch_mp.Manager() mp_shared_space = mp_manager.Namespace() ds = SerializeArray(mp_shared_space) - loader = data.DataLoader(ds, - num_workers=32, - batch_size=16, + loader = data.DataLoader(ds, **loader_kwargs, drop_last=False, persistent_workers=True, ) - for i in range (1000000): - patch_info_list = np.array([[[0, 0], [180, 180]]]*256, dtype=np.int32) - tile_img = np.full([5000, 5000, 3], i) - mp_shared_space.tile_img = torch.from_numpy(tile_img).share_memory_() - mp_shared_space.patch_info_list = torch.from_numpy(patch_info_list).share_memory_() + for tile_idx, tile_info in enumerate(tile_info_list): - # change the data without changing the loader - # so that the loader multiproc stays alives - for batch_idx, batch_data in enumerate(loader): - sample_data_list, sample_info_list = batch_data - print(sample_data_list[:,0,0,0]) - break + tile_info, patch_info_list = tile_info - while len(mp_queue_input > 0): - tile_info, patch_info_list = mp_queue_input.pop() - tile_tl, tile_br = tile_info - tile_img = wsi_handler.read_region( - tile_tl[::-1], (tile_br - tile_tl)[::-1] - ) - # change the data without changing the loader - # so that the loader multiproc stays alives + tile_input_info = tile_info[0] + tile_input_tl, tile_input_br = tile_input_info + # shift from wsi system to tile input system + # ! (tile output is within tile input system) + # ! this will shift both patch input and output placement to tile input system + # ! hence, output placement need to be shifted (corrected) later for post proc + patch_info_list -= np.reshape(tile_input_tl, [1, 1, 1, 2]) + + tile_img = wsi_handler.read_region(tile_input_tl[::-1], + (tile_input_br - tile_input_tl)[::-1]) + + # change the data in namespace to sync across persistent loader worker + # also no need to do locking as these are assumed to be read only from worker + mp_shared_space.tile_img = torch.from_numpy(tile_img).share_memory_() + mp_shared_space.patch_info_list = torch.from_numpy(patch_info_list).share_memory_() accumulated_patch_output = [] for batch_idx, batch_data in enumerate(loader): sample_data_list, sample_info_list = batch_data - sample_output_list = run_step(sample_data_list) + sample_output_list = run_step(sample_data_list, model) sample_info_list = sample_info_list.numpy() - curr_batch_size = sample_output_list.shape[0] - sample_output_list = np.split(sample_output_list, curr_batch_size, axis=0) - sample_info_list = np.split(sample_info_list, curr_batch_size, axis=0) - sample_output_list = list(zip(sample_info_list, sample_output_list)) - accumulated_patch_output.extend(sample_output_list) - mp_queue_output.append(accumulated_patch_output) - run inference on input patches - return - -def mp_dispatcher(data_list, nr_worker=0, show_pbar=True): - """ - data_list is alist of [[func, arg1, arg2, etc.]] - Resutls are alway sorted according to source position - """ - if nr_worker > 0: - proc_pool = ProcessPoolExecutor(nr_worker) - - result_list = [] - future_list = [] - - if show_pbar: - pbar = tqdm(total=len(data_list), ascii=True, position=0) - for run_idx, dat in enumerate(data_list): - func = dat[0] - args = dat[1:] - if nr_worker > 0: - future = proc_pool.submit(func, run_idx, *args) - future_list.append(future) - else: - # ! assume 1st return is alwasy run_id - result = func(run_idx, *args) - result_list.append(result) - if show_pbar: - pbar.update() - if nr_worker > 0: - for future in as_completed(future_list): - if future.exception() is not None: - logging.info(future.exception()) - else: - result = future.result() - result_list.append(result) - if show_pbar: - pbar.update() - if show_pbar: - pbar.close() - - result_list = sorted(result_list, key=lambda k : k[0]) - result_list = [v[1:] for v in result_list] - return result_list + accumulated_patch_output.append([sample_info_list, sample_output_list]) + forward_output_queue.put(accumulated_patch_output) + return True #### -def check_valid(run_idx, info, wsi_mask): - output_bbox = np.rint(info[1]).astype(np.int64) - output_roi = wsi_mask[ - output_bbox[0][0] : output_bbox[1][0], - output_bbox[0][1] : output_bbox[1][1], - ] - return run_idx, (torch.sum(output_roi) > 0).item() +def postproc_tile(tile_info, patch_info_list, func_opt): + # output pos of the tile within the source wsi + tile_input_tl, tile_output_br = tile_info[0] + tile_output_tl, tile_output_br = tile_info[1] # Y, X + offset = tile_output_tl - tile_input_tl + + # ! shape may be uneven hence just detach all into a big list + patch_pos_list = [] + patch_feat_list = [] + split_inst = lambda x : np.split(x, x.shape[0], axis=0) + for batch_pos, batch_feat in patch_info_list: + patch_pos_list.extend(split_inst(batch_pos)) + patch_feat_list.extend(split_inst(batch_feat)) + + nr_ch = patch_feat_list[-1].shape[-1] + tile_shape = (tile_output_br - tile_output_tl).tolist() + pred_map = np.zeros(tile_shape + [nr_ch], dtype=np.float32) + for idx in range(len(patch_pos_list)): + # zero idx to remove singleton, squeeze may kill h/w/c + patch_pos = patch_pos_list[idx][0].copy() + # ! assume patch pos alrd aligned to be within tile input system + patch_pos = patch_pos - offset # shift from wsi to tile output system + pos_tl, pos_br = patch_pos[1] # retrieve ouput placement + pred_map[ + pos_tl[0] : pos_br[0], + pos_tl[1] : pos_br[1] + ] = patch_feat_list[idx][0] + + postproc_func, postproc_kwargs = func_opt + pred_inst, inst_info_dict = postproc_func(pred_map, **postproc_kwargs) + return inst_info_dict #### class InferManager(base.InferManager): @@ -410,13 +387,20 @@ def __select_valid_patches(self, patch_info_list): has_output_info: whether output information is given """ + def check_valid(info, wsi_mask): + output_bbox = np.rint(info[1]).astype(np.int64) + output_roi = wsi_mask[ + output_bbox[0][0] : output_bbox[1][0], + output_bbox[0][1] : output_bbox[1][1], + ] + return (torch.sum(output_roi) > 0).item() down_sample_ratio = self.wsi_mask.shape[0] / self.wsi_proc_shape[0] torch_mask = torch.from_numpy(self.wsi_mask).share_memory_() - run_list = [[check_valid, info * down_sample_ratio, torch_mask] for info in patch_info_list] + valid_indices = [check_valid(info * down_sample_ratio, torch_mask) + for info in patch_info_list] # somehow multiproc is slower than single thread - valid_indices = mp_dispatcher(run_list, nr_worker=0, show_pbar=False) - valid_indices = np.concatenate(valid_indices, axis=0) + valid_indices = np.array(valid_indices) return patch_info_list[valid_indices] def __select_patches_in_tile(self, tile_info, patch_info_list): @@ -424,8 +408,8 @@ def __select_patches_in_tile(self, tile_info, patch_info_list): tile_tl, tile_br = tile_info[1] patch_tl_list = patch_info_list[:,1,0] patch_br_list = patch_info_list[:,1,1] - sel = (patch_tl_list[:,0] == tile_tl[0]) | (patch_tl_list[:,1] == tile_tl[1]) - sel |= (patch_br_list[:,0] == tile_br[0]) | (patch_br_list[:,1] == tile_br[1]) + sel = (patch_tl_list[:,0] >= tile_tl[0]) & (patch_tl_list[:,1] >= tile_tl[1]) + sel &= (patch_br_list[:,0] <= tile_br[0]) & (patch_br_list[:,1] <= tile_br[1]) return patch_info_list[sel] def _parse_args(self, run_args): @@ -518,8 +502,7 @@ def process_single_file(self, wsi_path, mask_path, output_dir): ) # * Async Inference - # * launch a seperate process to do forward and store the future, - # * the loader_workers may not need, to be spawned by the forward processs + # * launch a seperate process to do forward and store the result in a queue # * then while polling for forward result, launch separate process for # * doing the postproc # @@ -533,9 +516,9 @@ def process_single_file(self, wsi_path, mask_path, output_dir): patch_info_list = self.__select_valid_patches(patch_info_list) tile_info_list = self.__select_valid_patches(tile_info_list) - mp_manager = mp.Manager() + mp_manager = torch_mp.Manager() # contain at most 5 tile ouput before polling - mp_forward_o_queue = mp_manager.Queue(maxsize=5) + mp_forward_output_queue = mp_manager.Queue(maxsize=5) import collections forward_info_list = collections.deque() @@ -544,34 +527,55 @@ def process_single_file(self, wsi_path, mask_path, output_dir): for tile_info in tile_info_list: # retrieve valid patch within tile patch_in_tile_info_list = self.__select_patches_in_tile(tile_info, patch_info_list) - # shifting patch coord from wrt wsi to wrt to tile - tile_input_info = tile_info[0] - patch_input_info_list = patch_in_tile_info_list[:,0] - patch_input_info_list -= tile_input_info - forward_info_list.append([tile_input_info, patch_input_info_list]) + forward_info_list.append([tile_info, patch_in_tile_info_list]) + + loader_kwargs = dict( + num_workers=self.nr_inference_workers, + batch_size=self.batch_size, + ) + wsi_cache_path = "%s/src_wsi.npy" % self.cache_path forward_process = mp.Process(target=run_model, - args=(mp_forward_o_queue, forward_info_list, - wsi_path, wsi_ext, self.proc_mag)) + args=(mp_forward_output_queue, forward_info_list, + wsi_path, wsi_ext, self.proc_mag, wsi_cache_path, + self.run_step, self.net, loader_kwargs)) forward_process.start() + + post_proc_kwargs = { + "nr_types": self.method["model_args"]["nr_types"], + "return_centroids": True, + } + + proced_tile_counter = 0 + future_list = collections.deque() + proc_pool = ProcessPoolExecutor(self.nr_post_proc_workers) + # will this lead to infinite loop ? + while forward_process.exitcode is None: + if not mp_forward_output_queue.empty(): + # ! assume the forward result are in + # ! sequential as defined above + tile_info = tile_info_list[proced_tile_counter] + forward_output = mp_forward_output_queue.get() + future = proc_pool.submit(postproc_tile, + tile_info, forward_output, + (self.post_proc_func, post_proc_kwargs)) + # postproc_tile(tile_info, forward_output, + # (self.post_proc_func, post_proc_kwargs)) + future_list.append(future) # deal when forward finish or stick callback ? + proced_tile_counter += 1 + if forward_process.exitcode > 0: + raise ValueError(f'Forward process exited with code {forward_process.exitcode}') forward_process.join() - - # do double queue polling - # while len(future_list) > 0: - # while not future_list[0].done(): - # future_list.rotate() - # proc_future = future_list.popleft() - # if proc_future.exception() is not None: - # print(proc_future.exception()) - # proc_run_id, proc_result = proc_future.result() - # if proc_run_id == 'forward': - # proc_future = postproc_mp_pool.submit() - # future_list.append(proc_future) - # elif proc_run_id == 'postproc': - # proc_future.result() - # else: - # assert False + + while len(future_list) > 0: + if not future_list[0].done(): + future_list.rotate() + continue + proc_future = future_list.popleft() + if proc_future.exception() is not None: + print(proc_future.exception()) + proc_future.result() return From 6179a8df785d483b5429ab99ba42bf30f33f137b Mon Sep 17 00:00:00 2001 From: vqdang Date: Fri, 5 Mar 2021 17:45:35 +0000 Subject: [PATCH 07/19] UPD: track --- infer/super_wsi.py | 350 +++++++++++++++++++++++++++++++++------------ 1 file changed, 262 insertions(+), 88 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index b2d3d483..1869a6d6 100644 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -1,9 +1,11 @@ import multiprocessing as mp -from concurrent.futures import FIRST_EXCEPTION, ProcessPoolExecutor, as_completed, wait +from concurrent.futures import (FIRST_EXCEPTION, ProcessPoolExecutor, + as_completed, wait) mp.set_start_method("spawn", True) # ! must be at top for VScode debugging import argparse +import collections import glob import json import logging @@ -25,16 +27,11 @@ import torch.multiprocessing as torch_mp import torch.utils.data as data import tqdm -from docopt import docopt - -from misc.utils import ( - cropping_center, - get_bounding_box, - log_debug, - log_info, - rm_n_mkdir, -) + +from misc.utils import (cropping_center, get_bounding_box, log_debug, log_info, + rm_n_mkdir) from misc.wsi_handler import get_file_handler + from . import base @@ -118,19 +115,7 @@ def flat_mesh_grid_coord(y, x): # print(info_list.shape) return info_list -# info_list = _get_patch_info( -# np.array([70, 100]), -# np.array([40, 40]), -# np.array([20, 20])) -# print(info_list[:,1,0]) - -# info_list = _get_patch_info( -# np.array([100, 100]), -# np.array([80, 80]), -# np.array([60, 60])) -# print(info_list[:,0,1]) -# exit() - +#### def _get_tile_info(img_shape, input_size, output_size, margin_size, unit_size): """Get top left coordinate information of patches from original image. @@ -188,6 +173,8 @@ def get_info_stack(output_tl, output_br): np.stack([output_tl, output_br], axis=1), ], axis=1) return info_list + + # * Full Tile Grid info_list = get_info_stack(output_tl, output_br).astype(np.int64) # flag surrounding ambiguous (left margin, right margin) @@ -207,13 +194,12 @@ def get_info_stack(output_tl, output_br): removal_flag[(info_list[:,1,0,0] == np.min(output_tl[:,0])),2] = 0 # exclude those contain bot most boundary removal_flag[(info_list[:,1,1,0] == np.max(output_br[:,0])),3] = 0 - # print(removal_flag) - # print(info_list[...,::-1][:,1]) - # exit() + mode_list = np.full(info_list.shape[0], 0) + all_info = [[info_list, removal_flag, mode_list]] br_most = np.max(output_br, axis=0) tl_most = np.min(output_tl, axis=0) - # * ------------------------------- + # * Tile Boundary Redo with Margin # get the fix grid tile info y_fix_output_tl = output_tl - np.array([margin_size[0], 0])[None,:] y_fix_output_br = np.stack([output_tl[:,0], output_br[:,1]], axis=-1) @@ -225,21 +211,20 @@ def get_info_stack(output_tl, output_br): # sel position not on the image boundary sel = (output_tl[:,0] == np.min(output_tl[:,0])) y_info_list = get_info_stack(y_fix_output_tl[~sel], y_fix_output_br[~sel]) - # print(y_info_list[...,::-1][:,1],'\n') # flag horizontal ambiguous region for y (left margin, right margin) # |----|------------|----| # |\\\\| |\\\\| # |----|------------|----| - # ambiguous ambiguous (margin size) + # <----> ambiguous (margin size) removal_flag = np.zeros((y_info_list.shape[0], 4,)) # left, right, top, bot removal_flag[:,[0,1]] = 1 # exclude the left most boundary removal_flag[(y_info_list[:,1,0,1] == np.min(output_tl[:,1])),0] = 0 # exclude the right most boundary removal_flag[(y_info_list[:,1,1,1] == np.max(output_br[:,1])),1] = 0 - # print(removal_flag) - # print(y_info_list[...,::-1][:,1]) + mode_list = np.full(y_info_list.shape[0], 1) + # all_info.append([y_info_list, removal_flag, mode_list]) x_fix_output_br = output_br + np.array([0, margin_size[1]])[None,:] x_fix_output_tl = np.stack([output_tl[:,0], output_br[:,1]], axis=-1) @@ -250,11 +235,10 @@ def get_info_stack(output_tl, output_br): # sel position not on the image boundary sel = (output_br[:,1] == np.max(output_br[:,1])) x_info_list = get_info_stack(x_fix_output_tl[~sel], x_fix_output_br[~sel]) - # print(x_info_list[...,::-1][:,1],'\n') # flag vertical ambiguous region for x (top margin, bottom margin) - # |----| - # |\\\\| ambiguous - # |----| + # |----| ^ + # |\\\\| | ambiguous + # |----| V # | | # | | # |----| @@ -266,8 +250,8 @@ def get_info_stack(output_tl, output_br): removal_flag[(x_info_list[:,1,0,0] == np.min(output_tl[:,0])),2] = 0 # exclude the right most boundary removal_flag[(x_info_list[:,1,1,0] == np.max(output_br[:,0])),3] = 0 - # print(removal_flag) - # print(x_info_list[...,::-1][:,1]) + mode_list = np.full(x_info_list.shape[0], 2) + # all_info.append([x_info_list, removal_flag, mode_list]) # * define the tile cross section sel = np.any(output_br == br_most, axis=-1) @@ -279,18 +263,18 @@ def get_info_stack(output_tl, output_br): xsect_br[xsect_br[:,1] > br_most[1], 1] = br_most[1] xsect_tl[xsect_tl[:,0] < tl_most[0], 0] = tl_most[0] xsect_tl[xsect_tl[:,1] < tl_most[1], 1] = tl_most[1] - xsect_info = get_info_stack(xsect_tl, xsect_br) + xsect_info_list = get_info_stack(xsect_tl, xsect_br) + mode_list = np.full(xsect_info_list.shape[0], 3) + removal_flag = np.full((xsect_info_list.shape[0], 4,), 0) # left, right, top, bot + # all_info.append([xsect_info_list, removal_flag, mode_list]) - return info_list + # ! combine all + info_list, removal_flag, mode_list = list(zip(*all_info)) + info_list = np.concatenate(info_list, axis=0).astype(np.int32) + mode_list = np.concatenate(mode_list, axis=0) + removal_flag = np.concatenate(removal_flag, axis=0) -# info_list = _get_tile_info( -# np.array([170, 100]), # [130, 100], [140, 100] -# np.array([80, 80]), -# np.array([60, 60]), -# np.array([20, 20]), -# np.array([20, 20]),) -# print(info_list[:,0]) -# exit() + return info_list, removal_flag, mode_list #### def run_model( @@ -340,14 +324,106 @@ def run_model( sample_output_list = run_step(sample_data_list, model) sample_info_list = sample_info_list.numpy() accumulated_patch_output.append([sample_info_list, sample_output_list]) - forward_output_queue.put(accumulated_patch_output) - return True + forward_output_queue.put([tile_idx, accumulated_patch_output]) + print('%d/%d' % (tile_idx, len(tile_info_list))) + return +#### +#### +# ! seem to be 1 pix off at cross section or sthg +def get_inst_in_margin(arr, margin_size, tile_pp_info): + """ + include the margin line itself + """ + tile_pp_info = np.array(tile_pp_info) + + inst_in_margin = [] + # extract those lie within margin region + if tile_pp_info[0] == 1: # left edge + inst_in_margin.append(arr[:,:(margin_size+1)]) + if tile_pp_info[1] == 1: # right edge + inst_in_margin.append(arr[:,-(margin_size+1):]) + if tile_pp_info[2] == 1: # top edge + inst_in_margin.append(arr[:(margin_size+1),:]) + if tile_pp_info[3] == 1: # bottom edge + inst_in_margin.append(arr[-(margin_size+1):,:]) + inst_in_margin = [v.flatten() for v in inst_in_margin] + if len(inst_in_margin) > 0: + inst_in_margin = np.concatenate(inst_in_margin, axis=0) + inst_in_margin = np.unique(inst_in_margin) + else: + inst_in_margin = np.array([]) # empty array + return inst_in_margin +#### +# ! BUG: fix this, this create exclusive problem due wildcard mass selection +def get_inst_on_margin(arr, margin_size, tile_pp_info): + tile_pp_info = np.array(tile_pp_info) + # extract those lie on the margin line + # l r t b + # ! define the line crossing and derive the cross point replacement will be + # ! much easier to manage ! + inst_on_margin = [] # just need to do 1 pix check + if (tile_pp_info == [1, 1, 1, 1]).all(): + inst_on_margin.append(arr[margin_size:-margin_size, margin_size ]) + inst_on_margin.append(arr[margin_size:-margin_size,-(margin_size+1)]) + inst_on_margin.append(arr[ margin_size ,margin_size:-margin_size]) + inst_on_margin.append(arr[-(margin_size+1),margin_size:-margin_size]) + elif (tile_pp_info == [1, 0, 1, 0]).all(): + inst_on_margin.append(arr[ margin_size ,margin_size:-margin_size]) + inst_on_margin.append(arr[-(margin_size+1),margin_size:-margin_size]) + elif (tile_pp_info == [1, 0, 0, 1]).all(): + inst_on_margin.append(arr[ margin_size ,margin_size:-margin_size]) + inst_on_margin.append(arr[-(margin_size+1),margin_size:-margin_size]) + elif (tile_pp_info == [0, 1, 1, 0]).all(): + inst_on_margin.append(arr[ margin_size ,margin_size:-margin_size]) + inst_on_margin.append(arr[-(margin_size+1),margin_size:-margin_size]) + elif (tile_pp_info == [0, 1, 0, 1]).all(): + inst_on_margin.append(arr[ margin_size ,margin_size:-margin_size]) + inst_on_margin.append(arr[-(margin_size+1),margin_size:-margin_size]) + elif (tile_pp_info == [1, 0, 0, 0]).all() or \ + (tile_pp_info == [0, 1, 0, 0]).all() or \ + (tile_pp_info == [1, 1, 0, 0]).all() or \ + (tile_pp_info == [0, 0, 1, 1]).all() or \ + (tile_pp_info == [0, 0, 1, 0]).all() or \ + (tile_pp_info == [0, 0, 0, 1]).all(): + if tile_pp_info[0] == 1: + inst_on_margin.append(arr[:,margin_size]) + if tile_pp_info[1] == 1: + inst_on_margin.append(arr[:,-(margin_size+1)]) + if tile_pp_info[2] == 1: + inst_on_margin.append(arr[margin_size,:]) + if tile_pp_info[3] == 1: + inst_on_margin.append(arr[-(margin_size+1),:]) + else: + assert False + inst_on_margin = [v.flatten() for v in inst_on_margin] + + if len(inst_on_margin) > 0: + inst_on_margin = np.concatenate(inst_on_margin, axis=0) + inst_on_margin = np.unique(inst_on_margin) + else: + inst_on_margin = np.array([]) # empty array + + return inst_on_margin +# a = np.reshape(np.arange(0, 64), [8, 8]) +# print(a) +# print(get_inst_on_margin(a, 2, [1, 1 ,1, 1])) +# print(get_inst_on_margin(a, 1, [1, 1 ,1, 1])) +# print(get_inst_on_margin(a, 1, [1, 1 ,0, 0])) +# print(get_inst_on_margin(a, 1, [0, 0 ,1, 1])) +# print('') +# print(get_inst_in_margin(a, 1, [1, 1 ,1, 1])) +# print(get_inst_in_margin(a, 1, [0, 1 ,1, 1])) +# print(get_inst_in_margin(a, 1, [1, 0 ,1, 1])) +# print(get_inst_in_margin(a, 1, [0, 0 ,1, 1])) +# print(get_inst_in_margin(a, 1, [0, 1 ,1, 0])) +# exit() #### -def postproc_tile(tile_info, patch_info_list, func_opt): +def postproc_tile(tile_io_info, tile_pp_info, tile_mode, + margin_size, patch_info_list, func_opt): # output pos of the tile within the source wsi - tile_input_tl, tile_output_br = tile_info[0] - tile_output_tl, tile_output_br = tile_info[1] # Y, X + tile_input_tl, tile_output_br = tile_io_info[0] + tile_output_tl, tile_output_br = tile_io_info[1] # Y, X offset = tile_output_tl - tile_input_tl # ! shape may be uneven hence just detach all into a big list @@ -358,6 +434,7 @@ def postproc_tile(tile_info, patch_info_list, func_opt): patch_pos_list.extend(split_inst(batch_pos)) patch_feat_list.extend(split_inst(batch_feat)) + # * assemble patch to tile nr_ch = patch_feat_list[-1].shape[-1] tile_shape = (tile_output_br - tile_output_tl).tolist() pred_map = np.zeros(tile_shape + [nr_ch], dtype=np.float32) @@ -371,10 +448,81 @@ def postproc_tile(tile_info, patch_info_list, func_opt): pos_tl[0] : pos_br[0], pos_tl[1] : pos_br[1] ] = patch_feat_list[idx][0] + del patch_pos_list, patch_feat_list + # * retrieve actual output postproc_func, postproc_kwargs = func_opt pred_inst, inst_info_dict = postproc_func(pred_map, **postproc_kwargs) - return inst_info_dict + del pred_map + + # * perform removal for ambiguous region + + # Consider each symbol as 1 pixel + # This is margin inner area //// + # ----------------------------- ^ ^ + # |///////////////////////////| | | margin area + # |///////////////////////////| | | (yes including the inner and outer edge) + # |///|-------------------|///| | V + # |///| ^margin_size |///| | + # |///| |///| | + # |///| <-- margin line |///| | Image area + # |///| | |///| | + # |///| v |///| | + # |///|-------------------|///| | + # |///////////////////////////| | + # |///////////////////////////| | + # ----------------------------| V + + + if tile_mode == 0: + # for `full grid tile` + # -- extend from the boundary by the margin size, remove + # nuclei lie within the margin area but exclude those + # lie on the margin line + # also contain those lying on the edges + inst_in_margin = get_inst_in_margin(pred_inst, margin_size, tile_pp_info) + # those lying on the margin line + inst_on_margin = get_inst_on_margin(pred_inst, margin_size, tile_pp_info) + inst_within_margin = np.setdiff1d(inst_in_margin, inst_on_margin, assume_unique=True) + remove_inst_set = inst_within_margin.tolist() + elif tile_mode == 1 or tile_mode == 2: + # for `horizontal/vertical strip tiles` for fixing artifacts + # -- extend from the marked edges (top/bot or left/right) by the margin size, + # remove all nuclei lie within the margin area (including on the margin line) + # -- nuclei on all edges are removed (as these are alrd within `full grid tile`) + inst_in_margin = get_inst_in_margin(margin_size, tile_pp_info) # also contain those lying on the edges + if np.sum(tile_pp_info) == 1: + holder_flag = tile_pp_info.copy() + if tile_mode == 1: + holder_flag[[2, 3]] = 1 + else: + holder_flag[[0, 1]] = 1 + else: + holder_flag = [1, 1, 1, 1] + print(tile_mode, tile_pp_info, holder_flag) + inst_on_edge = get_inst_on_margin(0, holder_flag) + remove_inst_set = np.union1d(inst_in_margin, inst_on_edge) + remove_inst_set = remove_inst_set.tolist() + else: + # inst within the tile after excluding margin area out + # only for a tile at cross-section, which is designed such that + # their shape >= 3* margin size + all_inst = np.unique() + remove_inst_set = [] + + remove_inst_set = set(remove_inst_set) + + # * move pos back to wsi position + renew_id = 0 + new_inst_info_dict = {} + for k, v in inst_info_dict.items(): + if k not in remove_inst_set: + v['bbox'] += tile_output_tl[::-1] + v['centroid'] += tile_output_tl[::-1] + v['contour'] += tile_output_tl[::-1] + new_inst_info_dict[renew_id] = v + renew_id += 1 + return new_inst_info_dict #### class InferManager(base.InferManager): @@ -493,38 +641,37 @@ def process_single_file(self, wsi_path, mask_path, output_dir): ) patch_diff_shape = self.patch_input_shape - self.patch_output_shape # derive tile output placement as consecutive tiling with step size of 0 - # and tile output will have shape of multiple of patch_output_shape (round down) + # and tile output will have shape being of multiple of patch_output_shape (round down) tile_output_shape = np.floor(self.tile_shape / self.patch_output_shape) * self.patch_output_shape tile_input_shape = tile_output_shape + patch_diff_shape - tile_info_list = _get_tile_info( + + tile_io_info_list, \ + tile_pp_info_list, \ + tile_mode_list = _get_tile_info( self.wsi_proc_shape, tile_input_shape, tile_output_shape, self.ambiguous_size, self.patch_output_shape ) # * Async Inference - # * launch a seperate process to do forward and store the result in a queue - # * then while polling for forward result, launch separate process for - # * doing the postproc + # * launch a seperate process to do forward and store the result in a queue. + # * main thread will poll for forward result and launch separate process for + # * doing the postproc for every newly predicted tiles # # / forward \ (loop) # main ------------- main----------------main # \ postproc (loop)/ - # future_list = collections.deque() - # mp_pool = ProcessPoolExecutor(self.nr_post_proc_workers + 1) - patch_info_list = self.__select_valid_patches(patch_info_list) - tile_info_list = self.__select_valid_patches(tile_info_list) + tile_io_info_list = self.__select_valid_patches(tile_io_info_list) mp_manager = torch_mp.Manager() - # contain at most 5 tile ouput before polling + # contain at most 5 tile ouput before polling forward func mp_forward_output_queue = mp_manager.Queue(maxsize=5) - import collections forward_info_list = collections.deque() - nr_tile = tile_info_list.shape[0] - for tile_info in tile_info_list: + nr_tile = tile_io_info_list.shape[0] + for tile_info in tile_io_info_list: # retrieve valid patch within tile patch_in_tile_info_list = self.__select_patches_in_tile(tile_info, patch_info_list) forward_info_list.append([tile_info, patch_in_tile_info_list]) @@ -539,7 +686,6 @@ def process_single_file(self, wsi_path, mask_path, output_dir): args=(mp_forward_output_queue, forward_info_list, wsi_path, wsi_ext, self.proc_mag, wsi_cache_path, self.run_step, self.net, loader_kwargs)) - forward_process.start() post_proc_kwargs = { @@ -547,35 +693,63 @@ def process_single_file(self, wsi_path, mask_path, output_dir): "return_centroids": True, } - proced_tile_counter = 0 + self.wsi_inst_info = {} future_list = collections.deque() proc_pool = ProcessPoolExecutor(self.nr_post_proc_workers) # will this lead to infinite loop ? - while forward_process.exitcode is None: + while True: + if forward_process.exitcode is not None \ + and mp_forward_output_queue.empty(): + break + if not mp_forward_output_queue.empty(): - # ! assume the forward result are in - # ! sequential as defined above - tile_info = tile_info_list[proced_tile_counter] - forward_output = mp_forward_output_queue.get() - future = proc_pool.submit(postproc_tile, - tile_info, forward_output, - (self.post_proc_func, post_proc_kwargs)) - # postproc_tile(tile_info, forward_output, + # ! assume the forward result are spit out in + # ! same order as defined in `forward_info_list` + tile_idx, forward_output = mp_forward_output_queue.get() + tile_io_info = tile_io_info_list[tile_idx] + tile_pp_info = tile_pp_info_list[tile_idx] + tile_mode = tile_mode_list[tile_idx] + + # future = proc_pool.submit(postproc_tile, + # tile_info, forward_output, # (self.post_proc_func, post_proc_kwargs)) - future_list.append(future) # deal when forward finish or stick callback ? - proced_tile_counter += 1 + + tile_inst_dict = postproc_tile( + tile_io_info, tile_pp_info, tile_mode, + self.ambiguous_size[0], + forward_output, + (self.post_proc_func, post_proc_kwargs)) + + # ! may not work if output id range > maximum #inst in return dict + offset_id = len(self.wsi_inst_info) + for tile_inst_id, tile_inst_info in tile_inst_dict.items(): + inst_wsi_id = offset_id + tile_inst_id + 1 + self.wsi_inst_info[inst_wsi_id] = tile_inst_info + + print('Post proc %d' % tile_idx) + # future_list.append(future) # deal when forward finish or stick callback ? if forward_process.exitcode > 0: raise ValueError(f'Forward process exited with code {forward_process.exitcode}') forward_process.join() - while len(future_list) > 0: - if not future_list[0].done(): - future_list.rotate() - continue - proc_future = future_list.popleft() - if proc_future.exception() is not None: - print(proc_future.exception()) - proc_future.result() + # while len(future_list) > 0: + # if not future_list[0].done(): + # future_list.rotate() + # continue + # proc_future = future_list.popleft() + # if proc_future.exception() is not None: + # print(proc_future.exception()) + # proc_future.result() + + if self.save_mask or self.save_thumb: + json_path = "%s/json/%s.json" % (output_dir, wsi_name) + else: + json_path = "%s/%s.json" % (output_dir, wsi_name) + # self.__save_json(json_path, self.wsi_inst_info) + wsi_thumb_rgb = self.wsi_handler.get_full_img(read_mag=self.proc_mag) + from misc.viz_utils import visualize_instances_dict + wsi_overlay = visualize_instances_dict(wsi_thumb_rgb, self.wsi_inst_info, draw_dot=True) + cv2.imwrite('dump.png', cv2.cvtColor(wsi_overlay, cv2.COLOR_RGB2BGR)) return @@ -602,7 +776,7 @@ def process_wsi_list(self, run_args): wsi_path_list = glob.glob(self.input_dir + "/*") wsi_path_list.sort() # ensure ordering - for wsi_path in wsi_path_list[::-1]: + for wsi_path in wsi_path_list[1:2]: wsi_base_name = pathlib.Path(wsi_path).stem msk_path = "%s/%s.png" % (self.input_mask_dir, wsi_base_name) if self.save_thumb or self.save_mask: From 309e8d5b8f084d7042c588885df13dab30a86396 Mon Sep 17 00:00:00 2001 From: vqdang Date: Mon, 8 Mar 2021 13:48:26 +0000 Subject: [PATCH 08/19] UPD: fix extract inst on margin line, add viz test --- infer/super_wsi.py | 317 +++++++++++++++++++++++++++++++-------------- 1 file changed, 219 insertions(+), 98 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index 1869a6d6..4c82df2c 100644 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -224,7 +224,7 @@ def get_info_stack(output_tl, output_br): # exclude the right most boundary removal_flag[(y_info_list[:,1,1,1] == np.max(output_br[:,1])),1] = 0 mode_list = np.full(y_info_list.shape[0], 1) - # all_info.append([y_info_list, removal_flag, mode_list]) + all_info.append([y_info_list, removal_flag, mode_list]) x_fix_output_br = output_br + np.array([0, margin_size[1]])[None,:] x_fix_output_tl = np.stack([output_tl[:,0], output_br[:,1]], axis=-1) @@ -251,7 +251,7 @@ def get_info_stack(output_tl, output_br): # exclude the right most boundary removal_flag[(x_info_list[:,1,1,0] == np.max(output_br[:,0])),3] = 0 mode_list = np.full(x_info_list.shape[0], 2) - # all_info.append([x_info_list, removal_flag, mode_list]) + all_info.append([x_info_list, removal_flag, mode_list]) # * define the tile cross section sel = np.any(output_br == br_most, axis=-1) @@ -268,6 +268,7 @@ def get_info_stack(output_tl, output_br): removal_flag = np.full((xsect_info_list.shape[0], 4,), 0) # left, right, top, bot # all_info.append([xsect_info_list, removal_flag, mode_list]) + all_info = all_info[:1] # ! combine all info_list, removal_flag, mode_list = list(zip(*all_info)) info_list = np.concatenate(info_list, axis=0).astype(np.int32) @@ -275,7 +276,218 @@ def get_info_stack(output_tl, output_br): removal_flag = np.concatenate(removal_flag, axis=0) return info_list, removal_flag, mode_list +#### +# ! seem to be 1 pix off at cross section or sthg +def get_inst_in_margin(arr, margin_size, tile_pp_info): + """ + include the margin line itself + """ + assert margin_size > 0 + tile_pp_info = np.array(tile_pp_info) + + inst_in_margin = [] + # extract those lie within margin region + if tile_pp_info[0] == 1: # left edge + inst_in_margin.append(arr[:,:margin_size]) + if tile_pp_info[1] == 1: # right edge + inst_in_margin.append(arr[:,-margin_size:]) + if tile_pp_info[2] == 1: # top edge + inst_in_margin.append(arr[:margin_size,:]) + if tile_pp_info[3] == 1: # bottom edge + inst_in_margin.append(arr[-margin_size:,:]) + inst_in_margin = [v.flatten() for v in inst_in_margin] + if len(inst_in_margin) > 0: + inst_in_margin = np.concatenate(inst_in_margin, axis=0) + inst_in_margin = np.unique(inst_in_margin) + else: + inst_in_margin = np.array([]) # empty array + return inst_in_margin +#### +def get_inst_on_margin(arr, margin_size, tile_pp_info): + """ + """ + assert margin_size > 0 + # extract those lie on the margin line + tile_pp_info = np.array(tile_pp_info) + def line_intersection(line1, line2): + ydiff = (line1[0][0] - line1[1][0], line2[0][0] - line2[1][0]) + xdiff = (line1[0][1] - line1[1][1], line2[0][1] - line2[1][1]) + + def det(a, b): + return a[0] * b[1] - a[1] * b[0] + + div = det(xdiff, ydiff) + if div == 0: + return False # not intersect + + d = (det(*line1), det(*line2)) + x = det(d, xdiff) / div + y = det(d, ydiff) / div + # ! positive region only (due to indexing line) + return int(abs(y)), int(abs(x)) + + last_h, last_w = arr.shape[0]-1, arr.shape[1]-1 + line_list = [ + [[0, 0] , [last_h, 0] ], # left line + [[0, last_w], [last_h, last_w]], # right line + [[0, 0] , [0, last_w] ], # top line + [[last_h, 0], [last_h, last_w]], # bottom line + ] + + if tile_pp_info[0] == 1: + line_list[0] = [[0 , margin_size], + [last_h, margin_size]] + if tile_pp_info[1] == 1: + line_list[1] = [[0 , last_w-margin_size], + [last_h, last_w-margin_size]] + if tile_pp_info[2] == 1: + line_list[2] = [[margin_size, 0], + [margin_size, last_w]] + if tile_pp_info[3] == 1: + line_list[3] = [[last_h-margin_size, 0], + [last_h-margin_size, last_w]] + # x1 x2 + # x3 x4 + # all pts need to be valid idx ! + pts_list = [ + line_intersection(line_list[2], line_list[0]), # x1 + line_intersection(line_list[2], line_list[1]), # x2 + line_intersection(line_list[3], line_list[0]), # x3 + line_intersection(line_list[3], line_list[1]), # x4 + ] + + def sel_between_pts(p1, p2): + arr[p1[0]:p2[0]+1, + p1[1]:p2[1]+1] = 1 + + line_pix_list = [] + if tile_pp_info[0] == 1: + sel_between_pts(pts_list[0], pts_list[2]) + if tile_pp_info[1] == 1: + sel_between_pts(pts_list[1], pts_list[3]) + if tile_pp_info[2] == 1: + sel_between_pts(pts_list[0], pts_list[1]) + if tile_pp_info[3] == 1: + sel_between_pts(pts_list[2], pts_list[3]) + + # pt_index = lambda p1, p2: arr[p1[0]:p2[0]+1, + # p1[1]:p2[1]+1] + # line_pix_list = [] + # if tile_pp_info[0] == 1: + # line_pix_list.append(pt_index(pts_list[0], pts_list[2])), + # if tile_pp_info[1] == 1: + # line_pix_list.append(pt_index(pts_list[1], pts_list[3])), + # if tile_pp_info[2] == 1: + # line_pix_list.append(pt_index(pts_list[0], pts_list[1])), + # if tile_pp_info[3] == 1: + # line_pix_list.append(pt_index(pts_list[2], pts_list[3])), + + # inst_on_margin = [v.flatten() for v in line_pix_list] + + # if len(inst_on_margin) > 0: + # inst_on_margin = np.concatenate(inst_on_margin, axis=0) + # inst_on_margin = np.unique(inst_on_margin) + # else: + # inst_on_margin = np.array([]) # empty array + + # return inst_on_margin +#### +def get_inst_on_edge(arr, tile_pp_info): + inst_on_edge = [] + if tile_pp_info[0] == 1: + inst_on_edge.append(arr[:,0]) + if tile_pp_info[1] == 1: + inst_on_edge.append(arr[:,-1]) + if tile_pp_info[2] == 1: + inst_on_edge.append(arr[0,:]) + if tile_pp_info[3] == 1: + inst_on_edge.append(arr[-1,:]) + + inst_on_edge = [v.flatten() for v in inst_on_edge] + + if len(inst_on_edge) > 0: + inst_on_edge = np.concatenate(inst_on_edge, axis=0) + inst_on_edge = np.unique(inst_on_edge) + else: + inst_on_edge = np.array([]) # empty array + return inst_on_edge +#### +# a = np.reshape(np.arange(0, 64), [8, 8]) +# print(a) +# print(get_inst_on_margin(a, 2, [1, 1 ,1, 1])) +# print(get_inst_on_margin(a, 1, [1, 1 ,1, 1])) +# print(get_inst_on_margin(a, 1, [1, 1 ,0, 0])) +# print(get_inst_on_margin(a, 1, [0, 0 ,1, 1])) +# print(get_inst_on_margin(a, 1, [0, 1 ,1, 0])) +# print(get_inst_on_margin(a, 1, [0, 1 ,1, 1])) +# print(get_inst_on_margin(a, 1, [1, 1 ,0, 1])) +# print('') +# print(get_inst_in_margin(a, 1, [1, 1 ,1, 1])) +# print(get_inst_in_margin(a, 1, [0, 1 ,1, 1])) +# print(get_inst_in_margin(a, 1, [1, 0 ,1, 1])) +# print(get_inst_in_margin(a, 1, [0, 0 ,1, 1])) +# print(get_inst_in_margin(a, 1, [0, 1 ,1, 0])) +# exit() +#### +tile_shape = np.array([2048, 2048]) +image_shape = np.array([3600, 3600]) +patch_input_shape = np.array([256,256]) +patch_output_shape = np.array([164,164]) +ambiguous_size = patch_output_shape * 2 +patch_info_list = _get_patch_info(image_shape, patch_input_shape, patch_output_shape) + +patch_diff_shape = patch_input_shape - patch_output_shape +# derive tile output placement as consecutive tiling with step size of 0 +# and tile output will have shape being of multiple of patch_output_shape (round down) +tile_output_shape = np.floor(tile_shape / patch_output_shape) * patch_output_shape +tile_input_shape = tile_output_shape + patch_diff_shape + +tile_io_info_list, \ + tile_pp_info_list, \ + tile_mode_list = _get_tile_info( + image_shape, tile_input_shape, tile_output_shape, + ambiguous_size, patch_output_shape +) + +import matplotlib.pyplot as plt +cmap = plt.get_cmap('jet') + +# canvas = np.zeros(image_shape) +# nr_patch = patch_info_list.shape[0] +# for idx in range(0, nr_patch): +# tl, br = patch_info_list[idx][1] +# # canvas[tl[0]:br[0],tl[1]:br[1]] += (idx+1) +# patch = canvas[tl[0]:br[0],tl[1]:br[1]] +# patch[0,:] = 1 +# patch[-1,:] = 1 +# patch[:,0] = 1 +# patch[:,-1] = 1 +# # canvas[tl[0]:br[0],tl[1]:br[1]] += (idx+1) + +# canvas_color = (cmap(canvas / np.max(canvas)) * 255).astype(np.uint8)[...,:3] +# canvas_color[canvas==0] = 0 # background +# cv2.imwrite('patch.png', cv2.cvtColor(canvas_color, cv2.COLOR_RGB2BGR)) + +canvas = np.zeros(image_shape) +nr_tile = tile_io_info_list.shape[0] +for idx in range(0, nr_tile): + tl, br = tile_io_info_list[idx][1] + canvas[tl[0]:br[0],tl[1]:br[1]] += (idx+1) +canvas_color = (cmap(canvas / np.max(canvas)) * 255).astype(np.uint8)[...,:3] +canvas_color[canvas==0] = 0 # background +cv2.imwrite('tile.png', cv2.cvtColor(canvas_color, cv2.COLOR_RGB2BGR)) + +canvas = np.zeros(image_shape) +for idx in range(0, nr_tile): + tl, br = tile_io_info_list[idx][1] + print(tl, br, tile_pp_info_list[idx]) + sub_canvas = canvas[tl[0]:br[0],tl[1]:br[1]] + canvas[tl[0]:br[0],tl[1]:br[1]] = get_inst_on_margin(sub_canvas, ambiguous_size[0], tile_pp_info_list[idx]) +canvas_color = (cmap(canvas / np.max(canvas)) * 255).astype(np.uint8)[...,:3] +canvas_color[canvas==0] = 0 # background +cv2.imwrite('tile_fix.png', cv2.cvtColor(canvas_color, cv2.COLOR_RGB2BGR)) +exit() #### def run_model( forward_output_queue, @@ -329,96 +541,6 @@ def run_model( print('%d/%d' % (tile_idx, len(tile_info_list))) return #### -#### -# ! seem to be 1 pix off at cross section or sthg -def get_inst_in_margin(arr, margin_size, tile_pp_info): - """ - include the margin line itself - """ - tile_pp_info = np.array(tile_pp_info) - - inst_in_margin = [] - # extract those lie within margin region - if tile_pp_info[0] == 1: # left edge - inst_in_margin.append(arr[:,:(margin_size+1)]) - if tile_pp_info[1] == 1: # right edge - inst_in_margin.append(arr[:,-(margin_size+1):]) - if tile_pp_info[2] == 1: # top edge - inst_in_margin.append(arr[:(margin_size+1),:]) - if tile_pp_info[3] == 1: # bottom edge - inst_in_margin.append(arr[-(margin_size+1):,:]) - inst_in_margin = [v.flatten() for v in inst_in_margin] - if len(inst_in_margin) > 0: - inst_in_margin = np.concatenate(inst_in_margin, axis=0) - inst_in_margin = np.unique(inst_in_margin) - else: - inst_in_margin = np.array([]) # empty array - return inst_in_margin -#### -# ! BUG: fix this, this create exclusive problem due wildcard mass selection -def get_inst_on_margin(arr, margin_size, tile_pp_info): - tile_pp_info = np.array(tile_pp_info) - # extract those lie on the margin line - # l r t b - # ! define the line crossing and derive the cross point replacement will be - # ! much easier to manage ! - inst_on_margin = [] # just need to do 1 pix check - if (tile_pp_info == [1, 1, 1, 1]).all(): - inst_on_margin.append(arr[margin_size:-margin_size, margin_size ]) - inst_on_margin.append(arr[margin_size:-margin_size,-(margin_size+1)]) - inst_on_margin.append(arr[ margin_size ,margin_size:-margin_size]) - inst_on_margin.append(arr[-(margin_size+1),margin_size:-margin_size]) - elif (tile_pp_info == [1, 0, 1, 0]).all(): - inst_on_margin.append(arr[ margin_size ,margin_size:-margin_size]) - inst_on_margin.append(arr[-(margin_size+1),margin_size:-margin_size]) - elif (tile_pp_info == [1, 0, 0, 1]).all(): - inst_on_margin.append(arr[ margin_size ,margin_size:-margin_size]) - inst_on_margin.append(arr[-(margin_size+1),margin_size:-margin_size]) - elif (tile_pp_info == [0, 1, 1, 0]).all(): - inst_on_margin.append(arr[ margin_size ,margin_size:-margin_size]) - inst_on_margin.append(arr[-(margin_size+1),margin_size:-margin_size]) - elif (tile_pp_info == [0, 1, 0, 1]).all(): - inst_on_margin.append(arr[ margin_size ,margin_size:-margin_size]) - inst_on_margin.append(arr[-(margin_size+1),margin_size:-margin_size]) - elif (tile_pp_info == [1, 0, 0, 0]).all() or \ - (tile_pp_info == [0, 1, 0, 0]).all() or \ - (tile_pp_info == [1, 1, 0, 0]).all() or \ - (tile_pp_info == [0, 0, 1, 1]).all() or \ - (tile_pp_info == [0, 0, 1, 0]).all() or \ - (tile_pp_info == [0, 0, 0, 1]).all(): - if tile_pp_info[0] == 1: - inst_on_margin.append(arr[:,margin_size]) - if tile_pp_info[1] == 1: - inst_on_margin.append(arr[:,-(margin_size+1)]) - if tile_pp_info[2] == 1: - inst_on_margin.append(arr[margin_size,:]) - if tile_pp_info[3] == 1: - inst_on_margin.append(arr[-(margin_size+1),:]) - else: - assert False - inst_on_margin = [v.flatten() for v in inst_on_margin] - - if len(inst_on_margin) > 0: - inst_on_margin = np.concatenate(inst_on_margin, axis=0) - inst_on_margin = np.unique(inst_on_margin) - else: - inst_on_margin = np.array([]) # empty array - - return inst_on_margin -# a = np.reshape(np.arange(0, 64), [8, 8]) -# print(a) -# print(get_inst_on_margin(a, 2, [1, 1 ,1, 1])) -# print(get_inst_on_margin(a, 1, [1, 1 ,1, 1])) -# print(get_inst_on_margin(a, 1, [1, 1 ,0, 0])) -# print(get_inst_on_margin(a, 1, [0, 0 ,1, 1])) -# print('') -# print(get_inst_in_margin(a, 1, [1, 1 ,1, 1])) -# print(get_inst_in_margin(a, 1, [0, 1 ,1, 1])) -# print(get_inst_in_margin(a, 1, [1, 0 ,1, 1])) -# print(get_inst_in_margin(a, 1, [0, 0 ,1, 1])) -# print(get_inst_in_margin(a, 1, [0, 1 ,1, 0])) -# exit() -#### def postproc_tile(tile_io_info, tile_pp_info, tile_mode, margin_size, patch_info_list, func_opt): # output pos of the tile within the source wsi @@ -473,7 +595,8 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, # |///////////////////////////| | # ----------------------------| V - + # tile_mode = 3 # no fixing debug + print(tile_pp_info) if tile_mode == 0: # for `full grid tile` # -- extend from the boundary by the margin size, remove @@ -482,7 +605,7 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, # also contain those lying on the edges inst_in_margin = get_inst_in_margin(pred_inst, margin_size, tile_pp_info) # those lying on the margin line - inst_on_margin = get_inst_on_margin(pred_inst, margin_size, tile_pp_info) + inst_on_margin = get_inst_on_margin(pred_inst, margin_size-1, tile_pp_info) inst_within_margin = np.setdiff1d(inst_in_margin, inst_on_margin, assume_unique=True) remove_inst_set = inst_within_margin.tolist() elif tile_mode == 1 or tile_mode == 2: @@ -490,7 +613,7 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, # -- extend from the marked edges (top/bot or left/right) by the margin size, # remove all nuclei lie within the margin area (including on the margin line) # -- nuclei on all edges are removed (as these are alrd within `full grid tile`) - inst_in_margin = get_inst_in_margin(margin_size, tile_pp_info) # also contain those lying on the edges + inst_in_margin = get_inst_in_margin(pred_inst, margin_size, tile_pp_info) # also contain those lying on the edges if np.sum(tile_pp_info) == 1: holder_flag = tile_pp_info.copy() if tile_mode == 1: @@ -499,15 +622,13 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, holder_flag[[0, 1]] = 1 else: holder_flag = [1, 1, 1, 1] - print(tile_mode, tile_pp_info, holder_flag) - inst_on_edge = get_inst_on_margin(0, holder_flag) + inst_on_edge = get_inst_on_edge(pred_inst, holder_flag) remove_inst_set = np.union1d(inst_in_margin, inst_on_edge) remove_inst_set = remove_inst_set.tolist() else: # inst within the tile after excluding margin area out # only for a tile at cross-section, which is designed such that # their shape >= 3* margin size - all_inst = np.unique() remove_inst_set = [] remove_inst_set = set(remove_inst_set) From 801f41dae5d6cea0097cdf335edb3c1e3ff67e8b Mon Sep 17 00:00:00 2001 From: vqdang Date: Mon, 8 Mar 2021 14:13:41 +0000 Subject: [PATCH 09/19] UPD: add 1 pix off to inst on to resolve inclusion --- infer/super_wsi.py | 134 ++++++++------------------------------------- 1 file changed, 23 insertions(+), 111 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index 4c82df2c..68f991ad 100644 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -268,7 +268,7 @@ def get_info_stack(output_tl, output_br): removal_flag = np.full((xsect_info_list.shape[0], 4,), 0) # left, right, top, bot # all_info.append([xsect_info_list, removal_flag, mode_list]) - all_info = all_info[:1] + # all_info = all_info[:1] # ! combine all info_list, removal_flag, mode_list = list(zip(*all_info)) info_list = np.concatenate(info_list, axis=0).astype(np.int32) @@ -288,13 +288,13 @@ def get_inst_in_margin(arr, margin_size, tile_pp_info): inst_in_margin = [] # extract those lie within margin region if tile_pp_info[0] == 1: # left edge - inst_in_margin.append(arr[:,:margin_size]) + inst_in_margin.append(arr[:,:(margin_size+1)]) if tile_pp_info[1] == 1: # right edge - inst_in_margin.append(arr[:,-margin_size:]) + inst_in_margin.append(arr[:,-(margin_size+1):]) if tile_pp_info[2] == 1: # top edge - inst_in_margin.append(arr[:margin_size,:]) + inst_in_margin.append(arr[:(margin_size+1),:]) if tile_pp_info[3] == 1: # bottom edge - inst_in_margin.append(arr[-margin_size:,:]) + inst_in_margin.append(arr[-(margin_size+1):,:]) inst_in_margin = [v.flatten() for v in inst_in_margin] if len(inst_in_margin) > 0: inst_in_margin = np.concatenate(inst_in_margin, axis=0) @@ -302,6 +302,7 @@ def get_inst_in_margin(arr, margin_size, tile_pp_info): else: inst_in_margin = np.array([]) # empty array return inst_in_margin + #### def get_inst_on_margin(arr, margin_size, tile_pp_info): """ @@ -357,41 +358,27 @@ def det(a, b): line_intersection(line_list[3], line_list[1]), # x4 ] - def sel_between_pts(p1, p2): - arr[p1[0]:p2[0]+1, - p1[1]:p2[1]+1] = 1 - + pt_index = lambda p1, p2: arr[p1[0]:p2[0]+1, + p1[1]:p2[1]+1] line_pix_list = [] if tile_pp_info[0] == 1: - sel_between_pts(pts_list[0], pts_list[2]) + line_pix_list.append(pt_index(pts_list[0], pts_list[2])), if tile_pp_info[1] == 1: - sel_between_pts(pts_list[1], pts_list[3]) + line_pix_list.append(pt_index(pts_list[1], pts_list[3])), if tile_pp_info[2] == 1: - sel_between_pts(pts_list[0], pts_list[1]) + line_pix_list.append(pt_index(pts_list[0], pts_list[1])), if tile_pp_info[3] == 1: - sel_between_pts(pts_list[2], pts_list[3]) - - # pt_index = lambda p1, p2: arr[p1[0]:p2[0]+1, - # p1[1]:p2[1]+1] - # line_pix_list = [] - # if tile_pp_info[0] == 1: - # line_pix_list.append(pt_index(pts_list[0], pts_list[2])), - # if tile_pp_info[1] == 1: - # line_pix_list.append(pt_index(pts_list[1], pts_list[3])), - # if tile_pp_info[2] == 1: - # line_pix_list.append(pt_index(pts_list[0], pts_list[1])), - # if tile_pp_info[3] == 1: - # line_pix_list.append(pt_index(pts_list[2], pts_list[3])), - - # inst_on_margin = [v.flatten() for v in line_pix_list] - - # if len(inst_on_margin) > 0: - # inst_on_margin = np.concatenate(inst_on_margin, axis=0) - # inst_on_margin = np.unique(inst_on_margin) - # else: - # inst_on_margin = np.array([]) # empty array - - # return inst_on_margin + line_pix_list.append(pt_index(pts_list[2], pts_list[3])), + + inst_on_margin = [v.flatten() for v in line_pix_list] + + if len(inst_on_margin) > 0: + inst_on_margin = np.concatenate(inst_on_margin, axis=0) + inst_on_margin = np.unique(inst_on_margin) + else: + inst_on_margin = np.array([]) # empty array + + return inst_on_margin #### def get_inst_on_edge(arr, tile_pp_info): inst_on_edge = [] @@ -413,81 +400,7 @@ def get_inst_on_edge(arr, tile_pp_info): inst_on_edge = np.array([]) # empty array return inst_on_edge #### -# a = np.reshape(np.arange(0, 64), [8, 8]) -# print(a) -# print(get_inst_on_margin(a, 2, [1, 1 ,1, 1])) -# print(get_inst_on_margin(a, 1, [1, 1 ,1, 1])) -# print(get_inst_on_margin(a, 1, [1, 1 ,0, 0])) -# print(get_inst_on_margin(a, 1, [0, 0 ,1, 1])) -# print(get_inst_on_margin(a, 1, [0, 1 ,1, 0])) -# print(get_inst_on_margin(a, 1, [0, 1 ,1, 1])) -# print(get_inst_on_margin(a, 1, [1, 1 ,0, 1])) -# print('') -# print(get_inst_in_margin(a, 1, [1, 1 ,1, 1])) -# print(get_inst_in_margin(a, 1, [0, 1 ,1, 1])) -# print(get_inst_in_margin(a, 1, [1, 0 ,1, 1])) -# print(get_inst_in_margin(a, 1, [0, 0 ,1, 1])) -# print(get_inst_in_margin(a, 1, [0, 1 ,1, 0])) -# exit() -#### -tile_shape = np.array([2048, 2048]) -image_shape = np.array([3600, 3600]) -patch_input_shape = np.array([256,256]) -patch_output_shape = np.array([164,164]) -ambiguous_size = patch_output_shape * 2 -patch_info_list = _get_patch_info(image_shape, patch_input_shape, patch_output_shape) - -patch_diff_shape = patch_input_shape - patch_output_shape -# derive tile output placement as consecutive tiling with step size of 0 -# and tile output will have shape being of multiple of patch_output_shape (round down) -tile_output_shape = np.floor(tile_shape / patch_output_shape) * patch_output_shape -tile_input_shape = tile_output_shape + patch_diff_shape - -tile_io_info_list, \ - tile_pp_info_list, \ - tile_mode_list = _get_tile_info( - image_shape, tile_input_shape, tile_output_shape, - ambiguous_size, patch_output_shape -) - -import matplotlib.pyplot as plt -cmap = plt.get_cmap('jet') - -# canvas = np.zeros(image_shape) -# nr_patch = patch_info_list.shape[0] -# for idx in range(0, nr_patch): -# tl, br = patch_info_list[idx][1] -# # canvas[tl[0]:br[0],tl[1]:br[1]] += (idx+1) -# patch = canvas[tl[0]:br[0],tl[1]:br[1]] -# patch[0,:] = 1 -# patch[-1,:] = 1 -# patch[:,0] = 1 -# patch[:,-1] = 1 -# # canvas[tl[0]:br[0],tl[1]:br[1]] += (idx+1) - -# canvas_color = (cmap(canvas / np.max(canvas)) * 255).astype(np.uint8)[...,:3] -# canvas_color[canvas==0] = 0 # background -# cv2.imwrite('patch.png', cv2.cvtColor(canvas_color, cv2.COLOR_RGB2BGR)) - -canvas = np.zeros(image_shape) -nr_tile = tile_io_info_list.shape[0] -for idx in range(0, nr_tile): - tl, br = tile_io_info_list[idx][1] - canvas[tl[0]:br[0],tl[1]:br[1]] += (idx+1) -canvas_color = (cmap(canvas / np.max(canvas)) * 255).astype(np.uint8)[...,:3] -canvas_color[canvas==0] = 0 # background -cv2.imwrite('tile.png', cv2.cvtColor(canvas_color, cv2.COLOR_RGB2BGR)) - -canvas = np.zeros(image_shape) -for idx in range(0, nr_tile): - tl, br = tile_io_info_list[idx][1] - print(tl, br, tile_pp_info_list[idx]) - sub_canvas = canvas[tl[0]:br[0],tl[1]:br[1]] - canvas[tl[0]:br[0],tl[1]:br[1]] = get_inst_on_margin(sub_canvas, ambiguous_size[0], tile_pp_info_list[idx]) -canvas_color = (cmap(canvas / np.max(canvas)) * 255).astype(np.uint8)[...,:3] -canvas_color[canvas==0] = 0 # background -cv2.imwrite('tile_fix.png', cv2.cvtColor(canvas_color, cv2.COLOR_RGB2BGR)) -exit() + #### def run_model( forward_output_queue, @@ -596,7 +509,6 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, # ----------------------------| V # tile_mode = 3 # no fixing debug - print(tile_pp_info) if tile_mode == 0: # for `full grid tile` # -- extend from the boundary by the margin size, remove From e137e05eb6dd3778b78b6757674f2c7c878fdc66 Mon Sep 17 00:00:00 2001 From: vqdang Date: Mon, 8 Mar 2021 18:47:39 +0000 Subject: [PATCH 10/19] UPD: xsect fix complete --- infer/super_wsi.py | 255 +++++++++++++++++++++++++++++---------------- 1 file changed, 166 insertions(+), 89 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index 68f991ad..246cc55f 100644 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -4,7 +4,6 @@ mp.set_start_method("spawn", True) # ! must be at top for VScode debugging -import argparse import collections import glob import json @@ -12,6 +11,7 @@ import math import os import pathlib +import copy import re import shutil import sys @@ -210,7 +210,7 @@ def get_info_stack(output_tl, output_br): y_fix_output_br[y_fix_output_br[:,1] > br_most[1], 1] = br_most[1] # sel position not on the image boundary sel = (output_tl[:,0] == np.min(output_tl[:,0])) - y_info_list = get_info_stack(y_fix_output_tl[~sel], y_fix_output_br[~sel]) + y_info_list = get_info_stack(y_fix_output_tl[~sel], y_fix_output_br[~sel]).astype(np.int64) # flag horizontal ambiguous region for y (left margin, right margin) # |----|------------|----| @@ -234,7 +234,7 @@ def get_info_stack(output_tl, output_br): x_fix_output_br[x_fix_output_br[:,1] > br_most[1], 1] = br_most[1] # sel position not on the image boundary sel = (output_br[:,1] == np.max(output_br[:,1])) - x_info_list = get_info_stack(x_fix_output_tl[~sel], x_fix_output_br[~sel]) + x_info_list = get_info_stack(x_fix_output_tl[~sel], x_fix_output_br[~sel]).astype(np.int64) # flag vertical ambiguous region for x (top margin, bottom margin) # |----| ^ # |\\\\| | ambiguous @@ -242,7 +242,7 @@ def get_info_stack(output_tl, output_br): # | | # | | # |----| - # |\\\\| ambiguous + # |\\\\| # |----| removal_flag = np.zeros((x_info_list.shape[0], 4,)) # left, right, top, bot removal_flag[:,[2,3]] = 1 @@ -253,7 +253,6 @@ def get_info_stack(output_tl, output_br): mode_list = np.full(x_info_list.shape[0], 2) all_info.append([x_info_list, removal_flag, mode_list]) - # * define the tile cross section sel = np.any(output_br == br_most, axis=-1) xsect = output_br[~sel] xsect_tl = xsect - margin_size * 2 @@ -263,19 +262,12 @@ def get_info_stack(output_tl, output_br): xsect_br[xsect_br[:,1] > br_most[1], 1] = br_most[1] xsect_tl[xsect_tl[:,0] < tl_most[0], 0] = tl_most[0] xsect_tl[xsect_tl[:,1] < tl_most[1], 1] = tl_most[1] - xsect_info_list = get_info_stack(xsect_tl, xsect_br) + xsect_info_list = get_info_stack(xsect_tl, xsect_br).astype(np.int64) mode_list = np.full(xsect_info_list.shape[0], 3) removal_flag = np.full((xsect_info_list.shape[0], 4,), 0) # left, right, top, bot - # all_info.append([xsect_info_list, removal_flag, mode_list]) + all_info.append([xsect_info_list, removal_flag, mode_list]) - # all_info = all_info[:1] - # ! combine all - info_list, removal_flag, mode_list = list(zip(*all_info)) - info_list = np.concatenate(info_list, axis=0).astype(np.int32) - mode_list = np.concatenate(mode_list, axis=0) - removal_flag = np.concatenate(removal_flag, axis=0) - - return info_list, removal_flag, mode_list + return all_info #### # ! seem to be 1 pix off at cross section or sthg def get_inst_in_margin(arr, margin_size, tile_pp_info): @@ -455,7 +447,8 @@ def run_model( return #### def postproc_tile(tile_io_info, tile_pp_info, tile_mode, - margin_size, patch_info_list, func_opt): + margin_size, patch_info_list, func_opt, + prev_wsi_inst_dict=None): # output pos of the tile within the source wsi tile_input_tl, tile_output_br = tile_io_info[0] tile_output_tl, tile_output_br = tile_io_info[1] # Y, X @@ -508,7 +501,31 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, # |///////////////////////////| | # ----------------------------| V - # tile_mode = 3 # no fixing debug + def draw_prev_pred_inst(): + wsi_inst_uid_list = np.array(list(prev_wsi_inst_dict.keys())) + wsi_inst_com_list = np.array([v['centroid'] for v in prev_wsi_inst_dict.values()]) + wsi_inst_com_list = wsi_inst_com_list[:,::-1] # XY to YX + tile_output = tile_io_info[1] + sel = (wsi_inst_com_list[:,0] > tile_output[0,0]) + sel &= (wsi_inst_com_list[:,0] < tile_output[1,0]) + sel &= (wsi_inst_com_list[:,1] > tile_output[0,1]) + sel &= (wsi_inst_com_list[:,1] < tile_output[1,1]) + sel_idx = np.nonzero(sel.flatten())[0] + + tile_canvas = np.zeros(tile_output[1] - tile_output[0], dtype=np.int32) + for inst_idx in sel_idx: + inst_uid = wsi_inst_uid_list[inst_idx] + # shift from wsi system to tile output system + inst_info = prev_wsi_inst_dict[inst_uid] + inst_cnt = np.array(inst_info['contour']) + inst_cnt = inst_cnt - tile_output[0,[1,0]][None] + tile_canvas = cv2.drawContours(tile_canvas, [inst_cnt], -1, int(inst_uid), -1) + return tile_canvas + + # ! for correctness, cross section need specific dedicated sweep wrt to entire wsi to + # ! remove those lie within, or for simplicity, need caching the idx location of 4 corner + # ! tile to minimize the search effort ! + output_remove_inst_set = None if tile_mode == 0: # for `full grid tile` # -- extend from the boundary by the margin size, remove @@ -516,8 +533,10 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, # lie on the margin line # also contain those lying on the edges inst_in_margin = get_inst_in_margin(pred_inst, margin_size, tile_pp_info) - # those lying on the margin line - inst_on_margin = get_inst_on_margin(pred_inst, margin_size-1, tile_pp_info) + # those lying on the margin line, check 2pix toward margin area for sanity + inst_on_margin1 = get_inst_on_margin(pred_inst, margin_size-1, tile_pp_info) + inst_on_margin2 = get_inst_on_margin(pred_inst, margin_size , tile_pp_info) + inst_on_margin = np.union1d(inst_on_margin1, inst_on_margin2) inst_within_margin = np.setdiff1d(inst_in_margin, inst_on_margin, assume_unique=True) remove_inst_set = inst_within_margin.tolist() elif tile_mode == 1 or tile_mode == 2: @@ -540,11 +559,25 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, else: # inst within the tile after excluding margin area out # only for a tile at cross-section, which is designed such that - # their shape >= 3* margin size - remove_inst_set = [] + # their shape >= 3* margin size + + # ! for removal of current pred, just like tile_mode = 0 but with all side + inst_in_margin = get_inst_in_margin(pred_inst, margin_size, [1, 1, 1, 1]) + # those lying on the margin line, check 2pix toward margin area for sanity + inst_on_margin1 = get_inst_on_margin(pred_inst, margin_size-1, [1, 1, 1, 1]) + inst_on_margin2 = get_inst_on_margin(pred_inst, margin_size , [1, 1, 1, 1]) + inst_on_margin = np.union1d(inst_on_margin1, inst_on_margin2) + inst_within_margin = np.setdiff1d(inst_in_margin, inst_on_margin, assume_unique=True) + remove_inst_set = inst_within_margin.tolist() + # ! but we also need to remove prev inst exising on the global space of entire wsi + # ! on the margin + prev_pred_inst = draw_prev_pred_inst() + inst_on_margin1 = get_inst_on_margin(prev_pred_inst, margin_size-1, [1, 1, 1, 1]) + inst_on_margin2 = get_inst_on_margin(prev_pred_inst, margin_size , [1, 1, 1, 1]) + inst_on_margin = np.union1d(inst_on_margin1, inst_on_margin2) + output_remove_inst_set = inst_on_margin.tolist() remove_inst_set = set(remove_inst_set) - # * move pos back to wsi position renew_id = 0 new_inst_info_dict = {} @@ -555,8 +588,8 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, v['contour'] += tile_output_tl[::-1] new_inst_info_dict[renew_id] = v renew_id += 1 - return new_inst_info_dict + return new_inst_info_dict, output_remove_inst_set #### class InferManager(base.InferManager): @@ -678,9 +711,8 @@ def process_single_file(self, wsi_path, mask_path, output_dir): tile_output_shape = np.floor(self.tile_shape / self.patch_output_shape) * self.patch_output_shape tile_input_shape = tile_output_shape + patch_diff_shape - tile_io_info_list, \ - tile_pp_info_list, \ - tile_mode_list = _get_tile_info( + # [full_grid, vert/horiz, xsect] + all_tile_info = _get_tile_info( self.wsi_proc_shape, tile_input_shape, tile_output_shape, self.ambiguous_size, self.patch_output_shape ) @@ -695,75 +727,120 @@ def process_single_file(self, wsi_path, mask_path, output_dir): # \ postproc (loop)/ patch_info_list = self.__select_valid_patches(patch_info_list) - tile_io_info_list = self.__select_valid_patches(tile_io_info_list) - - mp_manager = torch_mp.Manager() - # contain at most 5 tile ouput before polling forward func - mp_forward_output_queue = mp_manager.Queue(maxsize=5) - forward_info_list = collections.deque() + def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, + prev_wsi_inst_dict=None): - nr_tile = tile_io_info_list.shape[0] - for tile_info in tile_io_info_list: - # retrieve valid patch within tile - patch_in_tile_info_list = self.__select_patches_in_tile(tile_info, patch_info_list) - forward_info_list.append([tile_info, patch_in_tile_info_list]) + tile_io_info_list = self.__select_valid_patches(tile_io_info_list) + + mp_manager = torch_mp.Manager() + # contain at most 5 tile ouput before polling forward func + mp_forward_output_queue = mp_manager.Queue(maxsize=5) - loader_kwargs = dict( - num_workers=self.nr_inference_workers, - batch_size=self.batch_size, - ) + forward_info_list = collections.deque() - wsi_cache_path = "%s/src_wsi.npy" % self.cache_path - forward_process = mp.Process(target=run_model, - args=(mp_forward_output_queue, forward_info_list, - wsi_path, wsi_ext, self.proc_mag, wsi_cache_path, - self.run_step, self.net, loader_kwargs)) - forward_process.start() - - post_proc_kwargs = { - "nr_types": self.method["model_args"]["nr_types"], - "return_centroids": True, - } - - self.wsi_inst_info = {} - future_list = collections.deque() - proc_pool = ProcessPoolExecutor(self.nr_post_proc_workers) - # will this lead to infinite loop ? - while True: - if forward_process.exitcode is not None \ - and mp_forward_output_queue.empty(): - break - - if not mp_forward_output_queue.empty(): - # ! assume the forward result are spit out in - # ! same order as defined in `forward_info_list` - tile_idx, forward_output = mp_forward_output_queue.get() - tile_io_info = tile_io_info_list[tile_idx] - tile_pp_info = tile_pp_info_list[tile_idx] - tile_mode = tile_mode_list[tile_idx] + nr_tile = tile_io_info_list.shape[0] + for tile_info in tile_io_info_list: + # retrieve valid patch within tile + patch_in_tile_info_list = self.__select_patches_in_tile(tile_info, patch_info_list) + forward_info_list.append([tile_info, patch_in_tile_info_list]) + loader_kwargs = dict( + num_workers=self.nr_inference_workers, + batch_size=self.batch_size, + ) + + wsi_cache_path = "%s/src_wsi.npy" % self.cache_path + forward_process = mp.Process(target=run_model, + args=(mp_forward_output_queue, forward_info_list, + wsi_path, wsi_ext, self.proc_mag, wsi_cache_path, + self.run_step, self.net, loader_kwargs)) + forward_process.start() + + post_proc_kwargs = { + "nr_types": self.method["model_args"]["nr_types"], + "return_centroids": True, + } + + wsi_inst_info = {} + if prev_wsi_inst_dict is not None: + wsi_inst_info = copy.deepcopy(prev_wsi_inst_dict) + + future_list = collections.deque() + proc_pool = ProcessPoolExecutor(self.nr_post_proc_workers) + # will this lead to infinite loop ? + offset_id = 0 + while True: + if forward_process.exitcode is not None \ + and mp_forward_output_queue.empty(): + break + + if not mp_forward_output_queue.empty(): + # ! assume the forward result are spit out in + # ! assume the forward result are spit out in + # ! assume the forward result are spit out in + # ! same order as defined in `forward_info_list` + tile_idx, forward_output = mp_forward_output_queue.get() + tile_io_info = tile_io_info_list[tile_idx] + tile_pp_info = tile_pp_info_list[tile_idx] + tile_mode = tile_mode_list[tile_idx] + + # future = proc_pool.submit(postproc_tile, # future = proc_pool.submit(postproc_tile, + # future = proc_pool.submit(postproc_tile, + # tile_info, forward_output, # tile_info, forward_output, - # (self.post_proc_func, post_proc_kwargs)) - - tile_inst_dict = postproc_tile( - tile_io_info, tile_pp_info, tile_mode, - self.ambiguous_size[0], - forward_output, - (self.post_proc_func, post_proc_kwargs)) - - # ! may not work if output id range > maximum #inst in return dict - offset_id = len(self.wsi_inst_info) - for tile_inst_id, tile_inst_info in tile_inst_dict.items(): - inst_wsi_id = offset_id + tile_inst_id + 1 - self.wsi_inst_info[inst_wsi_id] = tile_inst_info - - print('Post proc %d' % tile_idx) - # future_list.append(future) # deal when forward finish or stick callback ? - if forward_process.exitcode > 0: - raise ValueError(f'Forward process exited with code {forward_process.exitcode}') - forward_process.join() + # tile_info, forward_output, + # (self.post_proc_func, post_proc_kwargs)) + + new_inst_dict, \ + remove_uid_list = postproc_tile( + tile_io_info, tile_pp_info, tile_mode, + self.ambiguous_size[0], forward_output, + (self.post_proc_func, post_proc_kwargs), + prev_wsi_inst_dict) + + offset_id = max(offset_id, max(new_inst_dict.keys())) + for inst_id, inst_info in new_inst_dict.items(): + inst_wsi_id = offset_id + inst_id + 1 + wsi_inst_info[inst_wsi_id] = inst_info + + for inst_uid in remove_uid_list: + if inst_uid in wsi_inst_info: + wsi_inst_info.pop(inst_uid) + + print('Post proc %d' % tile_idx) + # future_list.append(future) # deal when forward finish or stick callback ? + if forward_process.exitcode > 0: + raise ValueError(f'Forward process exited with code {forward_process.exitcode}') + forward_process.join() + return wsi_inst_info + + # + # process full grid and vert/horiz fixing at the same time + # info_list = list(zip(*all_tile_info[:3])) + # tile_io_info_list = np.concatenate(info_list[0], axis=0).astype(np.int32) + # tile_pp_info_list = np.concatenate(info_list[1], axis=0) + # tile_mode_list = np.concatenate(info_list[2], axis=0) + # wsi_inst_info = run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list) + # print('here') + import joblib + wsi_inst_info = joblib.load('cache_output.dat') + + + # ! is this costly ?, may need to dispatch multiproc + tile_io_info_list = all_tile_info[-1][0] + tile_pp_info_list = all_tile_info[-1][1] + tile_mode_list = all_tile_info[-1][2] + wsi_inst_info = run_once(tile_io_info_list, + tile_pp_info_list, + tile_mode_list, + wsi_inst_info) + + # move search warrant for each tile into its own postproc call + + # retrieve all instance and check which instance belong to which xsect tile + # then dispacth re-inference and postproc again, inst needed to be copied # while len(future_list) > 0: # if not future_list[0].done(): @@ -781,7 +858,7 @@ def process_single_file(self, wsi_path, mask_path, output_dir): # self.__save_json(json_path, self.wsi_inst_info) wsi_thumb_rgb = self.wsi_handler.get_full_img(read_mag=self.proc_mag) from misc.viz_utils import visualize_instances_dict - wsi_overlay = visualize_instances_dict(wsi_thumb_rgb, self.wsi_inst_info, draw_dot=True) + wsi_overlay = visualize_instances_dict(wsi_thumb_rgb, wsi_inst_info, draw_dot=True) cv2.imwrite('dump.png', cv2.cvtColor(wsi_overlay, cv2.COLOR_RGB2BGR)) return From 5b68e2f972b01523fabe0b0725a3917d7f8dba4e Mon Sep 17 00:00:00 2001 From: vqdang Date: Tue, 9 Mar 2021 13:42:48 +0000 Subject: [PATCH 11/19] UPD: finish integrate multiproc forward/postproc --- infer/super_wsi.py | 173 ++++++++++++++++++++++++++------------------- 1 file changed, 99 insertions(+), 74 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index 246cc55f..f483334f 100644 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -442,9 +442,8 @@ def run_model( sample_info_list = sample_info_list.numpy() accumulated_patch_output.append([sample_info_list, sample_output_list]) forward_output_queue.put([tile_idx, accumulated_patch_output]) - - print('%d/%d' % (tile_idx, len(tile_info_list))) return + #### def postproc_tile(tile_io_info, tile_pp_info, tile_mode, margin_size, patch_info_list, func_opt, @@ -522,10 +521,8 @@ def draw_prev_pred_inst(): tile_canvas = cv2.drawContours(tile_canvas, [inst_cnt], -1, int(inst_uid), -1) return tile_canvas - # ! for correctness, cross section need specific dedicated sweep wrt to entire wsi to - # ! remove those lie within, or for simplicity, need caching the idx location of 4 corner - # ! tile to minimize the search effort ! output_remove_inst_set = None + # tile_mode = -1 # ! no fix, debud mode if tile_mode == 0: # for `full grid tile` # -- extend from the boundary by the margin size, remove @@ -556,7 +553,7 @@ def draw_prev_pred_inst(): inst_on_edge = get_inst_on_edge(pred_inst, holder_flag) remove_inst_set = np.union1d(inst_in_margin, inst_on_edge) remove_inst_set = remove_inst_set.tolist() - else: + elif tile_mode == 3: # inst within the tile after excluding margin area out # only for a tile at cross-section, which is designed such that # their shape >= 3* margin size @@ -576,6 +573,8 @@ def draw_prev_pred_inst(): inst_on_margin2 = get_inst_on_margin(prev_pred_inst, margin_size , [1, 1, 1, 1]) inst_on_margin = np.union1d(inst_on_margin1, inst_on_margin2) output_remove_inst_set = inst_on_margin.tolist() + else: + remove_inst_set = [] remove_inst_set = set(remove_inst_set) # * move pos back to wsi position @@ -593,7 +592,7 @@ def draw_prev_pred_inst(): #### class InferManager(base.InferManager): - def __select_valid_patches(self, patch_info_list): + def __get_valid_patch_idx(self, patch_info_list): """Select valid patches from the list of input patch information. Args: @@ -615,16 +614,16 @@ def check_valid(info, wsi_mask): for info in patch_info_list] # somehow multiproc is slower than single thread valid_indices = np.array(valid_indices) - return patch_info_list[valid_indices] + return valid_indices - def __select_patches_in_tile(self, tile_info, patch_info_list): + def __get_valid_patch_idx_in_tile(self, tile_info, patch_info_list): # checking basing on the output alignment tile_tl, tile_br = tile_info[1] patch_tl_list = patch_info_list[:,1,0] patch_br_list = patch_info_list[:,1,1] sel = (patch_tl_list[:,0] >= tile_tl[0]) & (patch_tl_list[:,1] >= tile_tl[1]) sel &= (patch_br_list[:,0] <= tile_br[0]) & (patch_br_list[:,1] <= tile_br[1]) - return patch_info_list[sel] + return sel def _parse_args(self, run_args): """Parse command line arguments and set as instance variables.""" @@ -726,23 +725,31 @@ def process_single_file(self, wsi_path, mask_path, output_dir): # main ------------- main----------------main # \ postproc (loop)/ - patch_info_list = self.__select_valid_patches(patch_info_list) + sel_index = self.__get_valid_patch_idx(patch_info_list) + patch_info_list = patch_info_list[sel_index] + + pbar_creator = lambda x, y, pos=0, leave=False: tqdm.tqdm( + desc=x, leave=leave, total=y, ncols=80, ascii=True, position=pos + ) def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, prev_wsi_inst_dict=None): - tile_io_info_list = self.__select_valid_patches(tile_io_info_list) + sel_index = self.__get_valid_patch_idx(tile_io_info_list) + tile_io_info_list = tile_io_info_list[sel_index] + tile_pp_info_list = tile_pp_info_list[sel_index] + tile_mode_list = tile_mode_list[sel_index] mp_manager = torch_mp.Manager() # contain at most 5 tile ouput before polling forward func mp_forward_output_queue = mp_manager.Queue(maxsize=5) forward_info_list = collections.deque() - nr_tile = tile_io_info_list.shape[0] - for tile_info in tile_io_info_list: - # retrieve valid patch within tile - patch_in_tile_info_list = self.__select_patches_in_tile(tile_info, patch_info_list) + for tile_idx in range(nr_tile): + tile_info = tile_io_info_list[tile_idx] + sel_index = self.__get_valid_patch_idx_in_tile(tile_info, patch_info_list) + patch_in_tile_info_list = patch_info_list[sel_index] forward_info_list.append([tile_info, patch_in_tile_info_list]) loader_kwargs = dict( @@ -762,94 +769,109 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, "return_centroids": True, } + proc_pool = None + # if self.nr_post_proc_workers > 0: + # proc_pool = ProcessPoolExecutor(self.nr_post_proc_workers) + + offset_id = 0 # offset id to increment overall saving dict wsi_inst_info = {} if prev_wsi_inst_dict is not None: wsi_inst_info = copy.deepcopy(prev_wsi_inst_dict) + offset_id = max(wsi_inst_info.keys()) + 1 + + foward_pbar = pbar_creator('Forward', nr_tile, pos=0) + postpr_pbar = pbar_creator('PostPro', nr_tile, pos=1) future_list = collections.deque() - proc_pool = ProcessPoolExecutor(self.nr_post_proc_workers) # will this lead to infinite loop ? - offset_id = 0 while True: if forward_process.exitcode is not None \ and mp_forward_output_queue.empty(): break if not mp_forward_output_queue.empty(): - # ! assume the forward result are spit out in - # ! assume the forward result are spit out in - # ! assume the forward result are spit out in - # ! same order as defined in `forward_info_list` tile_idx, forward_output = mp_forward_output_queue.get() + foward_pbar.update() + tile_io_info = tile_io_info_list[tile_idx] tile_pp_info = tile_pp_info_list[tile_idx] tile_mode = tile_mode_list[tile_idx] + args = [tile_io_info, tile_pp_info, tile_mode, + self.ambiguous_size[0], forward_output, + (self.post_proc_func, post_proc_kwargs), + prev_wsi_inst_dict] + + if proc_pool is not None: + future = proc_pool.submit(postproc_tile, *args) + future_list.append(future) + else: + new_inst_dict, remove_uid_list = postproc_tile(*args) + + # * aggregate + # ! the return id should be contiguous to maximuize + # ! counting range in int32 + for inst_id, inst_info in new_inst_dict.items(): + inst_wsi_id = offset_id + inst_id + 1 + wsi_inst_info[inst_wsi_id] = inst_info + offset_id = inst_wsi_id + 1 + + if remove_uid_list is not None: + for inst_uid in remove_uid_list: + if inst_uid in wsi_inst_info: + wsi_inst_info.pop(inst_uid) + postpr_pbar.update() - # future = proc_pool.submit(postproc_tile, - # future = proc_pool.submit(postproc_tile, - # future = proc_pool.submit(postproc_tile, - # tile_info, forward_output, - # tile_info, forward_output, - # tile_info, forward_output, - # (self.post_proc_func, post_proc_kwargs)) - - new_inst_dict, \ - remove_uid_list = postproc_tile( - tile_io_info, tile_pp_info, tile_mode, - self.ambiguous_size[0], forward_output, - (self.post_proc_func, post_proc_kwargs), - prev_wsi_inst_dict) - - offset_id = max(offset_id, max(new_inst_dict.keys())) - for inst_id, inst_info in new_inst_dict.items(): - inst_wsi_id = offset_id + inst_id + 1 - wsi_inst_info[inst_wsi_id] = inst_info + if forward_process.exitcode > 0: + raise ValueError(f'Forward process exited with code {forward_process.exitcode}') + forward_process.join() + while len(future_list) > 0: + if not future_list[0].done(): + future_list.rotate() + continue + proc_future = future_list.popleft() + if proc_future.exception() is not None: + print(proc_future.exception()) + + # * aggregate + # ! the return id should be contiguous to maximuize + # ! counting range in int32 + new_inst_dict, remove_uid_list = proc_future.result() + for inst_id, inst_info in new_inst_dict.items(): + inst_wsi_id = offset_id + inst_id + 1 + wsi_inst_info[inst_wsi_id] = inst_info + offset_id = inst_wsi_id + 1 + + if remove_uid_list is not None: for inst_uid in remove_uid_list: if inst_uid in wsi_inst_info: wsi_inst_info.pop(inst_uid) + postpr_pbar.update() + foward_pbar.close() + postpr_pbar.close() - print('Post proc %d' % tile_idx) - # future_list.append(future) # deal when forward finish or stick callback ? - if forward_process.exitcode > 0: - raise ValueError(f'Forward process exited with code {forward_process.exitcode}') - forward_process.join() return wsi_inst_info # - # process full grid and vert/horiz fixing at the same time - # info_list = list(zip(*all_tile_info[:3])) - # tile_io_info_list = np.concatenate(info_list[0], axis=0).astype(np.int32) - # tile_pp_info_list = np.concatenate(info_list[1], axis=0) - # tile_mode_list = np.concatenate(info_list[2], axis=0) - # wsi_inst_info = run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list) - # print('here') - import joblib - wsi_inst_info = joblib.load('cache_output.dat') - - - # ! is this costly ?, may need to dispatch multiproc + ## ** process full grid and vert/horiz fixing at the same time + info_list = list(zip(*all_tile_info[:3])) + tile_io_info_list = np.concatenate(info_list[0], axis=0).astype(np.int32) + tile_pp_info_list = np.concatenate(info_list[1], axis=0) + tile_mode_list = np.concatenate(info_list[2], axis=0) + wsi_inst_info = run_once(tile_io_info_list, + tile_pp_info_list, + tile_mode_list) + # import joblib + # wsi_inst_info = joblib.load('cache_output.dat') + + ## ** re-infer and redo postproc for xsect alone tile_io_info_list = all_tile_info[-1][0] tile_pp_info_list = all_tile_info[-1][1] tile_mode_list = all_tile_info[-1][2] wsi_inst_info = run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, - wsi_inst_info) - - # move search warrant for each tile into its own postproc call - - # retrieve all instance and check which instance belong to which xsect tile - # then dispacth re-inference and postproc again, inst needed to be copied - - # while len(future_list) > 0: - # if not future_list[0].done(): - # future_list.rotate() - # continue - # proc_future = future_list.popleft() - # if proc_future.exception() is not None: - # print(proc_future.exception()) - # proc_future.result() + wsi_inst_info) if self.save_mask or self.save_thumb: json_path = "%s/json/%s.json" % (output_dir, wsi_name) @@ -886,8 +908,11 @@ def process_wsi_list(self, run_args): wsi_path_list = glob.glob(self.input_dir + "/*") wsi_path_list.sort() # ensure ordering - for wsi_path in wsi_path_list[1:2]: + for wsi_path in wsi_path_list: wsi_base_name = pathlib.Path(wsi_path).stem + + if wsi_base_name != 'mini_2': continue + msk_path = "%s/%s.png" % (self.input_mask_dir, wsi_base_name) if self.save_thumb or self.save_mask: output_file = "%s/json/%s.json" % (self.output_dir, wsi_base_name) From 1ba04a5a13fb48d90123131c5b7c883b878ae807 Mon Sep 17 00:00:00 2001 From: vqdang Date: Tue, 9 Mar 2021 14:24:16 +0000 Subject: [PATCH 12/19] UPD: uncomment --- infer/super_wsi.py | 54 ++++++++++++++++++++++++++-------------------- 1 file changed, 31 insertions(+), 23 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index f483334f..8cbbfa0f 100644 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -394,7 +394,7 @@ def get_inst_on_edge(arr, tile_pp_info): #### #### -def run_model( +def run_model_forward( forward_output_queue, tile_info_list, wsi_path, wsi_ext, wsi_proc_mag, wsi_cache_path, @@ -661,7 +661,9 @@ def simple_get_mask(): mask = morphology.binary_dilation(mask, morphology.disk(16)) return mask - wsi_mask = np.array(simple_get_mask() > 0, dtype=np.uint8) + # ! using entire wsi, debugging purpose ! + # wsi_mask = np.array(simple_get_mask() > 0, dtype=np.uint8) + wsi_mask = simple_get_mask() return wsi_mask def process_single_file(self, wsi_path, mask_path, output_dir): @@ -678,6 +680,7 @@ def process_single_file(self, wsi_path, mask_path, output_dir): wsi_name = path_obj.stem # TODO: expose read mpp mode + start = time.perf_counter() self.wsi_handler = get_file_handler(wsi_path, backend=wsi_ext) self.wsi_proc_shape = self.wsi_handler.get_dimensions(self.proc_mag) # ! cache here is for legacy and to deal with esoteric internal wsi format @@ -715,6 +718,8 @@ def process_single_file(self, wsi_path, mask_path, output_dir): self.wsi_proc_shape, tile_input_shape, tile_output_shape, self.ambiguous_size, self.patch_output_shape ) + end = time.perf_counter() + log_info("Preparing Input Output Placement: {0}".format(end - start)) # * Async Inference # * launch a seperate process to do forward and store the result in a queue. @@ -742,7 +747,7 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, mp_manager = torch_mp.Manager() # contain at most 5 tile ouput before polling forward func - mp_forward_output_queue = mp_manager.Queue(maxsize=5) + mp_forward_output_queue = mp_manager.Queue(maxsize=16) forward_info_list = collections.deque() nr_tile = tile_io_info_list.shape[0] @@ -758,7 +763,7 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, ) wsi_cache_path = "%s/src_wsi.npy" % self.cache_path - forward_process = mp.Process(target=run_model, + forward_process = mp.Process(target=run_model_forward, args=(mp_forward_output_queue, forward_info_list, wsi_path, wsi_ext, self.proc_mag, wsi_cache_path, self.run_step, self.net, loader_kwargs)) @@ -770,8 +775,8 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, } proc_pool = None - # if self.nr_post_proc_workers > 0: - # proc_pool = ProcessPoolExecutor(self.nr_post_proc_workers) + if self.nr_post_proc_workers > 0: + proc_pool = ProcessPoolExecutor(self.nr_post_proc_workers) offset_id = 0 # offset id to increment overall saving dict wsi_inst_info = {} @@ -810,6 +815,7 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, # * aggregate # ! the return id should be contiguous to maximuize # ! counting range in int32 + inst_wsi_id = offset_id # barrier in case no output in tile! for inst_id, inst_info in new_inst_dict.items(): inst_wsi_id = offset_id + inst_id + 1 wsi_inst_info[inst_wsi_id] = inst_info @@ -837,6 +843,7 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, # ! the return id should be contiguous to maximuize # ! counting range in int32 new_inst_dict, remove_uid_list = proc_future.result() + inst_wsi_id = offset_id # barrier in case no output in tile! for inst_id, inst_info in new_inst_dict.items(): inst_wsi_id = offset_id + inst_id + 1 wsi_inst_info[inst_wsi_id] = inst_info @@ -854,6 +861,7 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, # ## ** process full grid and vert/horiz fixing at the same time + start = time.perf_counter() info_list = list(zip(*all_tile_info[:3])) tile_io_info_list = np.concatenate(info_list[0], axis=0).astype(np.int32) tile_pp_info_list = np.concatenate(info_list[1], axis=0) @@ -861,10 +869,14 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, wsi_inst_info = run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list) + end = time.perf_counter() + log_info("Proc Grid: {0}".format(end - start)) + # import joblib # wsi_inst_info = joblib.load('cache_output.dat') ## ** re-infer and redo postproc for xsect alone + start = time.perf_counter() tile_io_info_list = all_tile_info[-1][0] tile_pp_info_list = all_tile_info[-1][1] tile_mode_list = all_tile_info[-1][2] @@ -872,17 +884,14 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, tile_pp_info_list, tile_mode_list, wsi_inst_info) + end = time.perf_counter() + log_info("Proc XSect: {0}".format(end - start)) if self.save_mask or self.save_thumb: json_path = "%s/json/%s.json" % (output_dir, wsi_name) else: json_path = "%s/%s.json" % (output_dir, wsi_name) - # self.__save_json(json_path, self.wsi_inst_info) - wsi_thumb_rgb = self.wsi_handler.get_full_img(read_mag=self.proc_mag) - from misc.viz_utils import visualize_instances_dict - wsi_overlay = visualize_instances_dict(wsi_thumb_rgb, wsi_inst_info, draw_dot=True) - cv2.imwrite('dump.png', cv2.cvtColor(wsi_overlay, cv2.COLOR_RGB2BGR)) - + self.__save_json(json_path, wsi_inst_info) return def process_wsi_list(self, run_args): @@ -911,23 +920,22 @@ def process_wsi_list(self, run_args): for wsi_path in wsi_path_list: wsi_base_name = pathlib.Path(wsi_path).stem - if wsi_base_name != 'mini_2': continue - msk_path = "%s/%s.png" % (self.input_mask_dir, wsi_base_name) if self.save_thumb or self.save_mask: output_file = "%s/json/%s.json" % (self.output_dir, wsi_base_name) else: output_file = "%s/%s.json" % (self.output_dir, wsi_base_name) - # if os.path.exists(output_file): - # log_info("Skip: %s" % wsi_base_name) - # continue - # try: - # log_info("Process: %s" % wsi_base_name) - # self.process_single_file(wsi_path, msk_path, self.output_dir) - # log_info("Finish") - # except: - # logging.exception("Crash") + if os.path.exists(output_file): + log_info("Skip: %s" % wsi_base_name) + continue + try: + log_info("Process: %s" % wsi_base_name) + self.process_single_file(wsi_path, msk_path, self.output_dir) + log_info("Finish") + except: + logging.exception("Crash") + self.process_single_file(wsi_path, msk_path, self.output_dir) break rm_n_mkdir(self.cache_path) # clean up all cache From ce973487eeeee1f7fd0784dc539bd8924fd3f33e Mon Sep 17 00:00:00 2001 From: vqdang Date: Tue, 9 Mar 2021 18:01:40 +0000 Subject: [PATCH 13/19] UPD: update write permission --- .gitignore | 0 .vscode/launch.json | 15 +++++++++++++++ CHANGELOG.md | 0 LICENSE | 0 README.md | 0 __init__.py | 0 compute_stats.py | 0 config.py | 0 convert_chkpt_tf2pytorch.py | 0 convert_format.py | 2 +- dataloader/__init__.py | 0 dataloader/augs.py | 0 dataloader/infer_loader.py | 0 dataloader/train_loader.py | 0 dataset.py | 0 diagram.png | Bin docs/diagram.png | Bin docs/seg.gif | Bin environment.yml | 0 .../.ipynb_checkpoints/usage-checkpoint.ipynb | 0 examples/usage.ipynb | 0 extract_patches.py | 0 infer/__init__.py | 0 infer/base.py | 6 +++--- infer/super_wsi.py | 11 ++++++++++- infer/tile.py | 0 infer/wsi.py | 1 + metrics/README.md | 0 metrics/__init__.py | 0 metrics/stats_utils.py | 0 misc/__init__.py | 0 misc/patch_extractor.py | 0 misc/utils.py | 0 misc/viz_utils.py | 0 misc/wsi_handler.py | 0 models/__init__.py | 0 models/hovernet/__init__.py | 0 models/hovernet/net_desc.py | 0 models/hovernet/net_utils.py | 0 models/hovernet/opt.py | 0 models/hovernet/post_proc.py | 0 models/hovernet/run_desc.py | 0 models/hovernet/targets.py | 0 models/hovernet/utils.py | 0 requirements.txt | 0 run_infer.py | 3 ++- run_tile.sh | 0 run_train.py | 0 run_wsi.sh | 16 +++++++++------- seg.gif | Bin type_info.json | 0 variables_tf2pytorch.csv | 0 52 files changed, 41 insertions(+), 13 deletions(-) mode change 100644 => 100755 .gitignore mode change 100644 => 100755 .vscode/launch.json mode change 100644 => 100755 CHANGELOG.md mode change 100644 => 100755 LICENSE mode change 100644 => 100755 README.md mode change 100644 => 100755 __init__.py mode change 100644 => 100755 compute_stats.py mode change 100644 => 100755 config.py mode change 100644 => 100755 convert_chkpt_tf2pytorch.py mode change 100644 => 100755 convert_format.py mode change 100644 => 100755 dataloader/__init__.py mode change 100644 => 100755 dataloader/augs.py mode change 100644 => 100755 dataloader/infer_loader.py mode change 100644 => 100755 dataloader/train_loader.py mode change 100644 => 100755 dataset.py mode change 100644 => 100755 diagram.png mode change 100644 => 100755 docs/diagram.png mode change 100644 => 100755 docs/seg.gif mode change 100644 => 100755 environment.yml mode change 100644 => 100755 examples/.ipynb_checkpoints/usage-checkpoint.ipynb mode change 100644 => 100755 examples/usage.ipynb mode change 100644 => 100755 extract_patches.py mode change 100644 => 100755 infer/__init__.py mode change 100644 => 100755 infer/base.py mode change 100644 => 100755 infer/super_wsi.py mode change 100644 => 100755 infer/tile.py mode change 100644 => 100755 infer/wsi.py mode change 100644 => 100755 metrics/README.md mode change 100644 => 100755 metrics/__init__.py mode change 100644 => 100755 metrics/stats_utils.py mode change 100644 => 100755 misc/__init__.py mode change 100644 => 100755 misc/patch_extractor.py mode change 100644 => 100755 misc/utils.py mode change 100644 => 100755 misc/viz_utils.py mode change 100644 => 100755 misc/wsi_handler.py mode change 100644 => 100755 models/__init__.py mode change 100644 => 100755 models/hovernet/__init__.py mode change 100644 => 100755 models/hovernet/net_desc.py mode change 100644 => 100755 models/hovernet/net_utils.py mode change 100644 => 100755 models/hovernet/opt.py mode change 100644 => 100755 models/hovernet/post_proc.py mode change 100644 => 100755 models/hovernet/run_desc.py mode change 100644 => 100755 models/hovernet/targets.py mode change 100644 => 100755 models/hovernet/utils.py mode change 100644 => 100755 requirements.txt mode change 100644 => 100755 run_infer.py mode change 100644 => 100755 run_tile.sh mode change 100644 => 100755 run_train.py mode change 100644 => 100755 run_wsi.sh mode change 100644 => 100755 seg.gif mode change 100644 => 100755 type_info.json mode change 100644 => 100755 variables_tf2pytorch.csv diff --git a/.gitignore b/.gitignore old mode 100644 new mode 100755 diff --git a/.vscode/launch.json b/.vscode/launch.json old mode 100644 new mode 100755 index 0ee4859d..2b9f3e87 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -12,6 +12,21 @@ "program": "${file}", "console": "integratedTerminal", "cwd" : "${fileDirname}", // run at the location of the runfile + // hook terminal cmd to debug file + "args": [ + // "--gpu=0", + "--nr_types=6", + "--type_info_path=type_info.json", + "--batch_size=64", + "--model_mode=fast", + "--model_path=../../pretrained/hovernet_fast_pannuke_pytorch.tar", + "--nr_inference_workers=4", + "--nr_post_proc_workers=8", + "wsi", + "--ambiguous_size=328", + "--input_dir=exp_output/wsi/samples/", + "--output_dir=exp_output/wsi/pred/", + ] }, { "name": "Python: Remote Attach", diff --git a/CHANGELOG.md b/CHANGELOG.md old mode 100644 new mode 100755 diff --git a/LICENSE b/LICENSE old mode 100644 new mode 100755 diff --git a/README.md b/README.md old mode 100644 new mode 100755 diff --git a/__init__.py b/__init__.py old mode 100644 new mode 100755 diff --git a/compute_stats.py b/compute_stats.py old mode 100644 new mode 100755 diff --git a/config.py b/config.py old mode 100644 new mode 100755 diff --git a/convert_chkpt_tf2pytorch.py b/convert_chkpt_tf2pytorch.py old mode 100644 new mode 100755 diff --git a/convert_format.py b/convert_format.py old mode 100644 new mode 100755 index 88839608..7e32676c --- a/convert_format.py +++ b/convert_format.py @@ -55,7 +55,7 @@ def rgb2int(rgb): target_format = "qupath" # to rescale the coordinate set to match with lv0 mag of the wsi scale_factor = 1.0 - root_dir = "dataset/dummy/out/" + root_dir = "exp_output/wsi/pred/" # to define the name, and color conversion code for each target format type_info_dict = { diff --git a/dataloader/__init__.py b/dataloader/__init__.py old mode 100644 new mode 100755 diff --git a/dataloader/augs.py b/dataloader/augs.py old mode 100644 new mode 100755 diff --git a/dataloader/infer_loader.py b/dataloader/infer_loader.py old mode 100644 new mode 100755 diff --git a/dataloader/train_loader.py b/dataloader/train_loader.py old mode 100644 new mode 100755 diff --git a/dataset.py b/dataset.py old mode 100644 new mode 100755 diff --git a/diagram.png b/diagram.png old mode 100644 new mode 100755 diff --git a/docs/diagram.png b/docs/diagram.png old mode 100644 new mode 100755 diff --git a/docs/seg.gif b/docs/seg.gif old mode 100644 new mode 100755 diff --git a/environment.yml b/environment.yml old mode 100644 new mode 100755 diff --git a/examples/.ipynb_checkpoints/usage-checkpoint.ipynb b/examples/.ipynb_checkpoints/usage-checkpoint.ipynb old mode 100644 new mode 100755 diff --git a/examples/usage.ipynb b/examples/usage.ipynb old mode 100644 new mode 100755 diff --git a/extract_patches.py b/extract_patches.py old mode 100644 new mode 100755 diff --git a/infer/__init__.py b/infer/__init__.py old mode 100644 new mode 100755 diff --git a/infer/base.py b/infer/base.py old mode 100644 new mode 100755 index 81165ec2..4f79235f --- a/infer/base.py +++ b/infer/base.py @@ -67,11 +67,11 @@ def __load_model(self): net.load_state_dict(saved_state_dict, strict=True) net = torch.nn.DataParallel(net) - net = net.to("cuda") + self.net = net.to("cuda") module_lib = import_module("models.hovernet.run_desc") - run_step = getattr(module_lib, "infer_step") - self.run_step = lambda input_batch: run_step(input_batch, net) + self.run_step = getattr(module_lib, "infer_step") + # self.run_step = lambda input_batch: run_step(input_batch, net) module_lib = import_module("models.hovernet.post_proc") self.post_proc_func = getattr(module_lib, "process") diff --git a/infer/super_wsi.py b/infer/super_wsi.py old mode 100644 new mode 100755 index 8cbbfa0f..26cf5dcd --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -793,8 +793,9 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, if forward_process.exitcode is not None \ and mp_forward_output_queue.empty(): break - + if not mp_forward_output_queue.empty(): + print(mp_forward_output_queue.qsize()) tile_idx, forward_output = mp_forward_output_queue.get() foward_pbar.update() @@ -926,6 +927,14 @@ def process_wsi_list(self, run_args): else: output_file = "%s/%s.json" % (self.output_dir, wsi_base_name) + # ! cache in case of loading across network + start = time.perf_counter() + wsi_holder_path = self.cache_src_wsi_path + shutil.copyfile(wsi_path, wsi_holder_path) + end = time.perf_counter() + log_info('Move WSI: {0}'.format(end - start)) + + if os.path.exists(output_file): log_info("Skip: %s" % wsi_base_name) continue diff --git a/infer/tile.py b/infer/tile.py old mode 100644 new mode 100755 diff --git a/infer/wsi.py b/infer/wsi.py old mode 100644 new mode 100755 index a8409cdc..353ae7cc --- a/infer/wsi.py +++ b/infer/wsi.py @@ -2,6 +2,7 @@ from concurrent.futures import FIRST_EXCEPTION, ProcessPoolExecutor, as_completed, wait from multiprocessing import Lock, Pool +# also to ensure similar behavior between linux and windows mp.set_start_method("spawn", True) # ! must be at top for VScode debugging import argparse diff --git a/metrics/README.md b/metrics/README.md old mode 100644 new mode 100755 diff --git a/metrics/__init__.py b/metrics/__init__.py old mode 100644 new mode 100755 diff --git a/metrics/stats_utils.py b/metrics/stats_utils.py old mode 100644 new mode 100755 diff --git a/misc/__init__.py b/misc/__init__.py old mode 100644 new mode 100755 diff --git a/misc/patch_extractor.py b/misc/patch_extractor.py old mode 100644 new mode 100755 diff --git a/misc/utils.py b/misc/utils.py old mode 100644 new mode 100755 diff --git a/misc/viz_utils.py b/misc/viz_utils.py old mode 100644 new mode 100755 diff --git a/misc/wsi_handler.py b/misc/wsi_handler.py old mode 100644 new mode 100755 diff --git a/models/__init__.py b/models/__init__.py old mode 100644 new mode 100755 diff --git a/models/hovernet/__init__.py b/models/hovernet/__init__.py old mode 100644 new mode 100755 diff --git a/models/hovernet/net_desc.py b/models/hovernet/net_desc.py old mode 100644 new mode 100755 diff --git a/models/hovernet/net_utils.py b/models/hovernet/net_utils.py old mode 100644 new mode 100755 diff --git a/models/hovernet/opt.py b/models/hovernet/opt.py old mode 100644 new mode 100755 diff --git a/models/hovernet/post_proc.py b/models/hovernet/post_proc.py old mode 100644 new mode 100755 diff --git a/models/hovernet/run_desc.py b/models/hovernet/run_desc.py old mode 100644 new mode 100755 diff --git a/models/hovernet/targets.py b/models/hovernet/targets.py old mode 100644 new mode 100755 diff --git a/models/hovernet/utils.py b/models/hovernet/utils.py old mode 100644 new mode 100755 diff --git a/requirements.txt b/requirements.txt old mode 100644 new mode 100755 diff --git a/run_infer.py b/run_infer.py old mode 100644 new mode 100755 index 15477e23..61be79f8 --- a/run_infer.py +++ b/run_infer.py @@ -181,6 +181,7 @@ infer = InferManager(**method_args) infer.process_file_list(run_args) else: - from infer.wsi import InferManager + # from infer.wsi import InferManager + from infer.super_wsi import InferManager infer = InferManager(**method_args) infer.process_wsi_list(run_args) diff --git a/run_tile.sh b/run_tile.sh old mode 100644 new mode 100755 diff --git a/run_train.py b/run_train.py old mode 100644 new mode 100755 diff --git a/run_wsi.sh b/run_wsi.sh old mode 100644 new mode 100755 index 766421c6..863d0f28 --- a/run_wsi.sh +++ b/run_wsi.sh @@ -4,12 +4,14 @@ python run_infer.py \ --type_info_path=type_info.json \ --batch_size=64 \ --model_mode=fast \ ---model_path=../pretrained/hovernet_fast_pannuke_type_tf2pytorch.tar \ ---nr_inference_workers=8 \ ---nr_post_proc_workers=16 \ +--model_path=../../pretrained/hovernet_fast_pannuke_pytorch.tar \ +--nr_inference_workers=4 \ +--nr_post_proc_workers=32 \ wsi \ ---input_dir=dataset/sample_wsis/wsi/ \ ---output_dir=dataset/sample_wsis/out/ \ ---input_mask_dir=dataset/sample_wsis/msk/ \ +--input_dir=exp_output/wsi/samples/full/ \ +--output_dir=exp_output/wsi/pred/ \ +--input_mask_dir=exp_output/wsi/samples/mask/ \ --save_thumb \ ---save_mask +--save_mask \ +--ambiguous_size=328 \ +--tile_shape=2048 \ No newline at end of file diff --git a/seg.gif b/seg.gif old mode 100644 new mode 100755 diff --git a/type_info.json b/type_info.json old mode 100644 new mode 100755 diff --git a/variables_tf2pytorch.csv b/variables_tf2pytorch.csv old mode 100644 new mode 100755 From 7f6b8f4b42ab06a0fbdad715071bb9923d47cdea Mon Sep 17 00:00:00 2001 From: vqdang Date: Tue, 9 Mar 2021 18:46:10 +0000 Subject: [PATCH 14/19] UPD: add proto for serialize wsi, doesnt work --- infer/super_wsi.py | 95 ++++++++++++++++++++++++++++++---------------- 1 file changed, 63 insertions(+), 32 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index 26cf5dcd..57854f7a 100755 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -62,6 +62,39 @@ def __getitem__(self, idx): patch_data = self.preproc(patch_data) return patch_data, patch_info +#### +class SerializeWSI(data.Dataset): + """ + `mp_shared_space` must be from torch.multiprocessing, for example + + mp_manager = torch_mp.Manager() + mp_shared_space = mp_manager.Namespace() + mp_shared_space.image = torch.from_numpy(image) + """ + def __init__(self, wsi_path, wsi_ext, wsi_read_mag, + mp_shared_space, preproc=None): + super().__init__() + self.mp_shared_space = mp_shared_space + self.preproc = preproc + + self.wsi_handler = get_file_handler(wsi_path, backend=wsi_ext) + # ! cache here is for legacy and to deal with esoteric internal wsi format + self.wsi_handler.prepare_reading(read_mag=wsi_read_mag) + return + + def __len__(self): + return len(self.mp_shared_space.patch_info_list) + + def __getitem__(self, idx): + patch_info = self.mp_shared_space.patch_info_list[idx] + patch_info = patch_info.numpy() # else reverse op wont work + tl, br = patch_info[0] # retrieve input placement, [1] is output + patch_data = self.wsi_handler.read_region(tl[::-1], (br - tl)[::-1]) + if self.preproc is not None: + patch_dat = patch_data.copy() + patch_data = self.preproc(patch_data) + return patch_data, patch_info + #### def _remove_inst(inst_map, remove_id_list): """Remove instances with id in remove_id_list. @@ -401,15 +434,15 @@ def run_model_forward( run_step, model, loader_kwargs): wsi_handler = get_file_handler(wsi_path, backend=wsi_ext) - # ! cache here is for legacy and to deal with esoteric internal wsi format - wsi_handler.prepare_reading(read_mag=wsi_proc_mag, cache_path=wsi_cache_path) + wsi_handler.prepare_reading(read_mag=wsi_proc_mag) # using shared memory namespace so all the loader workers use same # underlying image data, also allow persistent worker and fast data switching mp_manager = torch_mp.Manager() mp_shared_space = mp_manager.Namespace() - ds = SerializeArray(mp_shared_space) + # ds = SerializeArray(mp_shared_space) + ds = SerializeWSI(wsi_path, wsi_ext, wsi_proc_mag, mp_shared_space) loader = data.DataLoader(ds, **loader_kwargs, drop_last=False, persistent_workers=True, @@ -425,14 +458,14 @@ def run_model_forward( # ! (tile output is within tile input system) # ! this will shift both patch input and output placement to tile input system # ! hence, output placement need to be shifted (corrected) later for post proc - patch_info_list -= np.reshape(tile_input_tl, [1, 1, 1, 2]) + # patch_info_list -= np.reshape(tile_input_tl, [1, 1, 1, 2]) - tile_img = wsi_handler.read_region(tile_input_tl[::-1], - (tile_input_br - tile_input_tl)[::-1]) + # tile_img = wsi_handler.read_region(tile_input_tl[::-1], + # (tile_input_br - tile_input_tl)[::-1]) # change the data in namespace to sync across persistent loader worker # also no need to do locking as these are assumed to be read only from worker - mp_shared_space.tile_img = torch.from_numpy(tile_img).share_memory_() + # mp_shared_space.tile_img = torch.from_numpy(tile_img).share_memory_() mp_shared_space.patch_info_list = torch.from_numpy(patch_info_list).share_memory_() accumulated_patch_output = [] @@ -468,8 +501,10 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, for idx in range(len(patch_pos_list)): # zero idx to remove singleton, squeeze may kill h/w/c patch_pos = patch_pos_list[idx][0].copy() - # ! assume patch pos alrd aligned to be within tile input system - patch_pos = patch_pos - offset # shift from wsi to tile output system + + # # ! assume patch pos alrd aligned to be within tile input system + # patch_pos = patch_pos - offset # shift from wsi to tile output system + pos_tl, pos_br = patch_pos[1] # retrieve ouput placement pred_map[ pos_tl[0] : pos_br[0], @@ -684,9 +719,7 @@ def process_single_file(self, wsi_path, mask_path, output_dir): self.wsi_handler = get_file_handler(wsi_path, backend=wsi_ext) self.wsi_proc_shape = self.wsi_handler.get_dimensions(self.proc_mag) # ! cache here is for legacy and to deal with esoteric internal wsi format - self.wsi_handler.prepare_reading( - read_mag=self.proc_mag, cache_path="%s/src_wsi.npy" % self.cache_path - ) + self.wsi_handler.prepare_reading(read_mag=self.proc_mag) self.wsi_proc_shape = np.array(self.wsi_proc_shape[::-1]) # to Y, X self.wsi_mask = self._get_wsi_mask(self.wsi_handler, mask_path) @@ -782,8 +815,9 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, wsi_inst_info = {} if prev_wsi_inst_dict is not None: wsi_inst_info = copy.deepcopy(prev_wsi_inst_dict) - offset_id = max(wsi_inst_info.keys()) + 1 - + if len(wsi_inst_info) > 0: + offset_id = max(wsi_inst_info.keys()) + 1 + foward_pbar = pbar_creator('Forward', nr_tile, pos=0) postpr_pbar = pbar_creator('PostPro', nr_tile, pos=1) @@ -795,7 +829,6 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, break if not mp_forward_output_queue.empty(): - print(mp_forward_output_queue.qsize()) tile_idx, forward_output = mp_forward_output_queue.get() foward_pbar.update() @@ -862,19 +895,20 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, # ## ** process full grid and vert/horiz fixing at the same time - start = time.perf_counter() - info_list = list(zip(*all_tile_info[:3])) - tile_io_info_list = np.concatenate(info_list[0], axis=0).astype(np.int32) - tile_pp_info_list = np.concatenate(info_list[1], axis=0) - tile_mode_list = np.concatenate(info_list[2], axis=0) - wsi_inst_info = run_once(tile_io_info_list, - tile_pp_info_list, - tile_mode_list) - end = time.perf_counter() + # start = time.perf_counter() + # info_list = list(zip(*all_tile_info[:3])) + # tile_io_info_list = np.concatenate(info_list[0], axis=0).astype(np.int32) + # tile_pp_info_list = np.concatenate(info_list[1], axis=0) + # tile_mode_list = np.concatenate(info_list[2], axis=0) + # wsi_inst_info = run_once(tile_io_info_list, + # tile_pp_info_list, + # tile_mode_list) + # end = time.perf_counter() log_info("Proc Grid: {0}".format(end - start)) # import joblib # wsi_inst_info = joblib.load('cache_output.dat') + wsi_inst_info = {} ## ** re-infer and redo postproc for xsect alone start = time.perf_counter() @@ -928,12 +962,11 @@ def process_wsi_list(self, run_args): output_file = "%s/%s.json" % (self.output_dir, wsi_base_name) # ! cache in case of loading across network - start = time.perf_counter() - wsi_holder_path = self.cache_src_wsi_path - shutil.copyfile(wsi_path, wsi_holder_path) - end = time.perf_counter() - log_info('Move WSI: {0}'.format(end - start)) - + # start = time.perf_counter() + # wsi_holder_path = self.cache_src_wsi_path + # shutil.copyfile(wsi_path, wsi_holder_path) + # end = time.perf_counter() + # log_info('Move WSI: {0}'.format(end - start)) if os.path.exists(output_file): log_info("Skip: %s" % wsi_base_name) @@ -945,7 +978,5 @@ def process_wsi_list(self, run_args): except: logging.exception("Crash") - self.process_single_file(wsi_path, msk_path, self.output_dir) - break rm_n_mkdir(self.cache_path) # clean up all cache return From 91592f957106f5f710b1f6e8140f5cb96c00e728 Mon Sep 17 00:00:00 2001 From: vqdang Date: Wed, 10 Mar 2021 11:36:19 +0000 Subject: [PATCH 15/19] UPD: sync handler implementation --- misc/wsi_handler.py | 230 ++++++++++++++++++++------------------------ 1 file changed, 106 insertions(+), 124 deletions(-) diff --git a/misc/wsi_handler.py b/misc/wsi_handler.py index 49aa4ea0..8376491d 100755 --- a/misc/wsi_handler.py +++ b/misc/wsi_handler.py @@ -1,194 +1,176 @@ -from collections import OrderedDict -import cv2 -import numpy as np -from skimage import img_as_ubyte -from skimage import color import re import subprocess +import warnings +from collections import OrderedDict +import cv2 +import glymur +import numpy as np import openslide +import tifffile +from skimage import color, img_as_ubyte class FileHandler(object): def __init__(self): - """The handler is responsible for storing the processed data, parsing + """ + The handler is responsible for storing the processed data, parsing the metadata from original file, and reading it from storage. """ self.metadata = { - ("available_mag", None), - ("base_mag", None), - ("vendor", None), - ("mpp ", None), - ("base_shape", None), + ('available_mag', None), + ('base_mag' , None), + ('vendor' , None), + ('base_mpp' , None), + ('base_shape' , None), } + pass def __load_metadata(self): - raise NotImplementedError - - def get_full_img(self, read_mag=None, read_mpp=None): - """Only use `read_mag` or `read_mpp`, not both, prioritize `read_mpp`. + pass - `read_mpp` is in X, Y format + def read_region(self): + pass + + def get_dimensions(self, read_mag=None, read_mpp=None): """ - raise NotImplementedError - - def read_region(self, coords, size): - """Must call `prepare_reading` before hand. - - Args: - coords (tuple): (dims_x, dims_y), - top left coordinates of image region at selected - `read_mag` or `read_mpp` from `prepare_reading` - size (tuple): (dims_x, dims_y) - width and height of image region at selected - `read_mag` or `read_mpp` from `prepare_reading` - + Will be in X, Y """ - raise NotImplementedError - - def get_dimensions(self, read_mag=None, read_mpp=None): - """Will be in X, Y.""" if read_mpp is not None: - read_scale = (self.metadata["base_mpp"] / read_mpp)[0] - read_mag = read_scale * self.metadata["base_mag"] - scale = read_mag / self.metadata["base_mag"] + read_scale = (self.metadata['base_mpp'] / read_mpp)[0] + read_mag = read_scale * self.metadata['base_mag'] + scale = read_mag / self.metadata['base_mag'] # may off some pixels wrt existing mag - return (self.metadata["base_shape"] * scale).astype(np.int32) + return (self.metadata['base_shape'] * scale).astype(np.int32) - def prepare_reading(self, read_mag=None, read_mpp=None, cache_path=None): - """Only use `read_mag` or `read_mpp`, not both, prioritize `read_mpp`. + def prepare_reading(self, read_mag=None, read_mpp=None): + """ + Only use either of these parameter, prioritize `read_mpp` - `read_mpp` is in X, Y format. + `read_mpp` is in X, Y format """ - read_lv, scale_factor = self._get_read_info( - read_mag=read_mag, read_mpp=read_mpp - ) - - if scale_factor is None: - self.image_ptr = None - self.read_lv = read_lv - else: - np.save(cache_path, self.get_full_img(read_mag=read_mag)) - self.image_ptr = np.load(cache_path, mmap_mode="r") + hires_lv, scale_from_hires_lv, scale_from_lv0 = self._get_read_info(read_mag=read_mag, read_mpp=read_mpp) + + self.read_lv = hires_lv + self.scale_from_hires_lv = scale_from_hires_lv + self.scale_from_lv0 = scale_from_lv0 return def _get_read_info(self, read_mag=None, read_mpp=None): if read_mpp is not None: - assert read_mpp[0] == read_mpp[1], "Not supported uneven `read_mpp`" - read_scale = (self.metadata["base_mpp"] / read_mpp)[0] - read_mag = read_scale * self.metadata["base_mag"] + assert read_mpp[0] == read_mpp[1], 'Not supported uneven `read_mpp`' + read_scale = (self.metadata['base_mpp'] / read_mpp)[0] + read_mag = read_scale * self.metadata['base_mag'] hires_mag = read_mag - scale_factor = None - if read_mag not in self.metadata["available_mag"]: - if read_mag > self.metadata["base_mag"]: - scale_factor = read_mag / self.metadata["base_mag"] - hires_mag = self.metadata["base_mag"] - else: - mag_list = np.array(self.metadata["available_mag"]) + scale_from_lv0 = read_mag / self.metadata['base_mag'] + scale_from_hires_lv = scale_from_lv0 + if read_mag not in self.metadata['available_mag']: + if read_mag > self.metadata['base_mag']: + scale_from_hires_lv = scale_from_lv0 + hires_mag = self.metadata['base_mag'] + else: + mag_list = np.array(self.metadata['available_mag']) mag_list = np.sort(mag_list)[::-1] hires_mag = mag_list - read_mag # only use higher mag as base for loading hires_mag = hires_mag[hires_mag > 0] # use the immediate higher to save compuration hires_mag = mag_list[np.argmin(hires_mag)] - scale_factor = read_mag / hires_mag - - hires_lv = self.metadata["available_mag"].index(hires_mag) - return hires_lv, scale_factor + scale_from_hires_lv = read_mag / hires_mag + hires_lv = self.metadata['available_mag'].index(hires_mag) + return hires_lv, scale_from_hires_lv, scale_from_lv0 class OpenSlideHandler(FileHandler): - """Class for handling OpenSlide supported whole-slide images.""" - + """ + Class for handling OpenSlide supported whole-slide images + """ def __init__(self, file_path): - """file_path (string): path to single whole-slide image.""" + """ + file_path (string): path to single whole-slide image + """ super().__init__() - self.file_ptr = openslide.OpenSlide(file_path) # load OpenSlide object + self.file_ptr = openslide.OpenSlide(file_path) # load OpenSlide object self.metadata = self.__load_metadata() # only used for cases where the read magnification is different from - self.image_ptr = None # the existing modes of the read file + self.image_ptr = None # the existing modes of the read file self.read_level = None def __load_metadata(self): metadata = {} wsi_properties = self.file_ptr.properties - level_0_magnification = wsi_properties[openslide.PROPERTY_NAME_OBJECTIVE_POWER] - level_0_magnification = float(level_0_magnification) + mpp = [float(wsi_properties[openslide.PROPERTY_NAME_MPP_X]), + float(wsi_properties[openslide.PROPERTY_NAME_MPP_Y])] + mpp = np.array(mpp) - downsample_level = self.file_ptr.level_downsamples + try : + level_0_magnification = wsi_properties[openslide.PROPERTY_NAME_OBJECTIVE_POWER] + level_0_magnification = float(level_0_magnification) + downsample_level = self.file_ptr.level_downsamples + magnification_level = [level_0_magnification / lv for lv in downsample_level] + except: + if mpp[0] > 0.1 and mpp[1] < 0.4: + level_0_magnification = 40.0 + if mpp[0] >= 0.4 and mpp[1] < 0.6: + level_0_magnification = 20.0 + downsample_level = self.file_ptr.level_downsamples + warnings.warn('Could not detect magnification, guess from `mpp`.') magnification_level = [level_0_magnification / lv for lv in downsample_level] - mpp = [ - wsi_properties[openslide.PROPERTY_NAME_MPP_X], - wsi_properties[openslide.PROPERTY_NAME_MPP_Y], - ] - mpp = np.array(mpp) - metadata = [ - ("available_mag", magnification_level), # highest to lowest mag - ("base_mag", magnification_level[0]), - ("vendor", wsi_properties[openslide.PROPERTY_NAME_VENDOR]), - ("mpp ", mpp), - ("base_shape", np.array(self.file_ptr.dimensions)), + ('available_mag', magnification_level), # highest to lowest mag + ('base_mag' , magnification_level[0]), + ('vendor' , wsi_properties[openslide.PROPERTY_NAME_VENDOR]), + ('base_mpp' , mpp), + ('base_shape' , np.array(self.file_ptr.dimensions)), ] return OrderedDict(metadata) - + def read_region(self, coords, size): - """Must call `prepare_reading` before hand. + """ + Must call `prepare_image` before hand Args: - coords (tuple): (dims_x, dims_y), - top left coordinates of image region at selected - `read_mag` or `read_mpp` from `prepare_reading` - size (tuple): (dims_x, dims_y) - width and height of image region at selected - `read_mag` or `read_mpp` from `prepare_reading` - + coords (tuple): top left coordinates of image region at requested mag/mpp in `prepare_reading` + read_level_size (tuple): dimensions of image region at requested mag/mpp in `prepare_reading` (dims_x, dims_y) """ - if self.image_ptr is None: - # convert coord from read lv to lv zero - lv_0_shape = np.array(self.file_ptr.level_dimensions[0]) - lv_r_shape = np.array(self.file_ptr.level_dimensions[self.read_lv]) - up_sample = (lv_0_shape / lv_r_shape)[0] - new_coord = [0, 0] - new_coord[0] = int(coords[0] * up_sample) - new_coord[1] = int(coords[1] * up_sample) - region = self.file_ptr.read_region(new_coord, self.read_lv, size) - else: - region = self.image_ptr[ - coords[1] : coords[1] + size[1], coords[0] : coords[0] + size[0] - ] - return np.array(region)[..., :3] - def get_full_img(self, read_mag=None, read_mpp=None): - """Only use `read_mag` or `read_mpp`, not both, prioritize `read_mpp`. + # convert coord from read lv to lv zero + coord_lv0 = [0, 0] + coord_lv0[0] = int(coords[0] / self.scale_from_lv0) + coord_lv0[1] = int(coords[1] / self.scale_from_lv0) - `read_mpp` is in X, Y format. - """ + size_at_read_lv = (size / self.scale_from_hires_lv).astype(np.int32) + region = self.file_ptr.read_region(coord_lv0, self.read_lv, size_at_read_lv) + + region = np.array(region)[...,:3] + if self.scale_from_hires_lv is not None: + interp = cv2.INTER_LINEAR + region = cv2.resize(region, tuple(size), interpolation=interp) + return region - read_lv, scale_factor = self._get_read_info( - read_mag=read_mag, read_mpp=read_mpp - ) + def get_full_img(self, read_mag=None, read_mpp=None): + + read_lv, scale_from_hires_lv, scale_from_lv0 = self._get_read_info( + read_mag=read_mag, + read_mpp=read_mpp) read_size = self.file_ptr.level_dimensions[read_lv] wsi_img = self.file_ptr.read_region((0, 0), read_lv, read_size) - wsi_img = np.array(wsi_img)[..., :3] # remove alpha channel - if scale_factor is not None: + wsi_img = np.array(wsi_img)[...,:3] # remove alpha channel + if scale_from_hires_lv is not None: # now rescale then return - if scale_factor > 1.0: - interp = cv2.INTER_CUBIC - else: - interp = cv2.INTER_LINEAR - wsi_img = cv2.resize( - wsi_img, (0, 0), fx=scale_factor, fy=scale_factor, interpolation=interp - ) - return wsi_img - + interp = cv2.INTER_LINEAR + wsi_img = cv2.resize(wsi_img, (0, 0), + fx=scale_from_hires_lv, + fy=scale_from_hires_lv, + interpolation=interp) + return wsi_img def get_file_handler(path, backend): if backend in [ From a37675866d09ba73bbf01ea61dddc29fee54a608 Mon Sep 17 00:00:00 2001 From: vqdang Date: Wed, 10 Mar 2021 11:41:05 +0000 Subject: [PATCH 16/19] UPD: sync protocol --- infer/base.py | 2 +- infer/tile.py | 2 +- infer/wsi.py | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/infer/base.py b/infer/base.py index 4f79235f..64e9993b 100755 --- a/infer/base.py +++ b/infer/base.py @@ -67,7 +67,7 @@ def __load_model(self): net.load_state_dict(saved_state_dict, strict=True) net = torch.nn.DataParallel(net) - self.net = net.to("cuda") + self.model = net.to("cuda") module_lib = import_module("models.hovernet.run_desc") self.run_step = getattr(module_lib, "infer_step") diff --git a/infer/tile.py b/infer/tile.py index f27c7d3d..1caaf360 100755 --- a/infer/tile.py +++ b/infer/tile.py @@ -299,7 +299,7 @@ def detach_items_of_uid(items_list, uid, nr_expected_items): accumulated_patch_output = [] for batch_idx, batch_data in enumerate(dataloader): sample_data_list, sample_info_list = batch_data - sample_output_list = self.run_step(sample_data_list) + sample_output_list = self.run_step(self.model, sample_data_list) sample_info_list = sample_info_list.numpy() curr_batch_size = sample_output_list.shape[0] sample_output_list = np.split( diff --git a/infer/wsi.py b/infer/wsi.py index 353ae7cc..a893935b 100755 --- a/infer/wsi.py +++ b/infer/wsi.py @@ -287,7 +287,7 @@ def __run_model(self, patch_top_left_list, pbar_desc): accumulated_patch_output = [] for batch_idx, batch_data in enumerate(dataloader): sample_data_list, sample_info_list = batch_data - sample_output_list = self.run_step(sample_data_list) + sample_output_list = self.run_step(self.model, sample_data_list) sample_info_list = sample_info_list.numpy() curr_batch_size = sample_output_list.shape[0] sample_output_list = np.split(sample_output_list, curr_batch_size, axis=0) From 9258369f30cc36e15814580ea8fd277291c6f08a Mon Sep 17 00:00:00 2001 From: vqdang Date: Wed, 10 Mar 2021 15:02:02 +0000 Subject: [PATCH 17/19] UPD: persistent workers for old wsi infer mode --- infer/wsi.py | 129 ++++++++++++++++++++++++++++++--------------------- 1 file changed, 77 insertions(+), 52 deletions(-) diff --git a/infer/wsi.py b/infer/wsi.py index a893935b..4e8f9db1 100755 --- a/infer/wsi.py +++ b/infer/wsi.py @@ -25,9 +25,9 @@ import scipy.io as sio import torch import torch.utils.data as data +import torch.multiprocessing as torch_mp import tqdm -from dataloader.infer_loader import SerializeArray, SerializeFileList -from docopt import docopt + from misc.utils import ( cropping_center, get_bounding_box, @@ -42,11 +42,6 @@ thread_lock = Lock() -#### -def _init_worker_child(lock_): - global lock - lock = lock_ - #### def _remove_inst(inst_map, remove_id_list): @@ -256,48 +251,38 @@ def _assemble_and_flush(wsi_pred_map_mmap_path, chunk_info, patch_output_list): # print(chunk_info.flatten(), 'pass') return - #### -class InferManager(base.InferManager): - def __run_model(self, patch_top_left_list, pbar_desc): - # TODO: the cost of creating dataloader may not be cheap ? - dataset = SerializeArray( - "%s/cache_chunk.npy" % self.cache_path, - patch_top_left_list, - self.patch_input_shape, - ) +class SerializeArray(data.Dataset): + """ + `mp_shared_space` must be from torch.multiprocessing, for example - dataloader = data.DataLoader( - dataset, - num_workers=self.nr_inference_workers, - batch_size=self.batch_size, - drop_last=False, - ) + mp_manager = torch_mp.Manager() + mp_shared_space = mp_manager.Namespace() + mp_shared_space.image = torch.from_numpy(image) + """ + def __init__(self, mp_shared_space, patch_size, preproc=None): + super().__init__() + self.patch_size = patch_size + self.mp_shared_space = mp_shared_space + self.preproc = preproc + return - pbar = tqdm.tqdm( - desc=pbar_desc, - leave=True, - total=int(len(dataloader)), - ncols=80, - ascii=True, - position=0, - ) + def __len__(self): + return len(self.mp_shared_space.patch_info_list) - # run inference on input patches - accumulated_patch_output = [] - for batch_idx, batch_data in enumerate(dataloader): - sample_data_list, sample_info_list = batch_data - sample_output_list = self.run_step(self.model, sample_data_list) - sample_info_list = sample_info_list.numpy() - curr_batch_size = sample_output_list.shape[0] - sample_output_list = np.split(sample_output_list, curr_batch_size, axis=0) - sample_info_list = np.split(sample_info_list, curr_batch_size, axis=0) - sample_output_list = list(zip(sample_info_list, sample_output_list)) - accumulated_patch_output.extend(sample_output_list) - pbar.update() - pbar.close() - return accumulated_patch_output + def __getitem__(self, idx): + patch_info = self.mp_shared_space.patch_info_list[idx] + tl, br = patch_info[0] # retrieve input placement, [1] is output + patch_data = self.mp_shared_space.tile_img[ + tl[0] : tl[0] + self.patch_size[0], + tl[1] : tl[1] + self.patch_size[1]] + if self.preproc is not None: + patch_dat = patch_data.copy() + patch_data = self.preproc(patch_data) + return patch_data, patch_info[0][0] +#### +class InferManager(base.InferManager): def __select_valid_patches(self, patch_info_list, has_output_info=True): """Select valid patches from the list of input patch information. @@ -335,12 +320,53 @@ def __get_raw_prediction(self, chunk_info_list, patch_info_list): patch_info_list: list of patch coordinate information """ + pbar_creator = lambda x, y, pos=0, leave=False: tqdm.tqdm( + desc=x, leave=leave, total=y, ncols=80, ascii=True, position=pos + ) + + mp_manager = torch_mp.Manager() + mp_shared_space = mp_manager.Namespace() + + dataset = SerializeArray(mp_shared_space, self.patch_input_shape) + loader = data.DataLoader( + dataset, + num_workers=self.nr_inference_workers, + batch_size=self.batch_size, + drop_last=False, + persistent_workers=self.nr_inference_workers > 0 + ) + def run_model_forward(mp_manager, chunk_image, patch_info_list): + + mp_shared_space.tile_img = torch.from_numpy(chunk_image).share_memory_() + mp_shared_space.patch_info_list = torch.from_numpy(patch_info_list).share_memory_() + + nr_batch = int(patch_info_list.shape[0] / self.batch_size) + batch_pbar = pbar_creator('Proc Batch', nr_batch, pos=1) + + # run inference on input patches + accumulated_patch_output = [] + for batch_idx, batch_data in enumerate(loader): + sample_data_list, sample_info_list = batch_data + sample_output_list = self.run_step(sample_data_list, self.model) + sample_info_list = sample_info_list.numpy() + curr_batch_size = sample_output_list.shape[0] + sample_output_list = np.split(sample_output_list, curr_batch_size, axis=0) + sample_info_list = np.split(sample_info_list, curr_batch_size, axis=0) + sample_output_list = list(zip(sample_info_list, sample_output_list)) + accumulated_patch_output.extend(sample_output_list) + batch_pbar.update() + batch_pbar.close() + return accumulated_patch_output + # 1 dedicated thread just to write results back to disk proc_pool = Pool(processes=1) wsi_pred_map_mmap_path = "%s/pred_map.npy" % self.cache_path - + + nr_chunk = chunk_info_list.shape[0] + chunk_pbar = pbar_creator('Proc Chunk', nr_chunk, pos=0) + masking = lambda x, a, b: (a <= x) & (x <= b) - for idx in range(0, chunk_info_list.shape[0]): + for idx in range(0, nr_chunk): chunk_info = chunk_info_list[idx] # select patch basing on top left coordinate of input start_coord = chunk_info[0, 0] @@ -370,15 +396,16 @@ def __get_raw_prediction(self, chunk_info_list, patch_info_list): chunk_data = np.array(chunk_data)[..., :3] np.save("%s/cache_chunk.npy" % self.cache_path, chunk_data) - pbar_desc = "Process Chunk %d/%d" % (idx, chunk_info_list.shape[0]) - patch_output_list = self.__run_model( - chunk_patch_info_list[:, 0, 0], pbar_desc + patch_output_list = run_model_forward( + mp_manager, chunk_data, chunk_patch_info_list ) proc_pool.apply_async( _assemble_and_flush, args=(wsi_pred_map_mmap_path, chunk_info, patch_output_list), ) + chunk_pbar.update() + chunk_pbar.close() proc_pool.close() proc_pool.join() return @@ -470,9 +497,7 @@ def process_single_file(self, wsi_path, msk_path, output_dir): start = time.perf_counter() self.wsi_handler = get_file_handler(wsi_path, backend=wsi_ext) self.wsi_proc_shape = self.wsi_handler.get_dimensions(self.proc_mag) - self.wsi_handler.prepare_reading( - read_mag=self.proc_mag, cache_path="%s/src_wsi.npy" % self.cache_path - ) + self.wsi_handler.prepare_reading(read_mag=self.proc_mag) self.wsi_proc_shape = np.array(self.wsi_proc_shape[::-1]) # to Y, X if msk_path is not None and os.path.isfile(msk_path): From 9e54d8868c4119731db635ab261c420bb30b6550 Mon Sep 17 00:00:00 2001 From: vqdang Date: Wed, 10 Mar 2021 15:03:48 +0000 Subject: [PATCH 18/19] UPD: remove debug code --- infer/super_wsi.py | 95 +++++++++++++++++----------------------------- 1 file changed, 34 insertions(+), 61 deletions(-) diff --git a/infer/super_wsi.py b/infer/super_wsi.py index 57854f7a..913abf0f 100755 --- a/infer/super_wsi.py +++ b/infer/super_wsi.py @@ -62,39 +62,6 @@ def __getitem__(self, idx): patch_data = self.preproc(patch_data) return patch_data, patch_info -#### -class SerializeWSI(data.Dataset): - """ - `mp_shared_space` must be from torch.multiprocessing, for example - - mp_manager = torch_mp.Manager() - mp_shared_space = mp_manager.Namespace() - mp_shared_space.image = torch.from_numpy(image) - """ - def __init__(self, wsi_path, wsi_ext, wsi_read_mag, - mp_shared_space, preproc=None): - super().__init__() - self.mp_shared_space = mp_shared_space - self.preproc = preproc - - self.wsi_handler = get_file_handler(wsi_path, backend=wsi_ext) - # ! cache here is for legacy and to deal with esoteric internal wsi format - self.wsi_handler.prepare_reading(read_mag=wsi_read_mag) - return - - def __len__(self): - return len(self.mp_shared_space.patch_info_list) - - def __getitem__(self, idx): - patch_info = self.mp_shared_space.patch_info_list[idx] - patch_info = patch_info.numpy() # else reverse op wont work - tl, br = patch_info[0] # retrieve input placement, [1] is output - patch_data = self.wsi_handler.read_region(tl[::-1], (br - tl)[::-1]) - if self.preproc is not None: - patch_dat = patch_data.copy() - patch_data = self.preproc(patch_data) - return patch_data, patch_info - #### def _remove_inst(inst_map, remove_id_list): """Remove instances with id in remove_id_list. @@ -441,13 +408,13 @@ def run_model_forward( mp_manager = torch_mp.Manager() mp_shared_space = mp_manager.Namespace() - # ds = SerializeArray(mp_shared_space) - ds = SerializeWSI(wsi_path, wsi_ext, wsi_proc_mag, mp_shared_space) + ds = SerializeArray(mp_shared_space) loader = data.DataLoader(ds, **loader_kwargs, drop_last=False, persistent_workers=True, ) + all_time = 0 for tile_idx, tile_info in enumerate(tile_info_list): tile_info, patch_info_list = tile_info @@ -458,15 +425,18 @@ def run_model_forward( # ! (tile output is within tile input system) # ! this will shift both patch input and output placement to tile input system # ! hence, output placement need to be shifted (corrected) later for post proc - # patch_info_list -= np.reshape(tile_input_tl, [1, 1, 1, 2]) + patch_info_list -= np.reshape(tile_input_tl, [1, 1, 1, 2]) - # tile_img = wsi_handler.read_region(tile_input_tl[::-1], - # (tile_input_br - tile_input_tl)[::-1]) + tile_img = wsi_handler.read_region(tile_input_tl[::-1], + (tile_input_br - tile_input_tl)[::-1]) # change the data in namespace to sync across persistent loader worker # also no need to do locking as these are assumed to be read only from worker - # mp_shared_space.tile_img = torch.from_numpy(tile_img).share_memory_() + start = time.perf_counter() + mp_shared_space.tile_img = torch.from_numpy(tile_img).share_memory_() mp_shared_space.patch_info_list = torch.from_numpy(patch_info_list).share_memory_() + end = time.perf_counter() + all_time += (end - start) accumulated_patch_output = [] for batch_idx, batch_data in enumerate(loader): @@ -475,6 +445,7 @@ def run_model_forward( sample_info_list = sample_info_list.numpy() accumulated_patch_output.append([sample_info_list, sample_output_list]) forward_output_queue.put([tile_idx, accumulated_patch_output]) + log_info('Load Time: {0}'.format(all_time)) return #### @@ -503,7 +474,7 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, patch_pos = patch_pos_list[idx][0].copy() # # ! assume patch pos alrd aligned to be within tile input system - # patch_pos = patch_pos - offset # shift from wsi to tile output system + patch_pos = patch_pos - offset # shift from wsi to tile output system pos_tl, pos_br = patch_pos[1] # retrieve ouput placement pred_map[ @@ -535,18 +506,21 @@ def postproc_tile(tile_io_info, tile_pp_info, tile_mode, # |///////////////////////////| | # ----------------------------| V + # ! extreme slow down may happen because the aggregated results are huge def draw_prev_pred_inst(): + tile_output = tile_io_info[1] + tile_canvas = np.zeros(tile_output[1] - tile_output[0], dtype=np.int32) + if len(prev_wsi_inst_dict) == 0: return tile_canvas + wsi_inst_uid_list = np.array(list(prev_wsi_inst_dict.keys())) wsi_inst_com_list = np.array([v['centroid'] for v in prev_wsi_inst_dict.values()]) wsi_inst_com_list = wsi_inst_com_list[:,::-1] # XY to YX - tile_output = tile_io_info[1] sel = (wsi_inst_com_list[:,0] > tile_output[0,0]) sel &= (wsi_inst_com_list[:,0] < tile_output[1,0]) sel &= (wsi_inst_com_list[:,1] > tile_output[0,1]) sel &= (wsi_inst_com_list[:,1] < tile_output[1,1]) sel_idx = np.nonzero(sel.flatten())[0] - tile_canvas = np.zeros(tile_output[1] - tile_output[0], dtype=np.int32) for inst_idx in sel_idx: inst_uid = wsi_inst_uid_list[inst_idx] # shift from wsi system to tile output system @@ -601,7 +575,7 @@ def draw_prev_pred_inst(): inst_on_margin = np.union1d(inst_on_margin1, inst_on_margin2) inst_within_margin = np.setdiff1d(inst_in_margin, inst_on_margin, assume_unique=True) remove_inst_set = inst_within_margin.tolist() - # ! but we also need to remove prev inst exising on the global space of entire wsi + # ! but we also need to remove prev inst exist in the global space of entire wsi # ! on the margin prev_pred_inst = draw_prev_pred_inst() inst_on_margin1 = get_inst_on_margin(prev_pred_inst, margin_size-1, [1, 1, 1, 1]) @@ -799,7 +773,7 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, forward_process = mp.Process(target=run_model_forward, args=(mp_forward_output_queue, forward_info_list, wsi_path, wsi_ext, self.proc_mag, wsi_cache_path, - self.run_step, self.net, loader_kwargs)) + self.run_step, self.model, loader_kwargs)) forward_process.start() post_proc_kwargs = { @@ -858,7 +832,8 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, if remove_uid_list is not None: for inst_uid in remove_uid_list: if inst_uid in wsi_inst_info: - wsi_inst_info.pop(inst_uid) + # faster than pop + del wsi_inst_info[inst_uid] postpr_pbar.update() if forward_process.exitcode > 0: @@ -874,7 +849,7 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, print(proc_future.exception()) # * aggregate - # ! the return id should be contiguous to maximuize + # ! the return id should be contiguous to maximize # ! counting range in int32 new_inst_dict, remove_uid_list = proc_future.result() inst_wsi_id = offset_id # barrier in case no output in tile! @@ -886,30 +861,28 @@ def run_once(tile_io_info_list, tile_pp_info_list, tile_mode_list, if remove_uid_list is not None: for inst_uid in remove_uid_list: if inst_uid in wsi_inst_info: - wsi_inst_info.pop(inst_uid) + # faster than pop + del wsi_inst_info[inst_uid] postpr_pbar.update() foward_pbar.close() postpr_pbar.close() - + proc_pool.shutdown() return wsi_inst_info # ## ** process full grid and vert/horiz fixing at the same time - # start = time.perf_counter() - # info_list = list(zip(*all_tile_info[:3])) - # tile_io_info_list = np.concatenate(info_list[0], axis=0).astype(np.int32) - # tile_pp_info_list = np.concatenate(info_list[1], axis=0) - # tile_mode_list = np.concatenate(info_list[2], axis=0) - # wsi_inst_info = run_once(tile_io_info_list, - # tile_pp_info_list, - # tile_mode_list) - # end = time.perf_counter() + start = time.perf_counter() + info_list = list(zip(*all_tile_info[:3])) + tile_io_info_list = np.concatenate(info_list[0], axis=0).astype(np.int32) + tile_pp_info_list = np.concatenate(info_list[1], axis=0) + tile_mode_list = np.concatenate(info_list[2], axis=0) + wsi_inst_info = run_once(tile_io_info_list, + tile_pp_info_list, + tile_mode_list) + end = time.perf_counter() log_info("Proc Grid: {0}".format(end - start)) - - # import joblib - # wsi_inst_info = joblib.load('cache_output.dat') - wsi_inst_info = {} + # wsi_inst_info = {} ## ** re-infer and redo postproc for xsect alone start = time.perf_counter() tile_io_info_list = all_tile_info[-1][0] From 09de5da4704e283fe0ff5f94a77696330d761467 Mon Sep 17 00:00:00 2001 From: vqdang Date: Wed, 10 Mar 2021 15:04:15 +0000 Subject: [PATCH 19/19] UPD: update protocol --- infer/tile.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/infer/tile.py b/infer/tile.py index 1caaf360..242e84f7 100755 --- a/infer/tile.py +++ b/infer/tile.py @@ -299,7 +299,7 @@ def detach_items_of_uid(items_list, uid, nr_expected_items): accumulated_patch_output = [] for batch_idx, batch_data in enumerate(dataloader): sample_data_list, sample_info_list = batch_data - sample_output_list = self.run_step(self.model, sample_data_list) + sample_output_list = self.run_step(sample_data_list, self.model) sample_info_list = sample_info_list.numpy() curr_batch_size = sample_output_list.shape[0] sample_output_list = np.split(