Merge pull request #1546 from wzxysf/main
Enable vae tiling with end frame
This commit is contained in:
@@ -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