Live preview + progress bar

This commit is contained in:
City
2023-09-05 19:06:44 +02:00
parent 03f06cc801
commit d01bf62c13
2 changed files with 15 additions and 10 deletions
+11 -9
View File
@@ -426,7 +426,8 @@ class GaussianDiffusion:
cond_fn=None, cond_fn=None,
model_kwargs=None, model_kwargs=None,
device=None, device=None,
progress=False, pbar=None,
previewer=None,
): ):
""" """
Generate samples from the model. Generate samples from the model.
@@ -456,7 +457,8 @@ class GaussianDiffusion:
cond_fn=cond_fn, cond_fn=cond_fn,
model_kwargs=model_kwargs, model_kwargs=model_kwargs,
device=device, device=device,
progress=progress, pbar=pbar,
previewer=previewer,
): ):
final = sample final = sample
return final["sample"] return final["sample"]
@@ -471,7 +473,8 @@ class GaussianDiffusion:
cond_fn=None, cond_fn=None,
model_kwargs=None, model_kwargs=None,
device=None, device=None,
progress=False, pbar=None,
previewer=None,
): ):
""" """
Generate samples from the model and yield intermediate samples from Generate samples from the model and yield intermediate samples from
@@ -489,12 +492,6 @@ class GaussianDiffusion:
img = th.randn(*shape, device=device) img = th.randn(*shape, device=device)
indices = list(range(self.num_timesteps))[::-1] indices = list(range(self.num_timesteps))[::-1]
if progress:
# Lazy import so that we don't depend on tqdm.
from tqdm.auto import tqdm
indices = tqdm(indices)
for i in indices: for i in indices:
t = th.tensor([i] * shape[0], device=device) t = th.tensor([i] * shape[0], device=device)
with th.no_grad(): with th.no_grad():
@@ -509,6 +506,11 @@ class GaussianDiffusion:
) )
yield out yield out
img = out["sample"] img = out["sample"]
if pbar:
preview_bytes = None
if previewer:
preview_bytes = previewer.decode_latent_to_preview_image("JPEG", img)
pbar.update_absolute(indices.index(i)+1, self.num_timesteps, preview_bytes)
def ddim_sample( def ddim_sample(
self, self,
+4 -1
View File
@@ -6,6 +6,7 @@ import comfy.model_management
import comfy.model_patcher import comfy.model_patcher
import comfy.utils import comfy.utils
import comfy.latent_formats import comfy.latent_formats
import latent_preview
from .models import DiT_models from .models import DiT_models
from .diffusion import create_diffusion from .diffusion import create_diffusion
@@ -121,6 +122,8 @@ class DiTSampler:
# pre # pre
comfy.model_management.load_model_gpu(model) comfy.model_management.load_model_gpu(model)
real_model = model.model real_model = model.model
pbar = comfy.utils.ProgressBar(steps)
previewer = latent_preview.get_previewer(device, model.model.latent_format)
# Create sampling noise: # Create sampling noise:
z = torch.randn(batch_size, 4, real_model.latent_size, real_model.latent_size, device=device) z = torch.randn(batch_size, 4, real_model.latent_size, real_model.latent_size, device=device)
@@ -134,7 +137,7 @@ class DiTSampler:
# Sample images: # Sample images:
samples = diffusion.p_sample_loop( samples = diffusion.p_sample_loop(
model.model.forward_with_cfg, z.shape, z, clip_denoised=False, model_kwargs=model_kwargs, progress=True, device=device model.model.forward_with_cfg, z.shape, z, clip_denoised=False, model_kwargs=model_kwargs, pbar=pbar, previewer=previewer, device=device
) )
samples, _ = samples.chunk(2, dim=0) # Remove null class samples samples, _ = samples.chunk(2, dim=0) # Remove null class samples
samples = real_model.latent_format.process_out(samples.to(torch.float32)) samples = real_model.latent_format.process_out(samples.to(torch.float32))