Fix VAEDecodeTiled crash on NestedTensor latents (MiniMax H3) (#15477)

VAEDecode unwraps a NestedTensor latent (video/audio pair) to its
video component before calling vae.decode(). VAEDecodeTiled skipped
this unwrap and passed the NestedTensor straight into
vae.decode_tiled(), which fails deep in the MiniMax H3 video VAE when
a real tensor's .to() is called with the NestedTensor as an argument.

Fixes #15468.
This commit is contained in:
chelsealong 2026-08-11 11:00:23 +08:00 committed by GitHub
parent 34744cd29e
commit 6233790c6d
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 32 additions and 1 deletions

View File

@ -364,8 +364,12 @@ class VAEDecodeTiled:
temporal_size = None
temporal_overlap = None
latent = samples["samples"]
if latent.is_nested:
latent = latent.unbind()[0]
compression = vae.spacial_compression_decode()
images = vae.decode_tiled(samples["samples"], tile_x=tile_size // compression, tile_y=tile_size // compression, overlap=overlap // compression, tile_t=temporal_size, overlap_t=temporal_overlap)
images = vae.decode_tiled(latent, tile_x=tile_size // compression, tile_y=tile_size // compression, overlap=overlap // compression, tile_t=temporal_size, overlap_t=temporal_overlap)
if len(images.shape) == 5: #Combine batches
images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1])
return (images, )

View File

@ -0,0 +1,27 @@
from unittest.mock import MagicMock
import torch
from comfy.cli_args import args as cli_args
if not torch.cuda.is_available():
cli_args.cpu = True
import comfy.nested_tensor # noqa: E402
import nodes # noqa: E402
def test_vae_decode_tiled_unwraps_nested_tensor():
video = torch.zeros(1, 4, 2, 8, 8)
audio = torch.zeros(1, 2, 2, 40)
samples = {"samples": comfy.nested_tensor.NestedTensor((video, audio))}
vae = MagicMock()
vae.temporal_compression_decode.return_value = None
vae.spacial_compression_decode.return_value = 8
vae.decode_tiled.return_value = torch.zeros(1, 3, 2, 8, 8)
nodes.VAEDecodeTiled().decode(vae, samples, tile_size=512)
decoded_arg = vae.decode_tiled.call_args[0][0]
assert decoded_arg is video