diff --git a/nodes_model_loading.py b/nodes_model_loading.py index baa2ff5..a35fd9b 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": 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"}), + "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"}), "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"}), }, @@ -878,10 +878,12 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, scale_key = key.replace(".weight", ".scale_weight") if scale_key in sd: dtype_to_use = value.dtype - if "modulation" in name or "norm" in name or "bias" in name or "img_emb" in name: + if "bias" in name or "img_emb" in name: dtype_to_use = base_dtype if "patch_embedding" in name or "motion_encoder" in name: dtype_to_use = torch.float32 + if "modulation" in name or "norm" in name: + dtype_to_use = value.dtype if value.dtype == torch.float32 else base_dtype load_device = transformer_load_device if block_swap_args is not None: diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 668135a..0d2d39e 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -473,9 +473,9 @@ class WanSelfAttention(nn.Module): def qkv_fn_longcat(self, x): b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim q = self.q(x).view(b, s, n, d) - q = self.norm_q(q) + q = self.norm_q(q.float()).to(x.dtype) k = self.k(x).view(b, s, n, d) - k = self.norm_k(k) + k = self.norm_k(k.float()).to(x.dtype) v = self.v(x).view(b, s, n, d) return q, k, v @@ -1106,11 +1106,12 @@ class WanAttentionBlock(nn.Module): shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation) del e + input_dtype = x.dtype B, N, C = x.shape T = num_latent_frames is_longcat = C == 4096 if is_longcat: - input_x = self.modulate(self.norm1(x.view(B, T, -1, C).float()).to(x.dtype), shift_msa, scale_msa, seg_idx=self.seg_idx).view(B, N, C) + input_x = self.modulate(self.norm1(x.view(B, T, -1, C).float()), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype).view(B, N, C) else: input_x = self.modulate(self.norm1(x), shift_msa, scale_msa, seg_idx=self.seg_idx) @@ -1276,7 +1277,7 @@ class WanAttentionBlock(nn.Module): if not is_longcat: x = x.addcmul(y, gate_msa) else: - x = x + (y.view(B, -1, N//T, C).float() * gate_msa).view(B, -1, C).to(x.dtype) + x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C) del y, gate_msa # cross-attention & ffn function @@ -1311,6 +1312,7 @@ class WanAttentionBlock(nn.Module): num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond, rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs, adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, num_cond_latents=num_cond_latents) + x = x.to(input_dtype) # MultiTalk if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock): x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding, @@ -1325,7 +1327,8 @@ class WanAttentionBlock(nn.Module): # HuMo Audio Cross-Attention if humo_audio_input is not None: x = self.audio_cross_attn_wrapper(x, humo_audio_input, grid_sizes, humo_audio_scale) - + + # ffn if self.rope_func == "comfy_chunked": x_ffn = self.ffn_chunked(x, shift_mlp, scale_mlp) @@ -1342,7 +1345,7 @@ class WanAttentionBlock(nn.Module): if not is_longcat: mod_x = torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp) else: - mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()).to(x.dtype), 1 + scale_mlp).view(B, -1, C) + mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C).to(input_dtype) x_ffn = self.ffn(mod_x) del shift_mlp, scale_mlp @@ -1357,14 +1360,13 @@ class WanAttentionBlock(nn.Module): if not is_longcat: x = x.addcmul(x_ffn, gate_mlp) else: - x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).view(B, -1, C).to(x.dtype) + x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C) del gate_mlp if x_ip is not None: #stand-in x_ip = x_ip.addcmul(y_ip, gate_msa_ip) y_ip = self.ffn(torch.addcmul(shift_mlp_ip, self.norm2(x_ip), 1 + scale_mlp_ip)) x_ip = x_ip.addcmul(y_ip, gate_mlp_ip) - return x, x_ip, lynx_ref_feature, x_ovi @torch.compiler.disable()