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