diff --git a/easycontrol/pipeline.py b/easycontrol/pipeline.py index 47d2177..7cd2473 100644 --- a/easycontrol/pipeline.py +++ b/easycontrol/pipeline.py @@ -529,6 +529,8 @@ class FluxPipeline(DiffusionPipeline, FluxLoraLoaderMixin, FromSingleFileMixin): spatial_images=[], subject_images=[], cond_size=512, + use_zero_init: Optional[bool] = True, + zero_steps: Optional[int] = 0, ): height = height or self.default_sample_size * self.vae_scale_factor @@ -702,6 +704,9 @@ class FluxPipeline(DiffusionPipeline, FluxLoraLoaderMixin, FromSingleFileMixin): return_dict=False, )[0] + if (i <= zero_steps) and use_zero_init: + noise_pred = noise_pred*0. + # compute the previous noisy sample x_t -> x_t-1 latents_dtype = latents.dtype latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0] diff --git a/nodes/comfy_nodes.py b/nodes/comfy_nodes.py index 4bc49e8..171a11a 100644 --- a/nodes/comfy_nodes.py +++ b/nodes/comfy_nodes.py @@ -138,6 +138,8 @@ class EasyControlGenerate: "num_inference_steps": ("INT", {"default": 25, "min": 1, "max": 100, "step": 1}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), "cond_size": ("INT", {"default": 512, "min": 256, "max": 1024, "step": 64}), + "use_zero_init": ("BOOLEAN", {"default": True}), + "zero_steps": ("INT", {"default": 1, "min": 0, "max": 100}), }, "optional": { "spatial_image": ("IMAGE", ), @@ -150,7 +152,7 @@ class EasyControlGenerate: CATEGORY = "EasyControl" def generate(self, pipe, transformer, prompt, prompt_2, height, width, guidance_scale, - num_inference_steps, seed, cond_size, spatial_image=None, subject_image=None): + num_inference_steps, seed, cond_size, use_zero_init, zero_steps, spatial_image=None, subject_image=None): # Clear cache before generation for name, attn_processor in transformer.attn_processors.items(): attn_processor.bank_kv.clear() @@ -220,6 +222,8 @@ class EasyControlGenerate: spatial_images=spatial_images, subject_images=subject_images, cond_size=cond_size, + use_zero_init=use_zero_init, + zero_steps=int(zero_steps) ) # Convert PIL image to numpy array, then to torch.Tensor