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