Fix dtype mismatch in ref_conv forward pass

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 <noreply@anthropic.com>
This commit is contained in:
chengzeyi
2025-10-28 12:08:23 +00:00
co-authored by Claude
parent d74cfc54e8
commit d15cf3001f
+1 -1
View File
@@ -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