add:remove_bg in easy icLightApply
This commit is contained in:
+10
-5
@@ -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
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user