Team Ai
Apppublic

KalbeDigitalLab/pathology_nuclei_segmentation_classification

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
app.py107 linesDownload Raw Back to root
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