From 6233790c6dff26bf35113d46d6d3367b7041b1d8 Mon Sep 17 00:00:00 2001 From: chelsealong Date: Tue, 11 Aug 2026 11:00:23 +0800 Subject: [PATCH] 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. --- nodes.py | 6 ++++- .../test_vae_decode_tiled_nested.py | 27 +++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) create mode 100644 tests-unit/comfy_test/test_vae_decode_tiled_nested.py diff --git a/nodes.py b/nodes.py index 432f04d89..a7f91720f 100644 --- a/nodes.py +++ b/nodes.py @@ -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, ) diff --git a/tests-unit/comfy_test/test_vae_decode_tiled_nested.py b/tests-unit/comfy_test/test_vae_decode_tiled_nested.py new file mode 100644 index 000000000..7c4b345b9 --- /dev/null +++ b/tests-unit/comfy_test/test_vae_decode_tiled_nested.py @@ -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