diff --git a/comfy/ldm/minimax/vae.py b/comfy/ldm/minimax/vae.py index 65d06f3e9..e1c146ec4 100644 --- a/comfy/ldm/minimax/vae.py +++ b/comfy/ldm/minimax/vae.py @@ -6,6 +6,7 @@ import torch import torch.nn as nn import torch.nn.functional as F +import comfy.model_management import comfy.ops import comfy.quant_ops import comfy.rmsnorm @@ -321,6 +322,8 @@ class ViT3DDecoder(nn.Module): # Full VAE class MiniMaxH3VideoVAE(nn.Module): + comfy_has_chunked_io = True + def __init__( self, in_channels=3, @@ -389,6 +392,23 @@ class MiniMaxH3VideoVAE(nn.Module): def _decode_pixels(self, z): return self.decoder(self.post_quant_conv(z)) + def _normalize_pixels(self, x): + return x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x)) + + def _finalize_pixels(self, part): + # raw decoder output -> float32 pixels in [0, 1] (the VAE wrapper's process_output is identity) + part = part * self.pixel_std.to(device=part.device, dtype=torch.float32) + return part.add_(self.pixel_mean.to(device=part.device, dtype=torch.float32)).clamp_(0.0, 1.0) + + def decode_output_shape(self, input_shape): + b, c, t, h, w = input_shape + if t == 1: + frames = 1 + else: + pad_tokens, num_chunks = self._decode_temporal_chunks(t) + frames = self._decode_temporal_frame_plan(t + pad_tokens, num_chunks, pad_tokens) + return (b, self.decoder.out_channels, frames, h * self.vae_ratio, w * self.vae_ratio) + def _adaptive_encode(self, x): if self.tiling: return self.tiled_encode(x) @@ -521,18 +541,15 @@ class MiniMaxH3VideoVAE(nn.Module): # temporal chunking - def encode_temporal(self, x): - if x.shape[2] % self.clip_length != 0: - pad_size = (-x.shape[2]) % self.clip_length - pad_frames = x[:, :, -1:].repeat(1, 1, pad_size, 1, 1) - x = torch.cat([x, pad_frames], dim=2) - - num_chunks = x.shape[2] // self.clip_length - + def encode_temporal(self, x, device): + # chunked input io: x may live on the CPU, clips move to the device as they encode z_list = [] - for i in range(num_chunks): - clip_x = x[:, :, i * self.clip_length:(i + 1) * self.clip_length, :, :] - z_list.append(self._adaptive_encode(clip_x)) + for i in range(math.ceil(x.shape[2] / self.clip_length)): + clip_x = x[:, :, i * self.clip_length:(i + 1) * self.clip_length, :, :].to(device) + if clip_x.shape[2] < self.clip_length: + pad_frames = clip_x[:, :, -1:].repeat(1, 1, self.clip_length - clip_x.shape[2], 1, 1) + clip_x = torch.cat([clip_x, pad_frames], dim=2) + z_list.append(self._adaptive_encode(self._normalize_pixels(clip_x))) z = torch.cat(z_list, dim=2) if self.token_drop > 0: @@ -577,43 +594,42 @@ class MiniMaxH3VideoVAE(nn.Module): total_frames += final_overlap_frames return total_frames - self._decode_temporal_pad_frames(z_len, pad_tokens) - def decode_temporal(self, z): - chunk_dec = self.tokens_chunk_size * self.vae_ratio_t - split_count = int(self.token_drop > 0) + 1 - - pseudo_total_tokens = z.shape[2] + self.token_drop - - pad_tokens = 0 - remainder = pseudo_total_tokens % self.tokens_chunk_size - if remainder != 0: - pad_tokens = self.tokens_chunk_size - remainder - pseudo_total_tokens += pad_tokens + def _decode_temporal_chunks(self, z_len): + pseudo_total_tokens = z_len + self.token_drop + pad_tokens = (-pseudo_total_tokens) % self.tokens_chunk_size + pseudo_total_tokens += pad_tokens num_chunks = pseudo_total_tokens // self.tokens_chunk_size - int(self.token_drop > 0) if num_chunks < 1: # too few tokens for one chunk (e.g. T_lat == 2): pad one extra chunk pad_tokens += self.tokens_chunk_size num_chunks += 1 + return pad_tokens, num_chunks + def decode_temporal(self, z, output_buffer=None): + chunk_dec = self.tokens_chunk_size * self.vae_ratio_t + split_count = int(self.token_drop > 0) + 1 + + if output_buffer is None: + # finalized chunks stream out of VRAM so the full video never sits on the GPU + output_buffer = torch.empty(self.decode_output_shape(z.shape), dtype=torch.float32, + device=comfy.model_management.intermediate_device()) + + pad_tokens, num_chunks = self._decode_temporal_chunks(z.shape[2]) if pad_tokens > 0: pad_z = z[:, :, -1:, :, :].repeat(1, 1, pad_tokens, 1, 1) z = torch.cat([z, pad_z], dim=2) - output_frames = self._decode_temporal_frame_plan(z.shape[2], num_chunks, pad_tokens) - - dec = None + dec = output_buffer dec_overlap = None write_pos = 0 def write_part(part): - nonlocal dec, write_pos + nonlocal write_pos part_frames = part.shape[2] if part_frames <= 0: return - if dec is None: - out_shape = list(part.shape) - out_shape[2] = output_frames - dec = torch.empty(out_shape, dtype=part.dtype, device=part.device) + part = self._finalize_pixels(part) copy_frames = min(part_frames, max(0, dec.shape[2] - write_pos)) if copy_frames > 0: dec[:, :, write_pos:write_pos + copy_frames, :, :].copy_( @@ -653,18 +669,18 @@ class MiniMaxH3VideoVAE(nn.Module): return dec - def encode(self, x): + def encode(self, x, device=None): # x: [B, 3, T, H, W] in [-1, 1] -> normalized latents [B, 24, T_lat, H/16, W/16] if x.ndim == 4: x = x.unsqueeze(2) - - x = x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x)) + if device is None: + device = x.device if x.shape[2] == 1: - moments = self._adaptive_encode(x) + moments = self._adaptive_encode(self._normalize_pixels(x.to(device))) moments = moments[:, :, -1:, :, :] else: - moments = self.encode_temporal(x) + moments = self.encode_temporal(x, device) mean = torch.chunk(moments.float(), 2, dim=1)[0] @@ -679,18 +695,16 @@ class MiniMaxH3VideoVAE(nn.Module): def decode_tiled(self, z, **kwargs): return self.decode(z) - def decode(self, z): - # z: [B, 24, T_lat, H_lat, W_lat] normalized latents -> pixels [B, 3, T, H, W] in [-1, 1] + def decode(self, z, output_buffer=None): + # z: [B, 24, T_lat, H_lat, W_lat] normalized latents -> float32 pixels [B, 3, T, H, W] in [0, 1] latents_mean = self.latents_mean.view(1, -1, 1, 1, 1).to(z) latents_std = self.latents_std.view(1, -1, 1, 1, 1).to(z) z = z * latents_std + latents_mean if z.shape[2] == 1: - dec = self._adaptive_decode(z) - dec = dec[:, :, -1:, :, :] - else: - dec = self.decode_temporal(z) - - dec = dec.float() - dec.mul_(self.pixel_std.to(dec)).add_(self.pixel_mean.to(dec)).clamp_(0.0, 1.0).mul_(2.0).sub_(1.0) - return dec + dec = self._finalize_pixels(self._adaptive_decode(z)[:, :, -1:, :, :]) + if output_buffer is None: + return dec + output_buffer.copy_(dec) + return output_buffer + return self.decode_temporal(z, output_buffer) diff --git a/comfy/sd.py b/comfy/sd.py index 9ccd561bc..5fed4ca9a 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -955,13 +955,21 @@ class VAE: self.working_dtypes = [torch.float16, torch.float32] # the model tiles internally (256px spatial, 17-frame temporal chunks) self.handles_tiling = True + # decode finalizes straight to [0, 1] while streaming chunks out + self.process_output = lambda image: image + # one decoded temporal chunk (with overlap) is all that ever sits in VRAM + chunk_frames = (self.first_stage_model.tokens_chunk_size + self.first_stage_model.token_overlap) * self.first_stage_model.vae_ratio_t + def estimate_encode_memory(frames, height, width, dtype): fixed = 110_000_000 if frames == 1 else 1_300_000_000 elements_per_pixel = 7 if frames == 1 else 9.5 + # only one clip of the input video is ever resident on the GPU + frames = min(frames, self.first_stage_model.clip_length) return (elements_per_pixel * frames * height * width + fixed) * model_management.dtype_size(dtype) * 1.03 def estimate_decode_memory(frames, height, width, dtype): fixed = 110_000_000 if frames <= 22 else 270_000_000 + frames = min(frames, chunk_frames + 2) return (9.5 * frames * height * width + fixed) * model_management.dtype_size(dtype) * 1.03 self.memory_used_encode = lambda shape, dtype: estimate_encode_memory(shape[2], shape[3], shape[4], dtype) @@ -1198,6 +1206,7 @@ class VAE: do_tile = True if do_tile: + pixel_samples = None comfy.model_management.soft_empty_cache() dims = samples_in.ndim - 2 if dims == 1 or self.extra_1d_channel is not None: