add:remove_bg in easy icLightApply

This commit is contained in:
yolain
2024-05-11 02:13:55 +08:00
parent a84f7c4a58
commit 1cea58c7cf
3 changed files with 34 additions and 29 deletions
+10 -5
View File
@@ -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':
+22 -24
View File
@@ -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()
+2
View File
@@ -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
}