mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-03 17:17:15 +08:00
fixed tensor shape errors
This commit is contained in:
parent
596c342668
commit
aa33d7978e
@ -200,6 +200,8 @@ class AutoencodingEngineLegacy(AutoencodingEngine):
|
|||||||
return z
|
return z
|
||||||
|
|
||||||
def decode(self, z: torch.Tensor, **decoder_kwargs) -> torch.Tensor:
|
def decode(self, z: torch.Tensor, **decoder_kwargs) -> torch.Tensor:
|
||||||
|
print('Decoding hunyuan latent. Received tensor:',z.shape)
|
||||||
|
|
||||||
if self.max_batch_size is None:
|
if self.max_batch_size is None:
|
||||||
dec = self.post_quant_conv(z)
|
dec = self.post_quant_conv(z)
|
||||||
dec = self.decoder(dec, **decoder_kwargs)
|
dec = self.decoder(dec, **decoder_kwargs)
|
||||||
|
|||||||
@ -409,6 +409,7 @@ class VAE:
|
|||||||
self.downscale_index_formula = (4, 8, 8)
|
self.downscale_index_formula = (4, 8, 8)
|
||||||
self.latent_dim = 3
|
self.latent_dim = 3
|
||||||
self.latent_channels = ddconfig['z_channels'] = sd["decoder.conv_in.conv.weight"].shape[1]
|
self.latent_channels = ddconfig['z_channels'] = sd["decoder.conv_in.conv.weight"].shape[1]
|
||||||
|
print('Loading Hunyuan VAE. Latent channels = ',self.latent_channels)
|
||||||
self.first_stage_model = AutoencoderKL(ddconfig=ddconfig, embed_dim=sd['post_quant_conv.weight'].shape[1])
|
self.first_stage_model = AutoencoderKL(ddconfig=ddconfig, embed_dim=sd['post_quant_conv.weight'].shape[1])
|
||||||
self.memory_used_decode = lambda shape, dtype: (1500 * shape[2] * shape[3] * shape[4] * (4 * 8 * 8)) * model_management.dtype_size(dtype)
|
self.memory_used_decode = lambda shape, dtype: (1500 * shape[2] * shape[3] * shape[4] * (4 * 8 * 8)) * model_management.dtype_size(dtype)
|
||||||
self.memory_used_encode = lambda shape, dtype: (900 * max(shape[2], 2) * shape[3] * shape[4]) * model_management.dtype_size(dtype)
|
self.memory_used_encode = lambda shape, dtype: (900 * max(shape[2], 2) * shape[3] * shape[4]) * model_management.dtype_size(dtype)
|
||||||
|
|||||||
@ -62,6 +62,7 @@ def apply_model_with_memblocks(model, x, parallel, show_progress_bar):
|
|||||||
"""
|
"""
|
||||||
assert x.ndim == 5, f"TAEHV operates on NTCHW tensors, but got {x.ndim}-dim tensor"
|
assert x.ndim == 5, f"TAEHV operates on NTCHW tensors, but got {x.ndim}-dim tensor"
|
||||||
N, T, C, H, W = x.shape
|
N, T, C, H, W = x.shape
|
||||||
|
print('Received tensor of shape:',x.shape)
|
||||||
if parallel:
|
if parallel:
|
||||||
x = x.reshape(N*T, C, H, W)
|
x = x.reshape(N*T, C, H, W)
|
||||||
# parallel over input timesteps, iterate over blocks
|
# parallel over input timesteps, iterate over blocks
|
||||||
@ -71,8 +72,10 @@ def apply_model_with_memblocks(model, x, parallel, show_progress_bar):
|
|||||||
T = NT // N
|
T = NT // N
|
||||||
_x = x.reshape(N, T, C, H, W)
|
_x = x.reshape(N, T, C, H, W)
|
||||||
mem = F.pad(_x, (0,0,0,0,0,0,1,0), value=0)[:,:T].reshape(x.shape)
|
mem = F.pad(_x, (0,0,0,0,0,0,1,0), value=0)[:,:T].reshape(x.shape)
|
||||||
|
print('Intermediate shape:',x.shape)
|
||||||
x = b(x, mem)
|
x = b(x, mem)
|
||||||
else:
|
else:
|
||||||
|
print('Intermediate shape:',x.shape)
|
||||||
x = b(x)
|
x = b(x)
|
||||||
NT, C, H, W = x.shape
|
NT, C, H, W = x.shape
|
||||||
T = NT // N
|
T = NT // N
|
||||||
@ -89,6 +92,7 @@ def apply_model_with_memblocks(model, x, parallel, show_progress_bar):
|
|||||||
mem = [None] * len(model)
|
mem = [None] * len(model)
|
||||||
while work_queue:
|
while work_queue:
|
||||||
xt, i = work_queue.pop(0)
|
xt, i = work_queue.pop(0)
|
||||||
|
print('Intermediate shape:', xt.shape)
|
||||||
if i == len(model):
|
if i == len(model):
|
||||||
# reached end of the graph, append result to output list
|
# reached end of the graph, append result to output list
|
||||||
out.append(xt)
|
out.append(xt)
|
||||||
@ -150,7 +154,7 @@ class TAEHV(nn.Module):
|
|||||||
decoder_space_upscale: whether spatial upsampling is enabled for each block. upsampling can be disabled for a cheaper preview.
|
decoder_space_upscale: whether spatial upsampling is enabled for each block. upsampling can be disabled for a cheaper preview.
|
||||||
"""
|
"""
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.taehv_encoder = nn.Sequential(
|
self.encoder = nn.Sequential(
|
||||||
conv(TAEHV.image_channels, 64), nn.ReLU(inplace=True),
|
conv(TAEHV.image_channels, 64), nn.ReLU(inplace=True),
|
||||||
TPool(64, 2), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64),
|
TPool(64, 2), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64),
|
||||||
TPool(64, 2), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64),
|
TPool(64, 2), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64),
|
||||||
@ -159,7 +163,7 @@ class TAEHV(nn.Module):
|
|||||||
)
|
)
|
||||||
n_f = [256, 128, 64, 64]
|
n_f = [256, 128, 64, 64]
|
||||||
self.frames_to_trim = 2**sum(decoder_time_upscale) - 1
|
self.frames_to_trim = 2**sum(decoder_time_upscale) - 1
|
||||||
self.taehv_decoder = nn.Sequential(
|
self.decoder = nn.Sequential(
|
||||||
Clamp(), conv(TAEHV.latent_channels, n_f[0]), nn.ReLU(inplace=True),
|
Clamp(), conv(TAEHV.latent_channels, n_f[0]), nn.ReLU(inplace=True),
|
||||||
MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), nn.Upsample(scale_factor=2 if decoder_space_upscale[0] else 1), TGrow(n_f[0], 1), conv(n_f[0], n_f[1], bias=False),
|
MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), nn.Upsample(scale_factor=2 if decoder_space_upscale[0] else 1), TGrow(n_f[0], 1), conv(n_f[0], n_f[1], bias=False),
|
||||||
MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), nn.Upsample(scale_factor=2 if decoder_space_upscale[1] else 1), TGrow(n_f[1], 2 if decoder_time_upscale[0] else 1), conv(n_f[1], n_f[2], bias=False),
|
MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), nn.Upsample(scale_factor=2 if decoder_space_upscale[1] else 1), TGrow(n_f[1], 2 if decoder_time_upscale[0] else 1), conv(n_f[1], n_f[2], bias=False),
|
||||||
@ -167,10 +171,10 @@ class TAEHV(nn.Module):
|
|||||||
nn.ReLU(inplace=True), conv(n_f[3], TAEHV.image_channels),
|
nn.ReLU(inplace=True), conv(n_f[3], TAEHV.image_channels),
|
||||||
)
|
)
|
||||||
if checkpoint_path is not None:
|
if checkpoint_path is not None:
|
||||||
self.load_state_dict((checkpoint_path, safe_load=True))
|
self.load_state_dict(comfy.utils.load_torch_file(checkpoint_path, safe_load=True))
|
||||||
|
|
||||||
|
|
||||||
def load_state_dict(state_dict, strict=True):
|
def load_state_dict(self, state_dict, strict=True):
|
||||||
return super().load_state_dict(self.patch_tgrow_layers(state_dict), strict=strict)
|
return super().load_state_dict(self.patch_tgrow_layers(state_dict), strict=strict)
|
||||||
|
|
||||||
|
|
||||||
@ -217,7 +221,7 @@ class TAEHV(nn.Module):
|
|||||||
sd[key] = sd[key][-new_sd[key].shape[0]:]
|
sd[key] = sd[key][-new_sd[key].shape[0]:]
|
||||||
return sd
|
return sd
|
||||||
|
|
||||||
def encode_video(self, x, parallel=True, show_progress_bar=False):
|
def encode_video(self, x, parallel=False, show_progress_bar=False):
|
||||||
"""Encode a sequence of frames.
|
"""Encode a sequence of frames.
|
||||||
Args:
|
Args:
|
||||||
x: input NTCHW RGB (C=3) tensor with values in [0, 1].
|
x: input NTCHW RGB (C=3) tensor with values in [0, 1].
|
||||||
@ -228,17 +232,26 @@ class TAEHV(nn.Module):
|
|||||||
"""
|
"""
|
||||||
return apply_model_with_memblocks(self.encoder, x, parallel, show_progress_bar)
|
return apply_model_with_memblocks(self.encoder, x, parallel, show_progress_bar)
|
||||||
|
|
||||||
def decode_video(self, x, parallel=True, show_progress_bar=False):
|
def decode_video(self, x, parallel=False, show_progress_bar=False):
|
||||||
"""Decode a sequence of frames.
|
"""Decode a sequence of frames.
|
||||||
Args:
|
Args:
|
||||||
x: input NTCHW latent (C=16) tensor with ~Gaussian values.
|
x: input NCTHW latent (C=16) tensor with ~Gaussian values.
|
||||||
parallel: if True, all frames will be processed at once.
|
parallel: if True, all frames will be processed at once.
|
||||||
(this is faster but may require more memory).
|
(this is faster but may require more memory).
|
||||||
if False, frames will be processed sequentially.
|
if False, frames will be processed sequentially.
|
||||||
Returns NTCHW RGB tensor with ~[0, 1] values.
|
Returns NTCHW RGB tensor with ~[0, 1] values.
|
||||||
"""
|
"""
|
||||||
|
#converting NCTHW to NTCHW
|
||||||
|
x = x.permute(0,2,1,3,4)
|
||||||
x = apply_model_with_memblocks(self.decoder, x, parallel, show_progress_bar)
|
x = apply_model_with_memblocks(self.decoder, x, parallel, show_progress_bar)
|
||||||
return x[:, self.frames_to_trim:]
|
|
||||||
|
|
||||||
|
x = x[:, self.frames_to_trim:] # trim the time dimension
|
||||||
|
|
||||||
|
#converting NTCHW to NCTHW
|
||||||
|
x = x.permute(0,2,1,3,4)
|
||||||
|
|
||||||
|
return x
|
||||||
|
|
||||||
def decode(self, x):
|
def decode(self, x):
|
||||||
"""Decode a single frame or batch of frames for preview."""
|
"""Decode a single frame or batch of frames for preview."""
|
||||||
@ -247,13 +260,17 @@ class TAEHV(nn.Module):
|
|||||||
x = x.unsqueeze(1)
|
x = x.unsqueeze(1)
|
||||||
|
|
||||||
# For preview, we'll just take the first frame after decoding
|
# For preview, we'll just take the first frame after decoding
|
||||||
decoded = self.decode_video(x, parallel=True, show_progress_bar=False)
|
decoded = self.decode_video(x, parallel=False, show_progress_bar=False)
|
||||||
|
print(decoded.shape)
|
||||||
|
|
||||||
|
#converting
|
||||||
|
|
||||||
|
return decoded
|
||||||
# Return single frame for preview
|
# Return single frame for preview
|
||||||
if decoded.shape[1] > 0:
|
# if decoded.shape[1] > 0:
|
||||||
return decoded[:, 0]
|
# return decoded[:, 0]
|
||||||
else:
|
# else:
|
||||||
return decoded.squeeze(1)
|
# return decoded.squeeze(1)
|
||||||
|
|
||||||
def encode(self, x):
|
def encode(self, x):
|
||||||
"""Encode a single frame or batch of frames."""
|
"""Encode a single frame or batch of frames."""
|
||||||
@ -261,7 +278,7 @@ class TAEHV(nn.Module):
|
|||||||
# Add temporal dimension for single frame
|
# Add temporal dimension for single frame
|
||||||
x = x.unsqueeze(1)
|
x = x.unsqueeze(1)
|
||||||
|
|
||||||
encoded = self.encode_video(x, parallel=True, show_progress_bar=False)
|
encoded = self.encode_video(x, parallel=False, show_progress_bar=False)
|
||||||
|
|
||||||
# Return single frame
|
# Return single frame
|
||||||
return encoded.squeeze(1)
|
return encoded.squeeze(1)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user