Skip to content

Commit ae81b37

Browse files
Merge pull request #5165 from klimaleksus/fix-sequential-vae
Make VAE step sequential to prevent VRAM spikes, will fix #3059, #2082, #2561, #3462
2 parents c377777 + 67efee3 commit ae81b37

File tree

1 file changed

+2
-2
lines changed

1 file changed

+2
-2
lines changed

modules/processing.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -530,8 +530,8 @@ def infotext(iteration=0, position_in_batch=0):
530530
with devices.autocast():
531531
samples_ddim = p.sample(conditioning=c, unconditional_conditioning=uc, seeds=seeds, subseeds=subseeds, subseed_strength=p.subseed_strength, prompts=prompts)
532532

533-
samples_ddim = samples_ddim.to(devices.dtype_vae)
534-
x_samples_ddim = decode_first_stage(p.sd_model, samples_ddim)
533+
x_samples_ddim = [decode_first_stage(p.sd_model, samples_ddim[i:i+1].to(dtype=devices.dtype_vae))[0].cpu() for i in range(samples_ddim.size(0))]
534+
x_samples_ddim = torch.stack(x_samples_ddim).float()
535535
x_samples_ddim = torch.clamp((x_samples_ddim + 1.0) / 2.0, min=0.0, max=1.0)
536536

537537
del samples_ddim

0 commit comments

Comments
 (0)