Add img2img support

This commit is contained in:
kijai
2024-06-26 22:44:12 +03:00
parent 1eaded2d7b
commit 532fee5e89
2 changed files with 24 additions and 13 deletions
+16 -10
View File
@@ -403,6 +403,7 @@ class LuminaT2ISampler:
}, },
"optional": { "optional": {
"keep_model_loaded": ("BOOLEAN", {"default": False}), "keep_model_loaded": ("BOOLEAN", {"default": False}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
} }
} }
@@ -412,23 +413,30 @@ class LuminaT2ISampler:
CATEGORY = "LuminaWrapper" CATEGORY = "LuminaWrapper"
def process(self, lumina_model, lumina_embeds, latent, seed, steps, cfg, proportional_attn, solver, t_shift, 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() device = mm.get_torch_device()
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
model = lumina_model['model'] model = lumina_model['model']
dtype = lumina_model['dtype'] dtype = lumina_model['dtype']
z = latent["samples"].clone() vae_scaling_factor = 0.13025 #SDXL scaling factor
B = z.shape[0] x1 = latent["samples"].clone() * vae_scaling_factor
W = z.shape[3] * 8
H = z.shape[2] * 8 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): for i in range(B):
torch.manual_seed(seed + i) torch.manual_seed(seed + i)
noise = torch.randn_like(z[i]) z[i] = torch.randn_like(x1[i])
z[i] = z[i] + noise #z[i] = z[i] + noise
z[i] = z[i] * (1 - ode.t[0]) + x1[i] * ode.t[0]
#torch.random.manual_seed(int(seed)) #torch.random.manual_seed(int(seed))
#z = torch.randn([1, 4, z.shape[2], z.shape[3]], device=device) #z = torch.randn([1, 4, z.shape[2], z.shape[3]], device=device)
@@ -481,7 +489,7 @@ class LuminaT2ISampler:
#inference #inference
model.to(device) 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: if not keep_model_loaded:
print("Offloading Lumina model...") print("Offloading Lumina model...")
@@ -490,8 +498,6 @@ class LuminaT2ISampler:
gc.collect() gc.collect()
samples = samples[:len(samples) // 2] samples = samples[:len(samples) // 2]
vae_scaling_factor = 0.13025 #SDXL scaling factor
samples = samples / vae_scaling_factor samples = samples / vae_scaling_factor
return ({'samples': samples},) return ({'samples': samples},)
+5
View File
@@ -63,9 +63,11 @@ class ODE:
num_steps, num_steps,
sampler_type="euler", sampler_type="euler",
time_shifting_factor=None, time_shifting_factor=None,
strength=1.0,
t0=0.0, t0=0.0,
t1=1.0, t1=1.0,
use_sd3=False, use_sd3=False,
): ):
if use_sd3: if use_sd3:
self.t = th.linspace(t1, t0, num_steps) self.t = th.linspace(t1, t0, num_steps)
@@ -76,6 +78,9 @@ class ODE:
if time_shifting_factor: if time_shifting_factor:
self.t = self.t / (self.t + time_shifting_factor - time_shifting_factor * self.t) 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.use_sd3 = use_sd3
self.sampler_type = sampler_type self.sampler_type = sampler_type
if self.sampler_type == "euler": if self.sampler_type == "euler":