Allow HuMo to work with start_image
This commit is contained in:
+10
-2
@@ -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([
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user