Merge branch 'main' of https://github.com/kijai/ComfyUI-WanVideoWrapper
This commit is contained in:
+13
-6
@@ -75,14 +75,21 @@ class VAE(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
if data_dim == 80:
|
||||
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_80D, dtype=torch.float32))
|
||||
self.data_std = nn.Buffer(torch.tensor(DATA_STD_80D, dtype=torch.float32))
|
||||
data_mean = torch.tensor(DATA_MEAN_80D, dtype=torch.float32)
|
||||
data_std = torch.tensor(DATA_STD_80D, dtype=torch.float32)
|
||||
elif data_dim == 128:
|
||||
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_128D, dtype=torch.float32))
|
||||
self.data_std = nn.Buffer(torch.tensor(DATA_STD_128D, dtype=torch.float32))
|
||||
data_mean = torch.tensor(DATA_MEAN_128D, dtype=torch.float32)
|
||||
data_std = torch.tensor(DATA_STD_128D, dtype=torch.float32)
|
||||
else:
|
||||
raise ValueError(f"Unsupported data_dim={data_dim}, expected 80 or 128")
|
||||
|
||||
self.data_mean = self.data_mean.view(1, -1, 1)
|
||||
self.data_std = self.data_std.view(1, -1, 1)
|
||||
# match old shape: (1, channels, 1)
|
||||
data_mean = data_mean.view(1, -1, 1)
|
||||
data_std = data_std.view(1, -1, 1)
|
||||
|
||||
# register as buffers so they move with .to(device) / .cuda()
|
||||
self.register_buffer("data_mean", data_mean)
|
||||
self.register_buffer("data_std", data_std)
|
||||
|
||||
self.encoder = Encoder1D(
|
||||
dim=hidden_dim,
|
||||
|
||||
@@ -2204,8 +2204,6 @@ class WanVideoDecode:
|
||||
images = images.permute(1, 2, 3, 0).cpu().float()
|
||||
return (images,)
|
||||
else:
|
||||
if end_image is not None:
|
||||
enable_vae_tiling = False
|
||||
images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0]
|
||||
|
||||
|
||||
|
||||
@@ -1231,7 +1231,7 @@ class WanVideoVAE(nn.Module):
|
||||
return mask
|
||||
|
||||
|
||||
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, pbar=True):
|
||||
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, end_=False, pbar=True):
|
||||
_, _, T, H, W = hidden_states.shape
|
||||
size_h, size_w = tile_size
|
||||
stride_h, stride_w = tile_stride
|
||||
@@ -1248,14 +1248,20 @@ class WanVideoVAE(nn.Module):
|
||||
data_device = "cpu"
|
||||
computation_device = device
|
||||
|
||||
out_T = T * 4 - 3
|
||||
weight = torch.zeros((1, 1, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
||||
values = torch.zeros((1, 3, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
||||
weight, values = None, None
|
||||
if pbar:
|
||||
pbar = ProgressBar(len(tasks))
|
||||
for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
|
||||
hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device)
|
||||
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
|
||||
if end_:
|
||||
hidden_states_batch = self.model.decode_2(hidden_states_batch).to(data_device)
|
||||
else:
|
||||
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
|
||||
|
||||
if weight is None:
|
||||
weight = torch.zeros((1, 1, hidden_states_batch.shape[2], H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
||||
if values is None:
|
||||
values = torch.zeros((1, 3, hidden_states_batch.shape[2], H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
||||
|
||||
mask = self.build_mask(
|
||||
hidden_states_batch,
|
||||
@@ -1399,7 +1405,7 @@ class WanVideoVAE(nn.Module):
|
||||
for hidden_state in hidden_states:
|
||||
hidden_state = hidden_state.unsqueeze(0)
|
||||
if tiled:
|
||||
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, pbar=pbar)
|
||||
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, end_=end_, pbar=pbar)
|
||||
else:
|
||||
if end_:
|
||||
video = self.double_decode(hidden_state, device)
|
||||
|
||||
Reference in New Issue
Block a user