From 532fee5e8951dec607054d3ef463bcb8f6cc15a9 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 26 Jun 2024 22:44:12 +0300 Subject: [PATCH] Add img2img support --- nodes.py | 30 ++++++++++++++++++------------ transport.py | 7 ++++++- 2 files changed, 24 insertions(+), 13 deletions(-) diff --git a/nodes.py b/nodes.py index 42bf02c..f6f3e78 100644 --- a/nodes.py +++ b/nodes.py @@ -403,6 +403,7 @@ class LuminaT2ISampler: }, "optional": { "keep_model_loaded": ("BOOLEAN", {"default": False}), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), } } @@ -412,27 +413,34 @@ class LuminaT2ISampler: CATEGORY = "LuminaWrapper" def process(self, lumina_model, lumina_embeds, latent, seed, steps, cfg, proportional_attn, solver, t_shift, - do_extrapolation, scaling_watershed, keep_model_loaded=False): + do_extrapolation, scaling_watershed, strength=1.0, keep_model_loaded=False): device = mm.get_torch_device() offload_device = mm.unet_offload_device() model = lumina_model['model'] dtype = lumina_model['dtype'] - - z = latent["samples"].clone() - B = z.shape[0] - W = z.shape[3] * 8 - H = z.shape[2] * 8 + vae_scaling_factor = 0.13025 #SDXL scaling factor + + x1 = latent["samples"].clone() * vae_scaling_factor + + ode = ODE(steps, solver, t_shift, strength) + + B = x1.shape[0] + W = x1.shape[3] * 8 + H = x1.shape[2] * 8 + + z = torch.zeros_like(x1) for i in range(B): torch.manual_seed(seed + i) - noise = torch.randn_like(z[i]) - z[i] = z[i] + noise + z[i] = torch.randn_like(x1[i]) + #z[i] = z[i] + noise + z[i] = z[i] * (1 - ode.t[0]) + x1[i] * ode.t[0] #torch.random.manual_seed(int(seed)) #z = torch.randn([1, 4, z.shape[2], z.shape[3]], device=device) - + z = z.repeat(2, 1, 1, 1) z = z.to(dtype).to(device) @@ -481,7 +489,7 @@ class LuminaT2ISampler: #inference model.to(device) - samples = ODE(steps, solver, t_shift).sample(z, model.forward_with_cfg, **model_kwargs)[-1] + samples = ode.sample(z, model.forward_with_cfg, **model_kwargs)[-1] if not keep_model_loaded: print("Offloading Lumina model...") @@ -490,8 +498,6 @@ class LuminaT2ISampler: gc.collect() samples = samples[:len(samples) // 2] - - vae_scaling_factor = 0.13025 #SDXL scaling factor samples = samples / vae_scaling_factor return ({'samples': samples},) diff --git a/transport.py b/transport.py index 5c1c358..8061333 100644 --- a/transport.py +++ b/transport.py @@ -63,9 +63,11 @@ class ODE: num_steps, sampler_type="euler", time_shifting_factor=None, + strength=1.0, t0=0.0, t1=1.0, use_sd3=False, + ): if use_sd3: self.t = th.linspace(t1, t0, num_steps) @@ -75,7 +77,10 @@ class ODE: self.t = th.linspace(t0, t1, num_steps) if time_shifting_factor: self.t = self.t / (self.t + time_shifting_factor - time_shifting_factor * self.t) - + + if strength != 1.0: + self.t = self.t[int(num_steps * (1 - strength)):] + self.use_sd3 = use_sd3 self.sampler_type = sampler_type if self.sampler_type == "euler":