Allow HuMo to work with start_image

This commit is contained in:
kijai
2025-12-07 18:18:16 +02:00
parent 2b62866945
commit 123c9ca312
2 changed files with 16 additions and 7 deletions
+10 -2
View File
@@ -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([
+6 -5
View File
@@ -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