KalbeDigitalLab/pathology_nuclei_segmentation_classification
1
1import json2import os3from pathlib import Path4 5import gradio as gr6import numpy as np7import torch8from monai.bundle import ConfigParser9 10from utils import page_utils11 12with open("configs/inference.json") as f:13 inference_config = json.load(f)14 15device = torch.device('cpu')16if torch.cuda.is_available():17 device = torch.device('cuda:0')18 19# * NOTE: device must be hardcoded, config file won't affect the device selection20inference_config["device"] = device21 22parser = ConfigParser()23parser.read_config(f=inference_config)24parser.read_meta(f="configs/metadata.json")25 26inference = parser.get_parsed_content("inferer")27# loader = parser.get_parsed_content("dataloader")28network = parser.get_parsed_content("network_def")29preprocess = parser.get_parsed_content("preprocessing")30postprocess = parser.get_parsed_content("postprocessing")31 32use_fp16 = os.environ.get('USE_FP16', False)33 34state_dict = torch.load("models/model.pt")35network.load_state_dict(state_dict, strict=True)36 37network = network.to(device)38network.eval()39 40if use_fp16 and torch.cuda.is_available():41 network = network.half()42 43label2color = {0: (0, 0, 0),44 1: (225, 24, 69), # RED45 2: (135, 233, 17), # GREEN46 3: (0, 87, 233), # BLUE47 4: (242, 202, 25), # YELLOW48 5: (137, 49, 239),} # PURPLE49 50example_files = list(Path("sample_data").glob("*.png"))51 52def visualize_instance_seg_mask(mask):53 image = np.zeros((mask.shape[0], mask.shape[1], 3))54 labels = np.unique(mask)55 for i in range(image.shape[0]):56 for j in range(image.shape[1]):57 image[i, j, :] = label2color[mask[i, j]]58 image = image / 25559 return image60 61def query_image(img):62 data = {"image": img}63 batch = preprocess(data)64 batch['image'] = batch['image'].to(device)65 66 if use_fp16 and torch.cuda.is_available():67 batch['image'] = batch['image'].half()68 69 with torch.no_grad():70 pred = inference(batch['image'].unsqueeze(dim=0), network)71 72 batch["pred"] = pred73 for k,v in batch["pred"].items():74 batch["pred"][k] = v.squeeze(dim=0)75 76 batch = postprocess(batch)77 78 result = visualize_instance_seg_mask(batch["type_map"].squeeze())79 80 # Combine image81 result = batch["image"].permute(1, 2, 0).cpu().numpy() * 0.5 + result * 0.582 83 # Solve rotating problem84 result = np.fliplr(result)85 result = np.rot90(result, k=1)86 87 return result88 89# load Markdown file90with open('index.html', encoding='utf-8') as f:91 html_content = f.read()92 93demo = gr.Interface(94 query_image,95 inputs=[gr.Image(type="filepath")],96 outputs="image",97 theme=gr.themes.Default(primary_hue=page_utils.KALBE_THEME_COLOR, secondary_hue=page_utils.KALBE_THEME_COLOR).set(98 button_primary_background_fill="*primary_600",99 button_primary_background_fill_hover="*primary_500",100 button_primary_text_color="white",101 ),102 description = html_content,103 examples=example_files,104)105 106demo.queue(max_size=10).launch()107 