diff --git a/py/easyNodes.py b/py/easyNodes.py index 4264d6e..3446215 100644 --- a/py/easyNodes.py +++ b/py/easyNodes.py @@ -2217,7 +2217,8 @@ class icLightApply: "image": ("IMAGE",), "vae": ("VAE",), "lighting": (['None', 'Left Light', 'Right Light', 'Top Light', 'Bottom Light', 'Circle Light'],{"default": "None"}), - "source": (['Use Background Image', 'Use Flipped Background Image', 'Left Light', 'Right Light', 'Top Light', 'Bottom Light', 'Ambient'],{"default": "Use Background Image"}) + "source": (['Use Background Image', 'Use Flipped Background Image', 'Left Light', 'Right Light', 'Top Light', 'Bottom Light', 'Ambient'],{"default": "Use Background Image"}), + "remove_bg": ("BOOLEAN", {"default": True}), }, } @@ -2243,16 +2244,20 @@ class icLightApply: image, _ = results['result'] return image - def apply(self, mode, model, image, vae, lighting, source): + def apply(self, mode, model, image, vae, lighting, source, remove_bg): model_type = get_sd_version(model) if model_type == 'sdxl': raise Exception("IC Light model is not supported for SDXL now") - batch_size = image.shape[0] - if image.shape[3] == 3: + batch_size, height, width, channel = image.shape + if channel == 3: # remove bg if mode == 'Foreground' or batch_size == 1: - image = self.removebg(image) + if remove_bg: + image = self.removebg(image) + else: + mask = torch.full((1, height, width), 1.0, dtype=torch.float32, device="cpu") + image, = JoinImageWithAlpha().join_image_with_alpha(image, mask) iclight = ICLight() if mode == 'Foreground': diff --git a/py/ic_light/func.py b/py/ic_light/func.py index ae301e0..5ef68d2 100644 --- a/py/ic_light/func.py +++ b/py/ic_light/func.py @@ -35,18 +35,17 @@ class VAEEncodeArgMax(VAEEncode): class ICLight: @staticmethod - def apply_c_concat(cond, uncond, c_concat: torch.Tensor): - def write_c_concat(cond): - new_cond = [] - for t in cond: - n = [t[0], t[1].copy()] - if "model_conds" not in n[1]: - n[1]["model_conds"] = {} - n[1]["model_conds"]["c_concat"] = CONDRegular(c_concat) - new_cond.append(n) - return new_cond - - return (write_c_concat(cond), write_c_concat(uncond)) + def apply_c_concat(params: UnetParams, concat_conds) -> UnetParams: + """Apply c_concat on unet call.""" + sample = params["input"] + params["c"]["c_concat"] = torch.cat( + ( + [concat_conds.to(sample.device)] + * (sample.shape[0] // concat_conds.shape[0]) + ), + dim=0, + ) + return params @staticmethod def create_custom_conv( @@ -161,20 +160,19 @@ class ICLight: # [1, 4 * B, H, W] concat_conds = torch.cat([c[None, ...] for c in concat_conds], dim=1) + def unet_dummy_apply(unet_apply: Callable, params: UnetParams): + """A dummy unet apply wrapper serving as the endpoint of wrapper + chain.""" + return unet_apply(x=params["input"], t=params["timestep"], **params["c"]) - def wrapped_unet(unet_apply: Callable, params: UnetParams): - # Apply concat. - sample = params["input"] - params["c"]["c_concat"] = torch.cat( - ( - [concat_conds.to(sample.device)] - * (sample.shape[0] // concat_conds.shape[0]) - ), - dim=0, - ) - return unet_apply(x=sample, t=params["timestep"], **params["c"]) + existing_wrapper = work_model.model_options.get( + "model_function_wrapper", unet_dummy_apply + ) - work_model.set_model_unet_function_wrapper(wrapped_unet) + def wrapper_func(unet_apply: Callable, params: UnetParams): + return existing_wrapper(unet_apply, params=self.apply_c_concat(params, concat_conds)) + + work_model.set_model_unet_function_wrapper(wrapper_func) ic_model = load_unet(ic_model_path) ic_model_state_dict = ic_model.model.diffusion_model.state_dict() diff --git a/web/js/easy/easyDynamicWidgets.js b/web/js/easy/easyDynamicWidgets.js index 4f74ad1..03a777a 100644 --- a/web/js/easy/easyDynamicWidgets.js +++ b/web/js/easy/easyDynamicWidgets.js @@ -128,10 +128,12 @@ function widgetLogic(node, widget) { case 'easy icLightApply': if (widget.value === "Foreground") { toggleWidget(node, findWidgetByName(node, 'lighting'), true) + toggleWidget(node, findWidgetByName(node, 'remove_bg'), true) toggleWidget(node, findWidgetByName(node, 'source')) } else { toggleWidget(node, findWidgetByName(node, 'lighting')) toggleWidget(node, findWidgetByName(node, 'source'), true) + toggleWidget(node, findWidgetByName(node, 'remove_bg')) } break }