From d15cf3001fe08a97655eec49689394f6333c974a Mon Sep 17 00:00:00 2001 From: chengzeyi Date: Tue, 28 Oct 2025 12:08:23 +0000 Subject: [PATCH] Fix dtype mismatch in ref_conv forward pass MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This commit fixes a RuntimeError that occurs when using Fun-Control reference images: "Input type (float) and bias type (c10::Half) should be the same" Root cause: - Commit 1ba1a16 changed the dtype handling strategy to convert the main latent `x` to `base_dtype` instead of converting embeddings to match `x.dtype` - This caused `fun_ref` input to be in a different dtype than the `ref_conv` layer's weights and bias - Line 2324 already handles this correctly for `attn_cond` by converting to `self.attn_conv_in.weight.dtype` Solution: - Convert `fun_ref` to match `self.ref_conv.weight.dtype` before passing through the convolution layer - This follows the same pattern used for `attn_cond` on line 2324 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude --- wanvideo/modules/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 9264a73..a1e27f1 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2329,7 +2329,7 @@ class WanModel(torch.nn.Module): block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=F+1) if self.ref_conv is not None and fun_ref is not None: - fun_ref = self.ref_conv(fun_ref).flatten(2).transpose(1, 2) + fun_ref = self.ref_conv(fun_ref.to(self.ref_conv.weight.dtype)).flatten(2).transpose(1, 2) grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) seq_len += fun_ref.size(1) F += 1