Live preview + progress bar
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|||||||
Reference in New Issue
Block a user