From 123c9ca3124ffdeffaf9c36cd9e90ef80fcb486e Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 7 Dec 2025 18:18:16 +0200 Subject: [PATCH] Allow HuMo to work with start_image --- nodes_sampler.py | 12 ++++++++++-- wanvideo/modules/model.py | 11 ++++++----- 2 files changed, 16 insertions(+), 7 deletions(-) diff --git a/nodes_sampler.py b/nodes_sampler.py index fab7dfd..03d887f 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1234,6 +1234,7 @@ class WanVideoSampler: (ati_end_percent > 0 and idx == 0 and current_step_percentage >= ati_start_percent)): image_cond_input = image_cond_ati.to(z) elif humo_image_cond is not None: + humo_image_cond_neg_input = None if context_window is not None: image_cond_input = humo_image_cond[:, context_window].to(z) humo_image_cond_neg_input = humo_image_cond_neg[:, context_window].to(z) @@ -1241,8 +1242,15 @@ class WanVideoSampler: image_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:] humo_image_cond_neg_input[:, -humo_reference_count:] = humo_image_cond_neg[:, -humo_reference_count:] else: - image_cond_input = humo_image_cond.to(z) - humo_image_cond_neg_input = humo_image_cond_neg.to(z) + if image_cond is not None: + image_cond_input = image_cond.to(z) + if humo_reference_count > 0: + image_cond_input = torch.cat([image_cond, humo_image_cond[:, -humo_reference_count:].to(z)], dim=1) + humo_image_cond_neg_input = torch.cat([image_cond, humo_image_cond_neg[:, -humo_reference_count:].to(z)], dim=1) + else: + image_cond_input = humo_image_cond.to(z) + humo_image_cond_neg_input = humo_image_cond_neg.to(z) + elif image_cond is not None: if reverse_time: # Flip the image condition image_cond_input = torch.cat([ diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index a31665d..fd856c8 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -800,16 +800,16 @@ class WanHuMoCrossAttention(WanSelfAttention): self.attention_mode = attention_mode def forward(self, x, context, grid_sizes, **kwargs): - + b, n, d = x.size(0), self.num_heads, self.head_dim - q = self.norm_q(self.q(x)).view(b, -1, n, d) - k = self.norm_k(self.k(context)).view(b, -1, n, d) + q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype).to(x.dtype)).view(b, -1, n, d) + k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype).to(context.dtype)).view(b, -1, n, d) v = self.v(context).view(b, -1, n, d) # Handle video spatial structure hlen_wlen = grid_sizes[0][1] * grid_sizes[0][2] q = q.reshape(-1, hlen_wlen, n, d) - + # Handle audio temporal structure (16 tokens per frame) k = k.reshape(-1, 16, n, d) v = v.reshape(-1, 16, n, d) @@ -820,7 +820,7 @@ class WanHuMoCrossAttention(WanSelfAttention): x = x_text return self.o(x) - + class AudioCrossAttentionWrapper(nn.Module): def __init__(self, in_features, out_features, num_heads, qk_norm=True, eps=1e-6, kv_dim=None): super().__init__() @@ -829,6 +829,7 @@ class AudioCrossAttentionWrapper(nn.Module): self.norm1_audio = WanLayerNorm(out_features, eps, elementwise_affine=True) def forward(self, x, audio, grid_sizes, humo_audio_scale=1.0): + x = x.to(self.norm1_audio.weight.dtype) x = x + self.audio_cross_attn(self.norm1_audio(x), audio, grid_sizes) * humo_audio_scale return x