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:
parent
34744cd29e
commit
6233790c6d
6
nodes.py
6
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, )
|
||||
|
|
|
|||
|
|
@ -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
|
||||
Loading…
Reference in New Issue