Live preview + progress bar
This commit is contained in:
@@ -426,7 +426,8 @@ class GaussianDiffusion:
|
||||
cond_fn=None,
|
||||
model_kwargs=None,
|
||||
device=None,
|
||||
progress=False,
|
||||
pbar=None,
|
||||
previewer=None,
|
||||
):
|
||||
"""
|
||||
Generate samples from the model.
|
||||
@@ -456,7 +457,8 @@ class GaussianDiffusion:
|
||||
cond_fn=cond_fn,
|
||||
model_kwargs=model_kwargs,
|
||||
device=device,
|
||||
progress=progress,
|
||||
pbar=pbar,
|
||||
previewer=previewer,
|
||||
):
|
||||
final = sample
|
||||
return final["sample"]
|
||||
@@ -471,7 +473,8 @@ class GaussianDiffusion:
|
||||
cond_fn=None,
|
||||
model_kwargs=None,
|
||||
device=None,
|
||||
progress=False,
|
||||
pbar=None,
|
||||
previewer=None,
|
||||
):
|
||||
"""
|
||||
Generate samples from the model and yield intermediate samples from
|
||||
@@ -489,12 +492,6 @@ class GaussianDiffusion:
|
||||
img = th.randn(*shape, device=device)
|
||||
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:
|
||||
t = th.tensor([i] * shape[0], device=device)
|
||||
with th.no_grad():
|
||||
@@ -509,6 +506,11 @@ class GaussianDiffusion:
|
||||
)
|
||||
yield out
|
||||
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(
|
||||
self,
|
||||
|
||||
@@ -6,6 +6,7 @@ import comfy.model_management
|
||||
import comfy.model_patcher
|
||||
import comfy.utils
|
||||
import comfy.latent_formats
|
||||
import latent_preview
|
||||
|
||||
from .models import DiT_models
|
||||
from .diffusion import create_diffusion
|
||||
@@ -121,6 +122,8 @@ class DiTSampler:
|
||||
# pre
|
||||
comfy.model_management.load_model_gpu(model)
|
||||
real_model = model.model
|
||||
pbar = comfy.utils.ProgressBar(steps)
|
||||
previewer = latent_preview.get_previewer(device, model.model.latent_format)
|
||||
|
||||
# Create sampling noise:
|
||||
z = torch.randn(batch_size, 4, real_model.latent_size, real_model.latent_size, device=device)
|
||||
@@ -134,7 +137,7 @@ class DiTSampler:
|
||||
|
||||
# Sample images:
|
||||
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 = real_model.latent_format.process_out(samples.to(torch.float32))
|
||||
|
||||
Reference in New Issue
Block a user