control mask: add auto-mask and auto-segment and support for algo masking and rembg masking

This commit is contained in:
Vladimir Mandic
2024-01-16 12:47:39 -05:00
parent 18087f1d14
commit b1fa002ea7
7 changed files with 163 additions and 76 deletions
+19 -7
View File
@@ -2,22 +2,34 @@
## Update for 2023-01-15
Another release with a lot more functionality in the **Control** module and **FaceID/FaceSwap** & **PAdapter** modules
Another big release, highlights being:
- A lot more functionality in the **Control** module:
- Inpaint and outpaint support, flexible resizing options, optional hires
- More processors and models
- Full support for scripts and extensions
- Fully baked **FaceID** / **FaceSwap** & **IPAdapter** modules
- Brand new intelligent masking, manual or automatic using ML models and with live previews
Plus welcome additions to **UI performance, usability and accessibility** and flexibility of deployment
And it also includes fixes for all reported issues so far
- **Control**:
- **Control**:
- add **inpaint** support
applies to both *img2img* and *controlnet* workflows
- add **outpaint** support
applies to both *img2img* and *controlnet* workflows
*note*: increase denoising strength since outpainted area is blank by default
- new **mask** module
- granular blur (gaussian), errode (reduce or remove noise) and dilate (pad or expand)
- granular blur (gaussian), erode (reduce or remove noise) and dilate (pad or expand)
- optional **live preview**
- optional **auto-segmentation** (e.g. segment-anything) using ml models
- optional **auto-segmentation** using ml models
auto-segmentation can be done using **segment-anything** models or **rembg** models
*note*: auto segmentation will automatically expand user-masked area to segments that include current user mask
- can be combined with control processors in which case mask is applied before processor
- optional **auto-mask**
if you don't provide mask or mask is empty, you can instead use auto-mask to automatically generate mask
this is especially useful if you want to use advanced masking on batch or video inputs and don't want to manually mask each image
*note*: such auto-created mask is also subject to all other selected settings such as auto-segmentation, blur, erode and dilate
- masking can be combined with control processors in which case mask is applied before processor
- allow **resize** both *before* and *after* generate operation
this allows for workflows such as: *image -> upscale or downscale -> generate -> upscale or downscale -> output*
providing more flexibility and than standard hires workflow
@@ -25,8 +37,8 @@ And it also includes fixes for all reported issues so far
- implicit **hires**
since hires is only used for txt2img, control reuses existing resize functionality
any image size is used as txt2img target size
but if resize scale is also set its used to additionally upscale image after initial txt2img and for hires pass
- add support for **scripts** and **extensions**
but if resize scale is also set its used to additionally upscale image after initial txt2img and for hires pass
- add support for **scripts** and **extensions**
you can now combine control workflow with your favorite script or extension
*note* extensions that are hard-coded for txt2img or img2img tabs may not work until they are updated
- add **marigold** depth map processor
+3 -3
View File
@@ -343,10 +343,10 @@ function setupExtraNetworksForTab(tabname) {
let searchTimer = null;
txtSearchValue.addEventListener('input', (evt) => {
if (searchTimer) clearTimeout(searchTimer);
searchTimer = setTimeout(() => {
filterExtraNetworksForTab(txtSearchValue.value.toLowerCase());
searchTimer = setTimeout(async () => {
await filterExtraNetworksForTab(txtSearchValue.value.toLowerCase());
searchTimer = null;
}, 150);
}, 50);
});
// card hover
+136 -52
View File
@@ -1,4 +1,5 @@
from types import SimpleNamespace
from typing import List
import os
import time
import gradio as gr
@@ -6,7 +7,7 @@ import numpy as np
import cv2
from PIL import Image, ImageFilter, ImageOps
from transformers import SamModel, SamImageProcessor, MaskGenerationPipeline
from modules import shared, errors, devices, ui_components, ui_symbols
from modules import shared, errors, devices, ui_components, ui_symbols, paths
def get_crop_region(mask, pad=0):
@@ -110,13 +111,16 @@ MODELS = {
'Facebook SAM ViT Huge': 'facebook/sam-vit-huge',
'SlimSAM Uniform': 'Zigeng/SlimSAM-uniform-50',
'SlimSAM Uniform Tiny': 'Zigeng/SlimSAM-uniform-77',
# 'Tiny Random': 'fxmarty/sam-vit-tiny-random',
'Rembg Silueta': 'silueta',
'Rembg U2Net': 'u2net',
'Rembg ISNet': 'isnet',
# "u2net_human_seg",
# "isnet-general-use",
# "isnet-anime",
}
COLORMAP = ['autumn', 'bone', 'jet', 'winter', 'rainbow', 'ocean', 'summer', 'spring', 'cool', 'hsv', 'pink', 'hot', 'parula', 'magma', 'inferno', 'plasma', 'viridis', 'cividis', 'twilight', 'shifted', 'turbo', 'deepgreen']
cache_dir = 'models/control/segment'
loaded_model = None
model: SamModel = None
processor: SamImageProcessor = None
generator: MaskGenerationPipeline = None
debug = shared.log.trace if os.environ.get('SD_MASK_DEBUG', None) is not None else lambda *args, **kwargs: None
debug('Trace: MASK')
@@ -124,6 +128,7 @@ busy = False
btn_segment = None
controls = []
opts = SimpleNamespace(**{
'auto_mask': 'None',
'mask_blur': 0.01,
'mask_erode': 0.01,
'mask_dilate': 0.01,
@@ -134,7 +139,7 @@ opts = SimpleNamespace(**{
'seg_points_per_batch': 64,
'seg_topK': 50,
'seg_colormap': 'pink',
'preview_type': 'composite',
'preview_type': 'Composite',
'seg_live': True,
'weight_original': 0.5,
'weight_mask': 0.5,
@@ -143,18 +148,21 @@ opts = SimpleNamespace(**{
def init_model(selected_model: str):
global busy, loaded_model, model, processor, generator # pylint: disable=global-statement
if selected_model == "None":
if model is not None:
global busy, loaded_model, generator # pylint: disable=global-statement
model_path = MODELS[selected_model]
if model_path is None: # none
if generator is not None:
shared.log.debug('Segment unloading model')
model = None
loaded_model = None
processor = None
generator = None
devices.torch_gc()
return selected_model
model_path = MODELS[selected_model]
if model_path is not None and (loaded_model != selected_model or model is None or processor is None):
if 'Rembg' in selected_model: # rembg
loaded_model = model_path
generator = None
devices.torch_gc()
return selected_model
if loaded_model != selected_model or generator is None: # sam pipeline
busy = True
t0 = time.time()
shared.log.debug(f'Segment loading: model={selected_model} path={model_path}')
@@ -169,6 +177,7 @@ def init_model(selected_model: str):
)
devices.torch_gc()
shared.log.debug(f'Segment loaded: model={selected_model} path={model_path} time={time.time()-t0:.2f}s')
loaded_model = selected_model
busy = False
return selected_model
@@ -195,7 +204,7 @@ def run_segment(input_image: gr.Image, input_mask: np.ndarray):
i = 1
combined_mask = np.zeros(input_mask.shape, dtype='uint8')
input_mask_size = np.count_nonzero(input_mask)
debug(f'Segment: {vars(opts)}')
debug(f'Segment SAM: {vars(opts)}')
for mask in outputs['masks']:
mask = mask.astype('uint8')
mask_size = np.count_nonzero(mask)
@@ -216,6 +225,87 @@ def run_segment(input_image: gr.Image, input_mask: np.ndarray):
return combined_mask
def run_rembg(input_image: Image, input_mask: np.ndarray):
try:
import rembg
except Exception as e:
shared.log.error(f'Segment Rembg load failed: {e}')
return input_mask
if "U2NET_HOME" not in os.environ:
os.environ["U2NET_HOME"] = os.path.join(paths.models_path, "Rembg")
args = {
'data': input_image,
'only_mask': True,
'post_process_mask': False,
'bgcolor': None,
'alpha_matting': False,
'alpha_matting_foreground_threshold': 240,
'alpha_matting_background_threshold': 10,
'alpha_matting_erode_size': int(opts.mask_erode * 40),
'session': rembg.new_session(loaded_model),
}
mask = rembg.remove(**args)
mask = np.array(mask)
if len(input_mask.shape) > 2:
mask = cv2.cvtColor(input_mask, cv2.COLOR_RGB2GRAY)
binary_input = cv2.threshold(input_mask, 127, 255, cv2.THRESH_BINARY | cv2.THRESH_OTSU)[1]
binary_output = cv2.threshold(mask, 127, 255, cv2.THRESH_BINARY | cv2.THRESH_OTSU)[1]
binary_overlap = cv2.bitwise_and(binary_input, binary_output)
input_size = np.count_nonzero(binary_input)
overlap_size = np.count_nonzero(binary_overlap)
debug(f'Segment Rembg: {args} overlap={overlap_size}')
if input_size > 0 and overlap_size == 0:
mask = np.invert(mask)
return mask
def get_mask(input_image: gr.Image, input_mask: gr.Image):
t0 = time.time()
if input_mask is not None:
output_mask = np.array(input_mask)
if len(output_mask.shape) > 2:
output_mask = cv2.cvtColor(output_mask, cv2.COLOR_RGB2GRAY)
binary_mask = cv2.threshold(output_mask, 127, 255, cv2.THRESH_BINARY | cv2.THRESH_OTSU)[1]
mask_size = np.count_nonzero(binary_mask)
else:
output_mask = None
mask_size = 0
if mask_size == 0 and opts.auto_mask != 'None': # mask_size == 0
output_mask = np.array(input_image)
if opts.auto_mask == 'Threshold':
output_mask = cv2.cvtColor(output_mask, cv2.COLOR_RGB2GRAY)
output_mask = cv2.threshold(output_mask, 127, 255, cv2.THRESH_BINARY | cv2.THRESH_OTSU)[1]
elif opts.auto_mask == 'Edge':
output_mask = cv2.cvtColor(output_mask, cv2.COLOR_RGB2GRAY)
output_mask = cv2.threshold(output_mask, 127, 255, cv2.THRESH_BINARY | cv2.THRESH_OTSU)[1]
# output_mask = cv2.Canny(output_mask, 50, 150) # run either canny or threshold before contouring
contours, _hierarchy = cv2.findContours(output_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
contours = sorted(contours, key=cv2.contourArea, reverse=True) # sort contours by area with largest first
contours = contours[:opts.seg_topK] # limit to top K contours
output_mask = np.zeros(output_mask.shape, dtype='uint8')
largest_size = cv2.contourArea(contours[0]) if len(contours) > 0 else 0
for i, contour in enumerate(contours):
area_size = cv2.contourArea(contour)
luminance = int(255.0 * area_size / largest_size)
if luminance < 1:
break
cv2.drawContours(output_mask, contours, i, (luminance), -1)
elif opts.auto_mask == 'Grayscale':
lab_image = cv2.cvtColor(output_mask, cv2.COLOR_RGB2LAB)
l_channel, a, b = cv2.split(lab_image)
clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8)) # applying CLAHE to L-channel
cl = clahe.apply(l_channel)
lab_image = cv2.merge((cl, a, b)) # merge the CLAHE enhanced L-channel with the a and b channel
lab_image = cv2.cvtColor(lab_image, cv2.COLOR_LAB2RGB)
output_mask = cv2.cvtColor(lab_image, cv2.COLOR_RGB2GRAY)
t1 = time.time()
debug(f'Segment auto-mask: mode={opts.auto_mask} time={t1-t0:.2f}')
return output_mask
else: # no mask or empty mask and no auto-mask
return output_mask
def run_mask(input_image: gr.Image, input_mask: gr.Image = None, return_type: str = None, mask_blur: int = None, mask_padding: int = None, segment_enable=True):
if input_image is None:
return input_mask
@@ -224,38 +314,32 @@ def run_mask(input_image: gr.Image, input_mask: gr.Image = None, return_type: st
if isinstance(input_image, dict):
input_mask = input_image.get('mask', None)
input_image = input_image.get('image', None)
if input_mask is None:
input_mask = input_image.convert('L')
input_mask = input_mask.point(lambda x: 255 if x > 127 else 0)
else:
input_mask = input_mask.convert('L')
shared.log.debug(f'Segment mask: input={input_image} mask={input_mask} type={return_type}')
input_mask = np.array(input_mask) // 255
t0 = time.time()
input_mask = get_mask(input_image, input_mask) # perform optional auto-masking
if input_mask is None:
return None
if mask_blur is not None:
if mask_blur is not None: # compatibility with old img2img values which have different range
opts.mask_blur = mask_blur / min(input_image.width, input_image.height)
if mask_padding is not None:
opts.mask_dilate = mask_padding / min(input_image.width, input_image.height)
if generator is None or not segment_enable:
mask = input_mask * 255
if loaded_model is None or not segment_enable:
mask = input_mask
elif generator is None:
mask = run_rembg(input_image, input_mask)
else:
mask = run_segment(input_image, input_mask)
mask = cv2.resize(mask, (input_image.width, input_image.height), interpolation=cv2.INTER_LINEAR)
if mask is None:
shared.log.error('Segment error: no mask')
return input_mask
debug(f'Segment mask: mask={mask.shape}')
if opts.mask_erode > 0:
try:
kernel = np.ones((int(opts.mask_erode * input_image.height / 4) + 1, int(opts.mask_erode * input_image.width / 4) + 1), np.uint8)
cv2_mask = cv2.erode(mask, kernel, iterations=opts.kernel_iterations) # remove noise
mask = cv2_mask
debug(f'Segment erode={opts.mask_erode} kernel={kernel} mask={mask.shape}')
debug(f'Segment erode={opts.mask_erode} kernel={kernel.shape} mask={mask.shape}')
except Exception as e:
shared.log.error(f'Segment erode: {e}')
if opts.mask_dilate > 0:
@@ -263,7 +347,7 @@ def run_mask(input_image: gr.Image, input_mask: gr.Image = None, return_type: st
kernel = np.ones((int(opts.mask_dilate * input_image.height / 4) + 1, int(opts.mask_dilate * input_image.width / 4) + 1), np.uint8)
cv2_mask = cv2.dilate(mask, kernel, iterations=opts.kernel_iterations) # expand area
mask = cv2_mask
debug(f'Segment dilate={opts.mask_dilate} kernel={kernel} mask={mask.shape}')
debug(f'Segment dilate={opts.mask_dilate} kernel={kernel.shape} mask={mask.shape}')
except Exception as e:
shared.log.error(f'Segment dilate: {e}')
if opts.mask_blur > 0:
@@ -281,23 +365,24 @@ def run_mask(input_image: gr.Image, input_mask: gr.Image = None, return_type: st
t1 = time.time()
return_type = return_type or opts.preview_type
shared.log.debug(f'Segment mask opts: size={input_image.width}x{input_image.height} masked={mask_size}px area={area_size/total_size:.2f} time={t1-t0:.2f}')
if return_type == 'none':
shared.log.debug(f'Segment mask: size={input_image.width}x{input_image.height} masked={mask_size}px area={area_size/total_size:.2f} auto={opts.auto_mask} type={return_type} time={t1-t0:.2f}')
if return_type == 'None':
return input_mask
elif return_type == 'binary':
elif return_type == 'Binary':
binary_mask = cv2.threshold(mask, 127, 255, cv2.THRESH_BINARY | cv2.THRESH_OTSU)[1] # otsu uses mean instead of threshold
return Image.fromarray(binary_mask)
elif return_type == 'masked':
elif return_type == 'Masked':
orig = np.array(input_image)
mask = cv2.cvtColor(mask, cv2.COLOR_GRAY2RGB)
masked_image = cv2.bitwise_and(orig, mask)
return Image.fromarray(masked_image)
elif return_type == 'grayscale':
elif return_type == 'Grayscale':
return Image.fromarray(mask)
elif return_type == 'color':
elif return_type == 'Color':
colored_mask = cv2.applyColorMap(mask, COLORMAP.index(opts.seg_colormap)) # recolor mask
return Image.fromarray(colored_mask)
elif return_type == 'composite':
elif return_type == 'Composite':
colored_mask = cv2.applyColorMap(mask, COLORMAP.index(opts.seg_colormap)) # recolor mask
orig = np.array(input_image)
combined_image = cv2.addWeighted(orig, opts.weight_original, colored_mask, opts.weight_mask, 0)
@@ -323,43 +408,42 @@ def create_segment_ui():
opts.mask_blur = args[1]
opts.mask_erode = args[2]
opts.mask_dilate = args[3]
opts.seg_score_thresh = args[4]
opts.seg_iou_thresh = args[5]
opts.seg_nms_thresh = args[6]
opts.preview_type = args[7]
opts.seg_colormap = args[8]
def display_controls(selected_model):
return 4 * [gr.update(visible=True)] + (len(controls) - 4) * [gr.update(visible=selected_model != 'None')]
opts.auto_mask = args[4]
opts.seg_score_thresh = args[5]
opts.seg_iou_thresh = args[6]
opts.seg_nms_thresh = args[7]
opts.preview_type = args[8]
opts.seg_colormap = args[9]
global btn_segment # pylint: disable=global-statement
with gr.Accordion(open=False, label="Mask", elem_id="control_mask", elem_classes=["small-accordion"]):
controls.clear()
with gr.Row():
controls.append(gr.Checkbox(label="Live update", value=True))
btn_segment = ui_components.ToolButton(value=ui_symbols.refresh, visible=True)
with gr.Row():
controls.append(gr.Slider(minimum=0.0, maximum=1.0, step=0.01, label='Blur', value=0.01, elem_id="control_mask_blur"))
controls.append(gr.Slider(minimum=0.0, maximum=1.0, step=0.01, label='Erode', value=0.01, elem_id="control_mask_erode"))
controls.append(gr.Slider(minimum=0.0, maximum=1.0, step=0.01, label='Dilate', value=0.01, elem_id="control_mask_dilate"))
with gr.Row():
controls.append(gr.Dropdown(label="Auto-mask", choices=['None', 'Threshold', 'Edge', 'Grayscale'], value='None'))
selected_model = gr.Dropdown(label="Auto-segment", choices=MODELS.keys(), value='None')
btn_segment = ui_components.ToolButton(value=ui_symbols.refresh, visible=False)
with gr.Row():
controls.append(gr.Slider(minimum=0.0, maximum=1.0, step=0.01, label='Score', value=0.5, visible=False))
controls.append(gr.Slider(minimum=0.0, maximum=1.0, step=0.01, label='IOU', value=0.5, visible=False))
controls.append(gr.Slider(minimum=0.0, maximum=1.0, step=0.01, label='NMS', value=0.5, visible=False))
with gr.Row():
controls.append(gr.Dropdown(label="Preview", choices=['none', 'masked', 'binary', 'grayscale', 'color', 'composite'], value='composite'))
controls.append(gr.Dropdown(label="Preview", choices=['None', 'Masked', 'Binary', 'Grayscale', 'Color', 'Composite'], value='Composite'))
controls.append(gr.Dropdown(label="Colormap", choices=COLORMAP, value='pink'))
selected_model.change(fn=init_model, inputs=[selected_model], outputs=[selected_model])
selected_model.change(fn=display_controls, inputs=[selected_model], outputs=controls)
for control in controls:
control.change(fn=update_opts, inputs=controls, outputs=[])
def bind_controls(input_image: gr.Image, preview_image: gr.Image):
btn_segment.click(run_mask, inputs=[input_image], outputs=[preview_image])
input_image.edit(fn=run_mask_live, inputs=[input_image], outputs=[preview_image])
for control in controls:
control.change(fn=run_mask_live, inputs=[input_image], outputs=[preview_image])
def bind_controls(image_controls: List[gr.Image], preview_image: gr.Image):
for image_control in image_controls:
btn_segment.click(run_mask, inputs=[image_control], outputs=[preview_image])
image_control.edit(fn=run_mask_live, inputs=[image_control], outputs=[preview_image])
for control in controls:
control.change(fn=run_mask_live, inputs=[image_control], outputs=[preview_image])
+2 -11
View File
@@ -148,7 +148,6 @@ def select_input(input_mode, input_image, selected_init, init_type, input_resize
input_mask = masking.run_mask(input_image=selected_input, input_mask=None, return_type='grayscale')
input_source = [selected_input]
input_type = 'PIL.Image'
shared.log.debug(f'Control input: type={input_type} input={input_source}')
status = f'Control input | Image | Size {selected_input.width}x{selected_input.height} | Mode {selected_input.mode}'
res = [gr.Tabs.update(selected='out-gallery'), status]
elif isinstance(selected_input, dict): # inpaint -> dict image+mask
@@ -156,18 +155,15 @@ def select_input(input_mode, input_image, selected_init, init_type, input_resize
selected_input = selected_input['image']
input_source = [selected_input]
input_type = 'PIL.Image'
shared.log.debug(f'Control input: type={input_type} input={input_source} mask={input_mask}')
status = f'Control input | Image | Size {selected_input.width}x{selected_input.height} | Mode {selected_input.mode}'
res = [gr.Tabs.update(selected='out-gallery'), status]
elif isinstance(selected_input, gr.components.image.Image): # not likely
input_source = [selected_input.value]
input_type = 'gr.Image'
shared.log.debug(f'Control input: type={input_type} input={input_source}')
res = [gr.Tabs.update(selected='out-gallery'), status]
elif isinstance(selected_input, str): # video via upload > tmp filepath to video
input_source = selected_input
input_type = 'gr.Video'
shared.log.debug(f'Control input: type={input_type} input={input_source}')
status = get_video(input_source)
res = [gr.Tabs.update(selected='out-video'), status]
elif isinstance(selected_input, list): # batch or folder via upload -> list of tmp filepaths
@@ -178,10 +174,10 @@ def select_input(input_mode, input_image, selected_init, init_type, input_resize
input_type = 'files'
input_source = selected_input
status = f'Control input | Images | Files {len(input_source)}'
shared.log.debug(f'Control input: type={input_type} input={input_source}')
res = [gr.Tabs.update(selected='out-gallery'), status]
else: # unknown
input_source = None
shared.log.debug(f'Control input: type={input_type} input={input_source}')
# init inputs: optional
if init_type == 0: # Control only
input_init = None
@@ -194,7 +190,6 @@ def select_input(input_mode, input_image, selected_init, init_type, input_resize
input_source = [selected_init]
input_init = [selected_init]
input_type = 'PIL.Image'
shared.log.debug(f'Control input: type={input_type} input={input_source}')
status = f'Control input | Image | Size {selected_init.width}x{selected_init.height} | Mode {selected_init.mode}'
res = [gr.Tabs.update(selected='out-gallery'), status]
elif isinstance(selected_init, dict): # inpaint -> dict image+mask
@@ -202,18 +197,15 @@ def select_input(input_mode, input_image, selected_init, init_type, input_resize
input_init = selected_init['image']
input_source = [selected_init]
input_type = 'PIL.Image'
shared.log.debug(f'Control input: type={input_type} input={input_source} mask={input_mask}')
status = f'Control input | Image | Size {selected_init.width}x{selected_init.height} | Mode {selected_input.mode}'
res = [gr.Tabs.update(selected='out-gallery'), status]
elif isinstance(selected_init, gr.components.image.Image): # not likely
input_init = [selected_init.value]
input_type = 'gr.Image'
shared.log.debug(f'Control input: type={input_type} input={input_init}')
res = [gr.Tabs.update(selected='out-gallery'), status]
elif isinstance(selected_init, str): # video via upload > tmp filepath to video
input_init = selected_init
input_type = 'gr.Video'
shared.log.debug(f'Control input: type={input_type} input={input_init}')
status = get_video(input_init)
res = [gr.Tabs.update(selected='out-video'), status]
elif isinstance(selected_init, list): # batch or folder via upload -> list of tmp filepaths
@@ -224,7 +216,6 @@ def select_input(input_mode, input_image, selected_init, init_type, input_resize
input_type = 'files'
input_init = selected_init
status = f'Control input | Images | Files {len(input_init)}'
shared.log.debug(f'Control input: type={input_type} input={input_init} mode={input_mode}')
res = [gr.Tabs.update(selected='out-gallery'), status]
else: # unknown
input_init = None
@@ -702,7 +693,7 @@ def create_ui(_blocks: gr.Blocks=None):
generation_parameters_copypaste.add_paste_fields("control", input_image, paste_fields, override_settings)
bindings = generation_parameters_copypaste.ParamBinding(paste_button=btn_paste, tabname="control", source_text_component=prompt, source_image_component=output_gallery)
generation_parameters_copypaste.register_paste_params_button(bindings)
masking.bind_controls(input_inpaint, preview_process)
masking.bind_controls([input_image, input_inpaint, input_resize], preview_process)
if os.environ.get('SD_CONTROL_DEBUG', None) is not None: # debug only
+1 -1
Submodule wiki updated: 68094a6d4e...cadf034fa7