diff --git a/modules/impact/core.py b/modules/impact/core.py index de83c7b..ee4bf8d 100644 --- a/modules/impact/core.py +++ b/modules/impact/core.py @@ -18,6 +18,7 @@ from comfy import model_management from impact import utils from impact import impact_sampling from concurrent.futures import ThreadPoolExecutor +from comfy.ldm.cascade.stage_c_coder import StageC_coder SEG = namedtuple("SEG", @@ -214,6 +215,9 @@ def enhance_detail(image, model, clip, vae, guide_size, guide_size_for_bbox, max new_w = w new_h = h + new_w += new_w % 2 + new_h += new_h % 2 + if detailer_hook is not None: new_w, new_h = detailer_hook.touch_scaled_size(new_w, new_h) @@ -232,11 +236,14 @@ def enhance_detail(image, model, clip, vae, guide_size, guide_size_for_bbox, max if noise_mask is not None and inpaint_model: positive, negative, latent_image = nodes.InpaintModelConditioning().encode(positive, negative, upscaled_image, vae, noise_mask) else: - if isinstance(vae, cascade.StageC_coder): - nodes_stable_cascade.StableCascade_StageC_VAEEncode().generate(pixels, vae, compression=compression) - return vae_encode.encode(vae, pixels)[0] + if isinstance(vae.first_stage_model, StageC_coder): + latent_image = detailer_hook.stable_cascade_vae_encode(vae, upscaled_image) + if latent_image is None: + print(f"[Impact Pack] When using the StableCascade model, it is necessary to connect the StableCascade_DetailerHook.") + raise Exception("StableCascade_DetailerHook is not provided.") + else: + latent_image = to_latent_image(upscaled_image, vae) - latent_image = to_latent_image(upscaled_image, vae) if noise_mask is not None: latent_image['noise_mask'] = noise_mask @@ -265,7 +272,7 @@ def enhance_detail(image, model, clip, vae, guide_size, guide_size_for_bbox, max # non-latent downscale - latent downscale cause bad quality if detailer_hook is not None: refined_latent = detailer_hook.pre_decode(refined_latent) - stage_b = detailer_hook.stable_cascade_stage_b(vae, image, positive, negative, refined_latent) + stage_b = detailer_hook.stable_cascade_stage_b(image, positive, negative, refined_latent) else: stage_b = None diff --git a/modules/impact/hook_nodes.py b/modules/impact/hook_nodes.py index d4e3183..03e42da 100644 --- a/modules/impact/hook_nodes.py +++ b/modules/impact/hook_nodes.py @@ -96,8 +96,7 @@ class StableCascade_DetailerHookProvider: "b_cfg": ("FLOAT", {"default": 1.1, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.01}), "b_sampler_name": (comfy.samplers.KSampler.SAMPLERS,), "b_scheduler": (comfy.samplers.KSampler.SCHEDULERS,), - "b_denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), - "c_compression": ("INT", {"default": 4, "min": 4, "max": 128, "step": 1}), + "c_compression": ("INT", {"default": 42, "min": 4, "max": 128, "step": 1}), }, } @@ -106,7 +105,7 @@ class StableCascade_DetailerHookProvider: CATEGORY = "ImpactPack/Util" - def doit(self, b_model, b_vae, b_seed, b_steps, b_cfg, b_sampler_name, b_scheduler, b_denoise, c_compression): - hook = hooks.StableCascade_DetailerHook(b_model, b_vae, b_seed, b_steps, b_cfg, b_sampler_name, b_scheduler, b_denoise, c_compression) + def doit(self, b_model, b_vae, b_seed, b_steps, b_cfg, b_sampler_name, b_scheduler, c_compression): + hook = hooks.StableCascade_DetailerHook(b_model, b_vae, b_seed, b_steps, b_cfg, b_sampler_name, b_scheduler, c_compression) return (hook, ) diff --git a/modules/impact/hooks.py b/modules/impact/hooks.py index 3abc196..4735bbb 100644 --- a/modules/impact/hooks.py +++ b/modules/impact/hooks.py @@ -105,6 +105,13 @@ class DetailerHookCombine(PixelKSampleHookCombine): image = self.hook2.post_paste(image) return image + def stable_cascade_vae_encode(self, vae, pixels): + latent = self.hook1.stable_cascade_vae_encode(vae, pixels) + if latent is not None: + return latent + + return self.hook2.stable_cascade_vae_encode(vae, pixels) + def stable_cascade_stage_b(self, image, positive, negative, latent): image = self.hook1.stable_cascade_stage_b(image, positive, negative, latent) if image is not None: @@ -173,12 +180,15 @@ class DetailerHook(PixelKSampleHook): def post_paste(self, image): return image + def stable_cascade_vae_encode(self, vae, pixels): + return None + def stable_cascade_stage_b(self, image, positive, negative, latent): return None class StableCascade_DetailerHook(DetailerHook): - def __init__(self, b_model, b_vae, b_seed, b_steps, b_cfg, b_sampler_name, b_scheduler, b_denoise, c_compression): + def __init__(self, b_model, b_vae, b_seed, b_steps, b_cfg, b_sampler_name, b_scheduler, c_compression): super().__init__() self.b_model = b_model self.b_vae = b_vae @@ -187,28 +197,25 @@ class StableCascade_DetailerHook(DetailerHook): self.b_cfg = b_cfg self.b_sampler_name = b_sampler_name self.b_scheduler = b_scheduler - self.b_denoise = b_denoise self.c_compression = c_compression - self.orig_size = None + self.b_latent = None - def touch_scaled_size(self, w, h): - compression = min(self.c_compression, w, h) - self.orig_size = w, h - return w // compression, h // compression + def stable_cascade_vae_encode(self, vae, pixels): + obj = comfy_extras.nodes_stable_cascade.StableCascade_StageC_VAEEncode() + stage_c, stage_b = obj.generate(pixels, vae, compression=self.c_compression) + self.b_latent = stage_b + return stage_c def stable_cascade_stage_b(self, image, positive, negative, latent): - w, h = self.orig_size - # prepare stage_b - upscaled_image = utils.tensor_resize(image, w, h) - b_latent = utils.to_latent_image(upscaled_image, self.b_vae) - b_latent['noise_mask'] = latent['noise_mask'] + # self.b_latent['noise_mask'] = latent['noise_mask'] b_positive = comfy_extras.nodes_stable_cascade.StableCascade_StageB_Conditioning().set_prior(positive, latent)[0] # stage_b sampling - b_latent = impact_sampling.ksampler_wrapper(self.b_model, self.b_seed, self.b_steps, self.b_cfg, self.b_sampler_name, self.b_scheduler, b_positive, negative, b_latent, self.b_denoise) + b_latent = impact_sampling.ksampler_wrapper(self.b_model, self.b_seed, self.b_steps, self.b_cfg, self.b_sampler_name, self.b_scheduler, b_positive, negative, self.b_latent, 1.0) # stage_b decoding + self.b_latent = None return self.b_vae.decode(b_latent['samples'])