Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 59 additions & 45 deletions comfy/ldm/minimax/vae.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
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
Expand Down Expand Up @@ -321,6 +322,8 @@ def forward(self, x):
# Full VAE

class MiniMaxH3VideoVAE(nn.Module):
comfy_has_chunked_io = True

def __init__(
self,
in_channels=3,
Expand Down Expand Up @@ -389,6 +392,23 @@ def _encode_moments(self, x):
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)
Expand Down Expand Up @@ -521,18 +541,15 @@ def tiled_decode(self, z):

# 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:
Expand Down Expand Up @@ -577,43 +594,42 @@ def _decode_temporal_frame_plan(self, z_len, num_chunks, pad_tokens):
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_(
Expand Down Expand Up @@ -653,18 +669,18 @@ def write_part(part):
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]

Expand All @@ -679,18 +695,16 @@ def encode_tiled(self, x, **kwargs):
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)
9 changes: 9 additions & 0 deletions comfy/sd.py
Original file line number Diff line number Diff line change
Expand Up @@ -955,13 +955,21 @@ def estimate_memory(shape, dtype, num_layers = 16, kv_cache_multiplier = 2):
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)
Expand Down Expand Up @@ -1198,6 +1206,7 @@ def decode(self, samples_in, vae_options={}):
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:
Expand Down
Loading