From a0bdf208175f7aaa8e936937b859b77cfe1582f0 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 26 Oct 2025 23:05:17 +0200 Subject: [PATCH] Some cleanup and allow full block swap --- nodes_model_loading.py | 2 +- wanvideo/modules/model.py | 25 ++++++++++++------------- 2 files changed, 13 insertions(+), 14 deletions(-) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index ca3da12..baa2ff5 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -278,7 +278,7 @@ class WanVideoBlockSwap: def INPUT_TYPES(s): return { "required": { - "blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 40, "step": 1, "tooltip": "Number of transformer blocks to swap, the 14B model has 40, while the 1.3B model has 30 blocks"}), + "blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 48, "step": 1, "tooltip": "Number of transformer blocks to swap, the 14B model has 40, while the 1.3B and 5B models have 30 blocks. LongCat-video has 48 blocks"}), "offload_img_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload img_emb to offload_device"}), "offload_txt_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload time_emb to offload_device"}), }, diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index b86937b..668135a 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1232,19 +1232,18 @@ class WanAttentionBlock(nn.Module): full_k = torch.cat([k, k_ip], dim=1) full_v = torch.cat([v, v_ip], dim=1) y = self.self_attn.forward(q, full_k, full_v, seq_lens) - elif is_longcat: - if num_cond_latents is not None and num_cond_latents > 0: - num_cond_latents_thw = num_cond_latents * (N // grid_sizes[0][0]) - # process the condition tokens - q_cond = q[:, :num_cond_latents_thw].contiguous() - k_cond = k[:, :num_cond_latents_thw].contiguous() - v_cond = v[:, :num_cond_latents_thw].contiguous() - x_cond = self.self_attn.forward(q_cond, k_cond, v_cond, seq_lens) - # process the noise tokens - q_noise = q[:, num_cond_latents_thw:].contiguous() - x_noise = self.self_attn.forward(q_noise, k, v, seq_lens) - # merge x_cond and x_noise - y = torch.cat([x_cond, x_noise], dim=1).contiguous() + elif is_longcat and num_cond_latents is not None and num_cond_latents > 0: + num_cond_latents_thw = num_cond_latents * (N // grid_sizes[0][0]) + # process the condition tokens + x_cond = self.self_attn.forward( + q[:, :num_cond_latents_thw].contiguous(), + k[:, :num_cond_latents_thw].contiguous(), + v[:, :num_cond_latents_thw].contiguous(), + seq_lens) + # process the noise tokens + x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens) + # merge x_cond and x_noise + y = torch.cat([x_cond, x_noise], dim=1).contiguous() else: y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale)