File size: 10,488 Bytes
3afd6d6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 | """Loaders for tiny VAE decoders that ComfyUI core can't build yet.
TinyVAEDecoder handles flat TAESD-style decoders. Core's Decoder hardcodes a 64-wide stack
with 3 upsamples, so anything else fails to load — the 2D taeh3 is 96 wide with 4. The
architecture is recoverable from the checkpoint: keys are positional module indices, so
`N.conv.0.weight` is a Block, `N.weight` a conv, and the gaps are parameterless modules.
TAEHVDecoder handles the temporal taehv format, which core builds but can't size for H3.
"""
import logging
import os
import torch
import torch.nn as nn
import comfy.model_management
import comfy.utils
from comfy.taesd.taesd import Block, Clamp, conv
def build_tae_decoder(sd):
by_index = {}
for k, v in sd.items():
head, _, rest = k.partition(".")
if not head.isdigit():
raise ValueError(f"not a flat TAE decoder state dict (unexpected key '{k}')")
by_index.setdefault(int(head), {})[rest] = v
modules = []
for i in range(max(by_index) + 1):
entry = by_index.get(i)
if entry is None:
# index 0 is the input Clamp, 2 the ReLU after the input conv, the rest are upsamples
modules.append(Clamp() if i == 0 else nn.ReLU() if i == 2 else nn.Upsample(scale_factor=2))
elif "conv.0.weight" in entry:
w = entry["conv.0.weight"]
# only pass the kwarg when it's needed — older ComfyUI has no midblock-GN variant
if "pool.0.weight" in entry:
modules.append(Block(w.shape[1], w.shape[0], use_midblock_gn=True))
else:
modules.append(Block(w.shape[1], w.shape[0]))
elif "weight" in entry:
w = entry["weight"]
modules.append(conv(w.shape[1], w.shape[0], bias="bias" in entry))
else:
raise ValueError(f"unrecognized TAE decoder module at index {i}: {sorted(entry)}")
return nn.Sequential(*modules)
class TinyVAEDecoder:
"""Decode-only tiny VAE. decode() is float32 in [0, 1] like the TAE family; decode_video() is uint8."""
decodes_prefix = False # honours frame_indices exactly
def __init__(self, sd, device=None, dtype=None):
# keys may carry a "taesd_decoder."/"decoder." prefix; strip whatever is common
prefix = ""
first = next(iter(sd))
if not first.split(".")[0].isdigit():
prefix = first.split(".")[0] + "."
sd = {k[len(prefix):]: v for k, v in sd.items() if k.startswith(prefix)}
self.device = device if device is not None else comfy.model_management.vae_device()
self.dtype = dtype if dtype is not None else comfy.model_management.vae_dtype(
self.device, [torch.float16, torch.bfloat16])
self.model = build_tae_decoder(sd)
self.model.load_state_dict(sd)
self.model = _place(self.model, self.device, self.dtype)
self.latent_channels = self.model[1].weight.shape[1]
self.upscale_ratio = 2 ** sum(isinstance(m, nn.Upsample) for m in self.model)
def decode(self, latent):
"""[B, C, H, W] -> [B, 3, H*ratio, W*ratio], float32 on the input device."""
out = self.model(latent.to(device=self.device, dtype=self.dtype))
return out.to(device=latent.device, dtype=torch.float32)
def decode_video(self, latent, frame_indices=None):
"""[B, C, T, H, W] -> [T, 3, H*ratio, W*ratio] uint8 on the CPU. Decodes one frame at a time
in the model dtype and casts before the host copy, so the GPU never holds more than one frame;
frames land directly in the preallocated output instead of a second stack copy."""
x = latent[0]
indices = list(range(x.shape[1])) if frame_indices is None else list(frame_indices)
out = None
for i, t in enumerate(indices):
f = self.model(x[:, t].unsqueeze(0).to(device=self.device, dtype=self.dtype))[0]
f = f.clamp_(0, 1).mul_(255).to(torch.uint8)
if out is None:
out = torch.empty((len(indices),) + tuple(f.shape), dtype=torch.uint8)
out[i].copy_(f)
return out
def _place(model, device, dtype):
model = model.eval().to(device=device, dtype=dtype)
if torch.device(device).type == "cuda":
model.to(memory_format=torch.channels_last)
return model
def is_taehv_state_dict(sd):
return "decoder.1.weight" in sd and "decoder.22.bias" in sd
class TAEHVDecoder:
"""Temporal tiny VAE (madebyollin/taehv), decode only."""
decodes_prefix = True # memblock state chains forward, so partial requests decode the prefix
def __init__(self, sd, device=None, dtype=None):
from comfy.taesd.taehv import TAEHV, conv
latent_channels = sd["decoder.1.weight"].shape[1]
# final conv emits image_channels * patch_size**2
patch_size = max(1, int(round((sd["decoder.22.bias"].shape[0] / 3) ** 0.5)))
model = TAEHV(latent_channels=latent_channels)
# core derives patch_size from the channel count and has no entry for H3's 24;
# those two convs are the only parts of the model that depend on it
if model.patch_size != patch_size:
model.patch_size = patch_size
model.encoder[0] = conv(3 * patch_size ** 2, model.encoder[0].out_channels)
model.decoder[-1] = conv(model.decoder[-1].in_channels, 3 * patch_size ** 2)
model.load_state_dict(sd)
del model.encoder # decode only, so keep it off the device entirely
self.device = device if device is not None else comfy.model_management.vae_device()
self.dtype = dtype if dtype is not None else comfy.model_management.vae_dtype(
self.device, [torch.float16, torch.bfloat16])
self.model = _place(model, self.device, self.dtype)
self.latent_channels = latent_channels
self.upscale_ratio = patch_size * 2 ** sum(isinstance(m, nn.Upsample) for m in model.decoder)
self.is_h3 = latent_channels == 24 and patch_size == 2
def _decode(self, latent):
# [B, C, T, H, W] -> [B, 3, T*t_upscale - trim, H*ratio, W*ratio] on the intermediate device
# (CPU unless --gpu-only); frames stream off the GPU one at a time inside the memblock loop
return self.model.decode(latent.to(device=self.device, dtype=self.dtype))
def decode(self, latent):
"""[B, C, H, W] -> [B, 3, H*ratio, W*ratio], decoded as a single frame."""
return self._decode(latent.unsqueeze(2))[:, :, 0]
def _decode_h3_full(self, latent):
"""Whole-clip decode in H3's own chunking, so the frame count comes out exact."""
# H3's VAE codes 17 pixel frames per 5 latent tokens: trim each chunk's prefix
# rather than once globally, then drop the encoder's 3-token tail pad
from comfy.taesd.taehv import apply_model_with_memblocks
m = self.model
x = m.process_in(latent.to(device=self.device, dtype=self.dtype)).movedim(2, 1)
x = apply_model_with_memblocks(m.decoder, x, m.parallel, False,
output_device=comfy.model_management.intermediate_device(),
patch_size=m.patch_size, decode=True)
chunk = 5 * m.t_upscale
x = torch.nn.functional.pad(x, (0, 0, 0, 0, 0, 0, 0, -x.shape[1] % chunk))
x = x.unflatten(1, (-1, chunk))[:, :, m.frames_to_trim:].flatten(1, 2)
x = x[:, :-3 * m.t_upscale]
return x.movedim(2, 1)
def decode_video(self, latent, frame_indices=None):
"""[B, C, T, H, W] -> [n, 3, H*ratio, W*ratio] on the intermediate device, no repack."""
t_total = latent.shape[2]
n = t_total if frame_indices is None else max(1, min(len(frame_indices), t_total))
if n == t_total:
# whole clip asked for, so decode it properly — this returns every pixel frame,
# which is more than the latent frame count (17 per 5 tokens for H3)
if self.is_h3:
out = self._decode_h3_full(latent[:1])
if out.shape[2] > 0:
return out[0].movedim(0, 1)
return self._decode(latent[:1])[0].movedim(0, 1)
# MemBlock state chains forward, so frames can't be sampled across the clip without
# decoding everything before them — take a prefix to keep the per-step cost bounded
out = self._decode(latent[:1, :, :n])[0].movedim(0, 1)
if out.shape[0] > n:
out = out[torch.linspace(0, out.shape[0] - 1, n).round().long()]
return out
# Tiny VAEs shipped with the node pack live in <pack>/models/vae_approx,
# the same pack-local convention as the Moxie Multimedia Loader's models/.
PACK_VAE_APPROX_DIR = os.path.join(
os.path.dirname(os.path.abspath(__file__)), "models", "vae_approx"
)
TINY_VAE_EXTENSIONS = (".safetensors", ".sft", ".ckpt", ".pt", ".pth", ".bin")
def pack_vae_approx_files():
"""Tiny VAE files in the pack's models/vae_approx folder, as relative paths (sorted)."""
results = []
if os.path.isdir(PACK_VAE_APPROX_DIR):
for root, _dirs, files in os.walk(PACK_VAE_APPROX_DIR):
for filename in files:
if filename.lower().endswith(TINY_VAE_EXTENSIONS):
rel = os.path.relpath(os.path.join(root, filename), PACK_VAE_APPROX_DIR)
results.append(rel.replace("\\", "/"))
results.sort()
return results
def load_tiny_vae_decoder(name, device=None, dtype=None):
"""Load by vae_approx filename, searching the node pack's models/vae_approx
folder first and ComfyUI's standard models/vae_approx second.
Returns None (and logs) if it can't be used."""
import folder_paths
pack_path = os.path.join(PACK_VAE_APPROX_DIR, name)
if os.path.isfile(pack_path):
path = pack_path
else:
path = folder_paths.get_full_path("vae_approx", name)
if path is None:
logging.warning(
f"[Moxie TinyVAE] '{name}' not found in the pack's models/vae_approx "
"or ComfyUI's models/vae_approx."
)
return None
try:
sd = comfy.utils.load_torch_file(path, safe_load=True)
if is_taehv_state_dict(sd):
return TAEHVDecoder(sd, device=device, dtype=dtype)
return TinyVAEDecoder(sd, device=device, dtype=dtype)
except Exception as e:
logging.warning(f"[KJ TinyVAE] Could not load '{name}': {e}")
return None
|