diff --git a/comfy/sd.py b/comfy/sd.py index 4bdaa978c..d00fdb73f 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -31,6 +31,7 @@ import comfy.weight_adapter import yaml import math import os +from tqdm import tqdm import comfy.utils import comfy.ops @@ -1112,11 +1113,11 @@ class VAE: def vae_output_dtype(self): return model_management.intermediate_dtype() - def decode_tiled_(self, samples, tile_x=64, tile_y=64, overlap = 16): + def decode_tiled_(self, samples, tile_x=64, tile_y=64, overlap = 16, term_pbar_desc=None): steps = samples.shape[0] * comfy.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x, tile_y, overlap) steps += samples.shape[0] * comfy.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x // 2, tile_y * 2, overlap) steps += samples.shape[0] * comfy.utils.get_tiled_scale_steps(samples.shape[3], samples.shape[2], tile_x * 2, tile_y // 2, overlap) - pbar = comfy.utils.ProgressBar(steps) + pbar = comfy.utils.ProgressBar(steps, term_desc=term_pbar_desc) decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) output = self.process_output( @@ -1126,7 +1127,7 @@ class VAE: / 3.0) return output - def decode_tiled_1d(self, samples, tile_x=256, overlap=32): + def decode_tiled_1d(self, samples, tile_x=256, overlap=32, term_pbar_desc=None): if samples.ndim == 3: decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) else: @@ -1134,21 +1135,21 @@ class VAE: samples = samples.reshape((og_shape[0], og_shape[1] * og_shape[2], -1)) decode_fn = lambda a: self.first_stage_model.decode(a.reshape((-1, og_shape[1], og_shape[2], a.shape[-1])).to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) - return self.process_output(comfy.utils.tiled_scale_multidim(samples, decode_fn, tile=(tile_x,), overlap=overlap, upscale_amount=self.upscale_ratio, out_channels=self.output_channels, output_device=self.output_device)) + return self.process_output(comfy.utils.tiled_scale_multidim(samples, decode_fn, tile=(tile_x,), overlap=overlap, upscale_amount=self.upscale_ratio, out_channels=self.output_channels, output_device=self.output_device, term_pbar_desc=term_pbar_desc)) - def decode_tiled_3d(self, samples, tile_t=999, tile_x=32, tile_y=32, overlap=(1, 8, 8)): + def decode_tiled_3d(self, samples, tile_t=999, tile_x=32, tile_y=32, overlap=(1, 8, 8), term_pbar_desc=None): decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) - return self.process_output(comfy.utils.tiled_scale_multidim(samples, decode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.upscale_ratio, out_channels=self.output_channels, index_formulas=self.upscale_index_formula, output_device=self.output_device)) + return self.process_output(comfy.utils.tiled_scale_multidim(samples, decode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.upscale_ratio, out_channels=self.output_channels, index_formulas=self.upscale_index_formula, output_device=self.output_device, term_pbar_desc=term_pbar_desc)) def _decode_tiled_owned(self, samples, **kwargs): out = self.first_stage_model.decode_tiled(samples.to(self.vae_dtype).to(self.device), **kwargs) return self.process_output(out.to(device=self.output_device, dtype=self.vae_output_dtype(), copy=True)) - def encode_tiled_(self, pixel_samples, tile_x=512, tile_y=512, overlap = 64): + def encode_tiled_(self, pixel_samples, tile_x=512, tile_y=512, overlap = 64, term_pbar_desc=None): steps = pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x, tile_y, overlap) steps += pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x // 2, tile_y * 2, overlap) steps += pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x * 2, tile_y // 2, overlap) - pbar = comfy.utils.ProgressBar(steps) + pbar = comfy.utils.ProgressBar(steps, term_desc=term_pbar_desc) encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) samples = comfy.utils.tiled_scale(pixel_samples, encode_fn, tile_x, tile_y, overlap, upscale_amount = (1/self.downscale_ratio), out_channels=self.latent_channels, output_device=self.output_device, pbar=pbar) @@ -1157,7 +1158,7 @@ class VAE: samples /= 3.0 return samples - def encode_tiled_1d(self, samples, tile_x=256 * 2048, overlap=64 * 2048): + def encode_tiled_1d(self, samples, tile_x=256 * 2048, overlap=64 * 2048, term_pbar_desc=None): if self.latent_dim == 1: encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) out_channels = self.latent_channels @@ -1170,15 +1171,15 @@ class VAE: upscale_amount = 1 / self.downscale_ratio encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).reshape(1, out_channels, -1).to(dtype=self.vae_output_dtype()) - out = comfy.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_x,), overlap=overlap, upscale_amount=upscale_amount, out_channels=out_channels, output_device=self.output_device) + out = comfy.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_x,), overlap=overlap, upscale_amount=upscale_amount, out_channels=out_channels, output_device=self.output_device, term_pbar_desc=term_pbar_desc) if self.latent_dim == 1: return out else: return out.reshape(samples.shape[0], self.latent_channels, extra_channel_size, -1) - def encode_tiled_3d(self, samples, tile_t=9999, tile_x=512, tile_y=512, overlap=(1, 64, 64)): + def encode_tiled_3d(self, samples, tile_t=9999, tile_x=512, tile_y=512, overlap=(1, 64, 64), term_pbar_desc=None): encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) - return comfy.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.downscale_ratio, out_channels=self.latent_channels, downscale=True, index_formulas=self.downscale_index_formula, output_device=self.output_device) + return comfy.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.downscale_ratio, out_channels=self.latent_channels, downscale=True, index_formulas=self.downscale_index_formula, output_device=self.output_device, term_pbar_desc=term_pbar_desc) def _encode_tiled_owned(self, pixel_samples, **kwargs): x = self.process_input(pixel_samples).to(self.vae_dtype).to(self.device) @@ -1199,7 +1200,7 @@ class VAE: args["overlap_t"] = overlap_t return args - def decode(self, samples_in, vae_options={}): + def decode(self, samples_in, vae_options={}, term_pbar_desc="VaeDecode"): self.throw_exception_if_invalid() pixel_samples = None do_tile = False @@ -1220,7 +1221,7 @@ class VAE: pixel_samples = torch.empty(self.first_stage_model.decode_output_shape(samples_in.shape), device=self.output_device, dtype=self.vae_output_dtype()) preallocated = True - for x in range(0, samples_in.shape[0], batch_number): + for x in tqdm(range(0, samples_in.shape[0], batch_number), desc=term_pbar_desc, disable=term_pbar_desc is None): samples = samples_in[x:x + batch_number].to(device=self.device, dtype=self.vae_dtype) if preallocated: self.first_stage_model.decode(samples, output_buffer=pixel_samples[x:x+batch_number], **vae_options) @@ -1245,14 +1246,14 @@ class VAE: comfy.model_management.soft_empty_cache() dims = samples_in.ndim - 2 if dims == 1 or self.extra_1d_channel is not None: - pixel_samples = self.decode_tiled_1d(samples_in) + pixel_samples = self.decode_tiled_1d(samples_in, term_pbar_desc=term_pbar_desc) elif dims == 2: if self.handles_tiling: tile = 256 // self.spacial_compression_decode() overlap = tile // 4 pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap) else: - pixel_samples = self.decode_tiled_(samples_in) + pixel_samples = self.decode_tiled_(samples_in, term_pbar_desc=term_pbar_desc) elif dims == 3: tile = 256 // self.spacial_compression_decode() overlap = tile // 4 @@ -1272,7 +1273,7 @@ class VAE: while tile * 2 <= max(samples_in.shape[3], samples_in.shape[4]) and est(tile_t, tile * 2) <= budget: tile *= 2 overlap = tile // 4 - pixel_samples = self.decode_tiled_3d(samples_in, tile_t=tile_t, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) + pixel_samples = self.decode_tiled_3d(samples_in, tile_t=tile_t, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap), term_pbar_desc=term_pbar_desc) pixel_samples = pixel_samples.to(self.output_device).movedim(1,-1) return pixel_samples @@ -1296,7 +1297,7 @@ class VAE: s[-1] = min(s[-1], tile_x) return tuple(s) - def decode_tiled(self, samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): + def decode_tiled(self, samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None, term_pbar_desc="VaeDecode"): self.throw_exception_if_invalid() memory_used = self.memory_used_decode(self._tile_bounded_shape(samples.shape, tile_x, tile_y, tile_t), self.vae_dtype) model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload) @@ -1314,9 +1315,9 @@ class VAE: output = self._decode_tiled_owned(samples, **self._owned_tiled_args(tile_x, tile_y, overlap, tile_t, overlap_t)) elif dims == 1 or self.extra_1d_channel is not None: args.pop("tile_y") - output = self.decode_tiled_1d(samples, **args) + output = self.decode_tiled_1d(samples, **args, term_pbar_desc=term_pbar_desc) elif dims == 2: - output = self.decode_tiled_(samples, **args) + output = self.decode_tiled_(samples, **args, term_pbar_desc=term_pbar_desc) elif dims == 3: if overlap_t is None: args["overlap"] = (1, overlap, overlap) @@ -1325,10 +1326,10 @@ class VAE: if tile_t is not None: args["tile_t"] = max(2, tile_t) - output = self.decode_tiled_3d(samples, **args) + output = self.decode_tiled_3d(samples, **args, term_pbar_desc=term_pbar_desc) return output.movedim(1, -1) - def encode(self, pixel_samples): + def encode(self, pixel_samples, term_pbar_desc="VaeEncode"): self.throw_exception_if_invalid() pixel_samples = self.vae_encode_crop_pixels(pixel_samples) pixel_samples = pixel_samples.movedim(-1, 1) @@ -1347,7 +1348,7 @@ class VAE: batch_number = int(free_memory / max(1, memory_used)) batch_number = max(1, batch_number) samples = None - for x in range(0, pixel_samples.shape[0], batch_number): + for x in tqdm(range(0, pixel_samples.shape[0], batch_number), desc=term_pbar_desc, disable=term_pbar_desc is None): pixels_in = self.process_input(pixel_samples[x:x + batch_number]).to(self.vae_dtype) if getattr(self.first_stage_model, 'comfy_has_chunked_io', False): out = self.first_stage_model.encode(pixels_in, device=self.device) @@ -1376,17 +1377,17 @@ class VAE: if self.handles_tiling: samples = self._encode_tiled_owned(pixel_samples, tile_x=tile, tile_y=tile, overlap=overlap) else: - samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) + samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap), term_pbar_desc=term_pbar_desc) elif self.latent_dim == 1 or self.extra_1d_channel is not None: - samples = self.encode_tiled_1d(pixel_samples) + samples = self.encode_tiled_1d(pixel_samples, term_pbar_desc=term_pbar_desc) else: - samples = self.encode_tiled_(pixel_samples) + samples = self.encode_tiled_(pixel_samples, term_pbar_desc=term_pbar_desc) if self.format_encoded is not None: samples = self.format_encoded(samples) return samples - def encode_tiled(self, pixel_samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): + def encode_tiled(self, pixel_samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None, term_pbar_desc="VaeEncode"): self.throw_exception_if_invalid() pixel_samples = self.vae_encode_crop_pixels(pixel_samples) dims = self.latent_dim @@ -1411,9 +1412,9 @@ class VAE: with model_management.cuda_device_context(self.device): if dims == 1: args.pop("tile_y") - samples = self.encode_tiled_1d(pixel_samples, **args) + samples = self.encode_tiled_1d(pixel_samples, **args, term_pbar_desc=term_pbar_desc) elif dims == 2: - samples = self.encode_tiled_(pixel_samples, **args) + samples = self.encode_tiled_(pixel_samples, **args, term_pbar_desc=term_pbar_desc) elif dims == 3: if self.handles_tiling: samples = self._encode_tiled_owned(pixel_samples, **self._owned_tiled_args(tile_x, tile_y, overlap, tile_t, overlap_t)) @@ -1432,7 +1433,7 @@ class VAE: maximum = pixel_samples.shape[2] maximum = self.upscale_ratio[0](self.downscale_ratio[0](maximum)) - samples = self.encode_tiled_3d(pixel_samples[:,:,:maximum], **args) + samples = self.encode_tiled_3d(pixel_samples[:,:,:maximum], **args, term_pbar_desc=term_pbar_desc) if self.format_encoded is not None: samples = self.format_encoded(samples) diff --git a/comfy/utils.py b/comfy/utils.py index 61c2a22dd..37786592d 100644 --- a/comfy/utils.py +++ b/comfy/utils.py @@ -32,6 +32,7 @@ from torch.nn.functional import interpolate from tqdm.auto import trange from einops import rearrange from comfy.cli_args import args +from tqdm import tqdm import json import time import threading @@ -1106,7 +1107,7 @@ def get_tiled_scale_steps(width, height, tile_x, tile_y, overlap): return rows * cols @torch.inference_mode() -def tiled_scale_multidim(samples, function, tile=(64, 64), overlap=8, upscale_amount=4, out_channels=3, output_device="cpu", downscale=False, index_formulas=None, pbar=None): +def tiled_scale_multidim(samples, function, tile=(64, 64), overlap=8, upscale_amount=4, out_channels=3, output_device="cpu", downscale=False, index_formulas=None, pbar=None, term_pbar_desc=None): dims = len(tile) if not (isinstance(upscale_amount, (tuple, list))): @@ -1164,66 +1165,80 @@ def tiled_scale_multidim(samples, function, tile=(64, 64), overlap=8, upscale_am output = torch.empty([samples.shape[0], out_channels] + mult_list_upscale(samples.shape[2:]), device=output_device) - for b in range(samples.shape[0]): - s = samples[b:b+1] + term_pbar = None + try: + for b in range(samples.shape[0]): + s = samples[b:b+1] - # handle entire input fitting in a single tile - if all(s.shape[d+2] <= tile[d] for d in range(dims)): - output[b:b+1] = function(s).to(output_device) - if pbar is not None: - pbar.update(1) - continue + # handle entire input fitting in a single tile + if all(s.shape[d+2] <= tile[d] for d in range(dims)): + if term_pbar_desc and term_pbar is None: + term_pbar = tqdm(desc=term_pbar_desc, total=samples.shape[0]) + output[b:b+1] = function(s).to(output_device) + if pbar is not None: + pbar.update(1) + if term_pbar is not None: + term_pbar.update(1) + continue - out = output[b:b+1].zero_() - out_div = torch.zeros([s.shape[0], 1] + mult_list_upscale(s.shape[2:]), device=output_device) + out = output[b:b+1].zero_() + out_div = torch.zeros([s.shape[0], 1] + mult_list_upscale(s.shape[2:]), device=output_device) - positions = [range(0, s.shape[d+2] - overlap[d], tile[d] - overlap[d]) if s.shape[d+2] > tile[d] else [0] for d in range(dims)] + positions = [range(0, s.shape[d+2] - overlap[d], tile[d] - overlap[d]) if s.shape[d+2] > tile[d] else [0] for d in range(dims)] - for it in itertools.product(*positions): - s_in = s - upscaled = [] + if term_pbar_desc and term_pbar is None: + term_pbar = tqdm(desc=term_pbar_desc, total=samples.shape[0] * sum(1 for e in itertools.product(*positions))) - for d in range(dims): - pos = max(0, min(s.shape[d + 2] - overlap[d], it[d])) - l = min(tile[d], s.shape[d + 2] - pos) - s_in = s_in.narrow(d + 2, pos, l) - upscaled.append(round(get_pos(d, pos))) + for it in itertools.product(*positions): + s_in = s + upscaled = [] - ps = function(s_in).to(output_device) - mask = torch.ones([1, 1] + list(ps.shape[2:]), device=output_device) + for d in range(dims): + pos = max(0, min(s.shape[d + 2] - overlap[d], it[d])) + l = min(tile[d], s.shape[d + 2] - pos) + s_in = s_in.narrow(d + 2, pos, l) + upscaled.append(round(get_pos(d, pos))) - for d in range(2, dims + 2): - feather = round(get_scale(d - 2, overlap[d - 2])) - if feather >= mask.shape[d]: - continue - for t in range(feather): - a = (t + 1) / feather - mask.narrow(d, t, 1).mul_(a) - mask.narrow(d, mask.shape[d] - 1 - t, 1).mul_(a) + ps = function(s_in).to(output_device) + mask = torch.ones([1, 1] + list(ps.shape[2:]), device=output_device) - o = out - o_d = out_div - ps_view = ps - mask_view = mask - for d in range(dims): - l = min(ps_view.shape[d + 2], o.shape[d + 2] - upscaled[d]) - o = o.narrow(d + 2, upscaled[d], l) - o_d = o_d.narrow(d + 2, upscaled[d], l) - if l < ps_view.shape[d + 2]: - ps_view = ps_view.narrow(d + 2, 0, l) - mask_view = mask_view.narrow(d + 2, 0, l) + for d in range(2, dims + 2): + feather = round(get_scale(d - 2, overlap[d - 2])) + if feather >= mask.shape[d]: + continue + for t in range(feather): + a = (t + 1) / feather + mask.narrow(d, t, 1).mul_(a) + mask.narrow(d, mask.shape[d] - 1 - t, 1).mul_(a) - o.add_(ps_view * mask_view) - o_d.add_(mask_view) + o = out + o_d = out_div + ps_view = ps + mask_view = mask + for d in range(dims): + l = min(ps_view.shape[d + 2], o.shape[d + 2] - upscaled[d]) + o = o.narrow(d + 2, upscaled[d], l) + o_d = o_d.narrow(d + 2, upscaled[d], l) + if l < ps_view.shape[d + 2]: + ps_view = ps_view.narrow(d + 2, 0, l) + mask_view = mask_view.narrow(d + 2, 0, l) - if pbar is not None: - pbar.update(1) + o.add_(ps_view * mask_view) + o_d.add_(mask_view) - out.div_(out_div) + if pbar is not None: + pbar.update(1) + if term_pbar: + term_pbar.update(1) + + out.div_(out_div) + finally: + if term_pbar: + term_pbar.close() return output -def tiled_scale(samples, function, tile_x=64, tile_y=64, overlap = 8, upscale_amount = 4, out_channels = 3, output_device="cpu", pbar = None): - return tiled_scale_multidim(samples, function, (tile_y, tile_x), overlap=overlap, upscale_amount=upscale_amount, out_channels=out_channels, output_device=output_device, pbar=pbar) +def tiled_scale(samples, function, tile_x=64, tile_y=64, overlap = 8, upscale_amount = 4, out_channels = 3, output_device="cpu", pbar = None, term_pbar_desc=None): + return tiled_scale_multidim(samples, function, (tile_y, tile_x), overlap=overlap, upscale_amount=upscale_amount, out_channels=out_channels, output_device=output_device, pbar=pbar, term_pbar_desc=term_pbar_desc) def model_trange(*args, **kwargs): if not comfy.memory_management.aimdo_enabled: @@ -1266,7 +1281,7 @@ PROGRESS_THROTTLE_MIN_INTERVAL = 0.1 # 100ms minimum between updates PROGRESS_THROTTLE_MIN_PERCENT = 0.5 # 0.5% minimum progress change class ProgressBar: - def __init__(self, total, node_id=None): + def __init__(self, total, node_id=None, term_desc=None): global PROGRESS_BAR_HOOK self.total = total self.current = 0 @@ -1274,13 +1289,24 @@ class ProgressBar: self.node_id = node_id self._last_update_time = 0.0 self._last_sent_value = -1 + self.term_pbar = None + if term_desc: + self.term_pbar = tqdm(total=total, desc=term_desc) def update_absolute(self, value, total=None, preview=None): if total is not None: self.total = total if value > self.total: value = self.total + inc = value - self.current self.current = value + + if self.term_pbar and inc > 0: + self.term_pbar.total = self.total + self.term_pbar.update(inc) + if value >= self.total: + self.term_pbar.close() + if self.hook is not None: current_time = time.perf_counter() is_first = (self._last_sent_value < 0) diff --git a/comfy_extras/nodes_upscale_model.py b/comfy_extras/nodes_upscale_model.py index d1e0401b9..58e5e32e9 100644 --- a/comfy_extras/nodes_upscale_model.py +++ b/comfy_extras/nodes_upscale_model.py @@ -86,7 +86,7 @@ class ImageUpscaleWithModel(io.ComfyNode): try: steps = in_img.shape[0] * comfy.utils.get_tiled_scale_steps(in_img.shape[3], in_img.shape[2], tile_x=tile, tile_y=tile, overlap=overlap) pbar = comfy.utils.ProgressBar(steps) - s = comfy.utils.tiled_scale(in_img, lambda a: upscale_model(a.float()), tile_x=tile, tile_y=tile, overlap=overlap, upscale_amount=upscale_model.scale, pbar=pbar, output_device=output_device) + s = comfy.utils.tiled_scale(in_img, lambda a: upscale_model(a.float()), tile_x=tile, tile_y=tile, overlap=overlap, upscale_amount=upscale_model.scale, pbar=pbar, output_device=output_device, term_pbar_desc="Upscale") oom = False except Exception as e: model_management.raise_non_oom(e) diff --git a/nodes.py b/nodes.py index 1a3dd3f48..eb1d10195 100644 --- a/nodes.py +++ b/nodes.py @@ -335,7 +335,7 @@ class VAEDecode: if latent.is_nested: latent = latent.unbind()[0] - images = vae.decode(latent) + images = vae.decode(latent, term_pbar_desc="VaeDecode") if len(images.shape) == 5: #Combine batches images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1]) return (images, )