From a163d1896b3dc38d99ac44750771d9286b4ad41a Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 13 Jan 2024 19:35:47 -0500 Subject: [PATCH] segment prototype --- modules/segment.py | 124 ++++++++++++++++++++++++++++++++++++++++++ modules/ui_control.py | 4 +- wiki | 2 +- 3 files changed, 128 insertions(+), 2 deletions(-) create mode 100644 modules/segment.py diff --git a/modules/segment.py b/modules/segment.py new file mode 100644 index 000000000..7f1405dc1 --- /dev/null +++ b/modules/segment.py @@ -0,0 +1,124 @@ +""" +[docs](https://huggingface.co/docs/transformers/v4.36.1/en/model_doc/sam#overview) +TODO: +- PerSAM +- transformers.pipeline.MaskGenerationPipeline: https://huggingface.co/models?pipeline_tag=mask-generation +- transformers.pipeline.ImageSegmentationPipeline: https://huggingface.co/models?pipeline_tag=image-segmentation +""" + +from transformers import SamModel, SamImageProcessor, MaskGenerationPipeline +from PIL import Image +import gradio as gr +import numpy as np +import cv2 +from modules import shared, devices + + +MODELS = { + 'None': None, + 'Facebook SAM ViT Base': 'facebook/sam-vit-base', + 'Facebook SAM ViT Large': 'facebook/sam-vit-large', + 'Facebook SAM ViT Huge': 'facebook/sam-vit-huge', + 'SlimSAM Uniform': 'Zigeng/SlimSAM-uniform-50', +} +cache_dir = 'models/control/segment' +loaded_model = None +model: SamModel = None +processor: SamImageProcessor = None + + +def init(selected_model: str, input_image: gr.Image): + global loaded_model, model, processor # pylint: disable=global-statement + if input_image is None or input_image.get('image', None) is None: + return False + if selected_model == "None": + return False + if selected_model == "None": + model = None + loaded_model = None + processor = None + model_path = MODELS[selected_model] + if model_path is not None and (loaded_model != selected_model or model is None or processor is None): + shared.log.debug(f'Segment loading: model={selected_model} path={model_path}') + model = SamModel.from_pretrained(model_path, cache_dir=cache_dir).to(device=devices.device) + processor = SamImageProcessor.from_pretrained(model_path, cache_dir=cache_dir) + shared.log.debug(f'Segment loaded: model={selected_model} path={model_path}') + if model is None or processor is None: + return False + return True + + +# run as auto-mask with all possible masks +def run_segment(selected_model: str, input_image: gr.Image): + if not init(selected_model, input_image): + return input_image + input_mask = input_image.get('mask', None) or Image.new('L', input_image.get('image', None).size, 255) + input_image = input_image.get('image', None) + generator: MaskGenerationPipeline = MaskGenerationPipeline(model=model, image_processor=processor, device=devices.device) + with devices.inference_context(): + outputs = generator( + input_image, + points_per_batch=64, + pred_iou_thresh=0.75, + stability_score_thresh=0.95, + crops_nms_thresh=0.7, + crop_overlap_ratio=0.3, + ) + combined_mask = np.zeros(input_mask.size, dtype='uint8') + for i, mask in enumerate(outputs['masks']): + mask = mask.astype('uint8') * i * 10 + combined_mask = combined_mask + mask + total_size = np.prod(mask.shape) + area_size = np.count_nonzero(mask) + shared.log.debug(f'Segment mask: i={i} area={area_size/total_size:.2f} score={outputs["scores"][i].item():.2f}') + if i > 25: + break + combined_mask = cv2.applyColorMap(combined_mask, cv2.COLORMAP_JET) + combined_image = cv2.addWeighted(np.array(input_image), 0.6, combined_mask, 0.4, 0) + combined_mask = Image.fromarray(combined_mask) + combined_image = Image.fromarray(combined_image) + combined_mask.save('/tmp/mask-combined.png') + combined_image.save('/tmp/mask-combined-image.png') + + +# run with sam model directly needing set of points +def run_segment_points(selected_model: str, input_image: gr.Image): + if not init(selected_model, input_image): + return input_image + input_mask = input_image.get('mask', None) or Image.new('L', input_image.get('image', None).size, 0) + input_image = input_image.get('image', None) + with devices.inference_context(): + inputs = processor( + input_image, + input_points=[[[256, 256]]], # TODO calculate points based on mask + return_tensors="pt" + ).to(device=devices.device) + outputs = model( + pixel_values=inputs['pixel_values'], + multimask_output=True, + ) + masks = processor.post_process_masks( + outputs.pred_masks.cpu(), + inputs["original_sizes"].cpu(), + inputs["reshaped_input_sizes"].cpu() + ) + scores = outputs.iou_scores + mask = masks[0].squeeze(0) + scores = scores[0].squeeze(0) + masks = mask.unbind(0) + output_masks = [] + for i, mask in enumerate(masks): + mask = mask.detach().cpu().numpy() + mask = mask.astype('uint8') * 255 + mask = cv2.dilate(mask, np.ones((3, 3), np.uint8), iterations=2) + total_size = np.prod(mask.shape) + area_size = np.count_nonzero(mask) + shared.log.debug(f'Segment mask: area={area_size/total_size:.2f} score={scores[i].item():.2f}') + mask = Image.fromarray(mask) + output_masks.append(mask) + + +def create_segment_ui(input_image: gr.Image): + selected = gr.Dropdown(label="Segment", choices=MODELS.keys(), value='None') + selected.change(fn=run_segment, inputs=[selected, input_image], outputs=[]) + return selected diff --git a/modules/ui_control.py b/modules/ui_control.py index 0d94c45d3..a39e323d4 100644 --- a/modules/ui_control.py +++ b/modules/ui_control.py @@ -12,7 +12,7 @@ from modules.control.units import lite # vislearn ControlNet-XS from modules.control.units import t2iadapter # TencentARC T2I-Adapter from modules.control.units import reference # reference pipeline from scripts import ipadapter # pylint: disable=no-name-in-module -from modules import errors, shared, progress, sd_samplers, ui_components, ui_symbols, ui_common, ui_sections, generation_parameters_copypaste, call_queue, scripts # pylint: disable=ungrouped-imports +from modules import errors, shared, progress, sd_samplers, ui_components, ui_symbols, ui_common, ui_sections, generation_parameters_copypaste, call_queue, scripts, segment # pylint: disable=ungrouped-imports gr_height = 512 @@ -368,6 +368,8 @@ def create_ui(_blocks: gr.Blocks=None): input_resize = gr.Image(label="Input", show_label=False, type="pil", source="upload", interactive=True, tool="select", height=gr_height, visible=False, image_mode='RGB', elem_id='control_input_resize') input_inpaint = gr.Image(label="Input", show_label=False, type="pil", source="upload", interactive=True, tool="sketch", height=gr_height, visible=False, image_mode='RGB', elem_id='control_input_inpaint', brush_radius=64, mask_opacity=0.6) interrogate_clip, interrogate_booru = ui_sections.create_interrogate_buttons('control') + # with gr.Row(): + # segment_ui = segment.create_segment_ui(input_inpaint) with gr.Row(): input_buttons = [gr.Button('Select', visible=True, interactive=False), gr.Button('Inpaint', visible=True, interactive=True), gr.Button('Outpaint', visible=True, interactive=True)] with gr.Tab('Video', id='in-video') as tab_video: diff --git a/wiki b/wiki index 43d93f689..8fd5ddfef 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 43d93f689a5ebdc666db4184e34871f3dcec3423 +Subproject commit 8fd5ddfefffba1d6ee4da8d3bfb30b755895a7c3