diff --git a/nodes.py b/nodes.py index ce27d2a..42ab4e1 100644 --- a/nodes.py +++ b/nodes.py @@ -1097,6 +1097,7 @@ class WanVideoContextOptions: "optional": { "image_cond_start_step": ("INT", {"default": 6, "min": 0, "max": 10000, "step": 1, "tooltip": "!EXPERIMENTAL! Start step of using previous window results as input instead of the init image"}), "image_cond_window_count": ("INT", {"default": 2, "min": 1, "max": 10000, "step": 1, "tooltip": "!EXPERIMENTAL! Number of image 'prompt windows'"}), + "vae": ("WANVAE",), } } @@ -1106,7 +1107,7 @@ class WanVideoContextOptions: CATEGORY = "WanVideoWrapper" DESCRIPTION = "Context options for WanVideo, allows splitting the video into context windows and attemps blending them for longer generations than the model and memory otherwise would allow." - def process(self, context_schedule, context_frames, context_stride, context_overlap, freenoise, verbose, image_cond_start_step=6, image_cond_window_count=2): + def process(self, context_schedule, context_frames, context_stride, context_overlap, freenoise, verbose, image_cond_start_step=6, image_cond_window_count=2, vae=None): context_options = { "context_schedule":context_schedule, "context_frames":context_frames, @@ -1115,7 +1116,8 @@ class WanVideoContextOptions: "freenoise":freenoise, "verbose":verbose, "image_cond_start_step": image_cond_start_step, - "image_cond_window_count": image_cond_window_count + "image_cond_window_count": image_cond_window_count, + "vae": vae, } return (context_options,) @@ -1307,6 +1309,9 @@ class WanVideoSampler: context_frames = (context_options["context_frames"] - 1) // 4 + 1 context_stride = context_options["context_stride"] // 4 context_overlap = context_options["context_overlap"] // 4 + context_vae = context_options.get("vae", None) + if context_vae is not None: + context_vae.to(device) self.window_tracker = WindowTracker(verbose=context_options["verbose"]) @@ -1725,11 +1730,20 @@ class WanVideoSampler: if idx >= context_options["image_cond_start_step"]: #strength = 0.5 #partial_image_cond *= strength + if context_vae is not None: + to_decode = self.previous_noise_pred_context[:,-1,:, :].unsqueeze(1).unsqueeze(0).to(context_vae.dtype) + #to_decode = to_decode.permute(0, 1, 3, 2) + print("to_decode.shape", to_decode.shape) + image = context_vae.decode(to_decode, device=device, tiled=False)[0] + print("decoded image.shape", image.shape) #torch.Size([3, 37, 832, 480]) + image = context_vae.encode(image.unsqueeze(0).to(context_vae.dtype), device=device, tiled=False) + print("encoded image.shape", image.shape) + #partial_img_emb[:, 0, :, :] = image[0][:,0,:,:] print("partial_img_emb.shape", partial_img_emb.shape) mask = torch.ones(4, partial_img_emb.shape[2], partial_img_emb.shape[3], device=partial_img_emb.device, dtype=partial_img_emb.dtype) #torch.Size([20, 10, 104, 60]) print("mask.shape", mask.shape) print("self.previous_noise_pred_context.shape", self.previous_noise_pred_context.shape) #torch.Size([16, 10, 104, 60]) - partial_img_emb[:, 0, :, :] = torch.cat([self.previous_noise_pred_context[:, -1, :, :], mask], dim=0) + partial_img_emb[:, 0, :, :] = torch.cat([image[0][:,0,:,:], mask], dim=0) else: partial_img_emb[:, 0, :, :] = partial_image_cond diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 4813b54..e22e67c 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -52,44 +52,46 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0): from comfy.model_management import get_torch_device, get_autocast_device @torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False) -@torch.compiler.disable() def rope_apply(x, grid_sizes, freqs): - batch_size, max_seq, n, c_total = x.shape - c = c_total // 2 + n, c = x.size(2), x.size(3) // 2 - # Static splits - c1 = c - 2 * (c // 3) - c2 = c // 3 - c3 = c // 3 - freqs_split = freqs.split([c1, c2, c3], dim=1) - - def process_chunk(x_chunk, f, h, w, freqs_split): - seq_len = f * h * w - x_complex = torch.view_as_complex(x_chunk.to(torch.float64).reshape(seq_len, n, c, 2)) - - f1 = freqs_split[0][:f].view(f, 1, 1, -1) - f2 = freqs_split[1][:h].view(1, h, 1, -1) - f3 = freqs_split[2][:w].view(1, 1, w, -1) - - freq_cat = torch.cat([ - f1.repeat(1, h, w, 1), - f2.repeat(f, 1, w, 1), - f3.repeat(f, h, 1, 1) - ], dim=-1).reshape(seq_len, 1, c) - - x_i = torch.view_as_real(x_complex * freq_cat) - return x_i.reshape(seq_len, n, c_total) + freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1) output = [] for i, (f, h, w) in enumerate(grid_sizes.tolist()): seq_len = f * h * w - x_i = process_chunk(x[i, :seq_len], f, h, w, freqs_split) - if seq_len < max_seq: + @torch.compiler.disable() + def view_as_complex_no_compile(x): + x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2)) + return x_i + + x_i = view_as_complex_no_compile(x) + + f_size = (f, 1, 1, -1) + h_size = (1, h, 1, -1) + w_size = (1, 1, w, -1) + + freq_cat = torch.cat([ + freqs[0][:f].view(*f_size).expand(f, h, w, -1), + freqs[1][:h].view(*h_size).expand(f, h, w, -1), + freqs[2][:w].view(*w_size).expand(f, h, w, -1) + ], dim=-1).reshape(seq_len, 1, -1) + + @torch.compiler.disable() + def view_as_real_no_compile(x_i): + x_i.mul_(freq_cat) + x_i = torch.view_as_real(x_i).flatten(2) + return x_i + + x_i = view_as_real_no_compile(x_i) + del freq_cat + + if seq_len < x.size(1): x_i = torch.cat([x_i, x[i, seq_len:]], dim=0) output.append(x_i) - return torch.stack(output) + return torch.stack(output).to(torch.float32) def rope_apply_original(x, grid_sizes, freqs): n, c = x.size(2), x.size(3) // 2