From 09b4c3a865ffecd73e39082d6682117f78b8e777 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 19 Jun 2025 20:29:32 +0300 Subject: [PATCH] support context windows, torch compile fixes --- multitalk/multitalk.py | 28 ++++++++++++++++++-------- multitalk/nodes.py | 2 +- nodes.py | 42 +++++++++++++++++++++++++-------------- wanvideo/modules/model.py | 2 +- 4 files changed, 49 insertions(+), 25 deletions(-) diff --git a/multitalk/multitalk.py b/multitalk/multitalk.py index 297cd7a..aa4427f 100644 --- a/multitalk/multitalk.py +++ b/multitalk/multitalk.py @@ -3,6 +3,7 @@ from einops import rearrange, repeat import torch import torch.nn as nn from functools import lru_cache +from ..wanvideo.modules.attention import attention from comfy import model_management as mm @@ -187,6 +188,7 @@ class AudioProjModel(ModelMixin, ConfigMixin): return context_tokens +#@torch.compiler.disable() class SingleStreamAttention(nn.Module): def __init__( self, @@ -199,6 +201,7 @@ class SingleStreamAttention(nn.Module): attn_drop: float = 0.0, proj_drop: float = 0.0, eps: float = 1e-6, + attention_mode: str = 'sdpa', ) -> None: super().__init__() assert dim % num_heads == 0, "dim should be divisible by num_heads" @@ -223,11 +226,12 @@ class SingleStreamAttention(nn.Module): self.add_q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() self.add_k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() + self.attention_mode = attention_mode + def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None) -> torch.Tensor: N_t, N_h, N_w = shape - if not enable_sp: - x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) + x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) # get q for hidden_state B, N, C = x.shape @@ -248,19 +252,23 @@ class SingleStreamAttention(nn.Module): if self.qk_norm: encoder_k = self.add_k_norm(encoder_k) - x = torch.nn.functional.scaled_dot_product_attention( - q, encoder_k, encoder_v, attn_mask=None, is_causal=False, dropout_p=0.0) + x = attention( + q.transpose(1, 2), + encoder_k.transpose(1, 2), + encoder_v.transpose(1, 2), + attention_mode=self.attention_mode + ) + #x = torch.nn.functional.scaled_dot_product_attention( + # q, encoder_k, encoder_v, attn_mask=None, is_causal=False, dropout_p=0.0) # linear transform x_output_shape = (B, N, C) - x = x.transpose(1, 2) + #x = x.transpose(1, 2) x = x.reshape(x_output_shape) x = self.proj(x) x = self.proj_drop(x) - if not enable_sp: - # reshape x to origin shape - x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t) + x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t) return x @@ -278,6 +286,7 @@ class SingleStreamMultiAttention(SingleStreamAttention): eps: float = 1e-6, class_range: int = 24, class_interval: int = 4, + attention_mode: str = 'sdpa', ) -> None: super().__init__( dim=dim, @@ -289,6 +298,7 @@ class SingleStreamMultiAttention(SingleStreamAttention): attn_drop=attn_drop, proj_drop=proj_drop, eps=eps, + attention_mode=attention_mode, ) self.class_interval = class_interval self.class_range = class_range @@ -298,6 +308,8 @@ class SingleStreamMultiAttention(SingleStreamAttention): self.rope_1d = RotaryPositionalEmbedding1D(self.head_dim) + self.attention_mode = attention_mode + def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, diff --git a/multitalk/nodes.py b/multitalk/nodes.py index b8f7003..5abc9a3 100644 --- a/multitalk/nodes.py +++ b/multitalk/nodes.py @@ -138,7 +138,7 @@ class MultiTalkWav2VecEmbeds: # audio encoder audio_duration = len(audio_segment) / sr - video_length = audio_duration * 25 # Assume the video fps is 25 + video_length = audio_duration * fps print("Audio duration:", audio_duration, "Video length:", video_length) embeddings = wav2vec(audio_feature.to(dtype), seq_len=int(video_length), output_hidden_states=True) diff --git a/nodes.py b/nodes.py index 048bbe7..3f6c3cc 100644 --- a/nodes.py +++ b/nodes.py @@ -750,7 +750,8 @@ class WanVideoModelLoader: eps=transformer.eps, norm_layer=WanRMSNorm, class_range=24, - class_interval=4 + class_interval=4, + attention_mode=attention_mode, ) block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True) if norm_input_visual else nn.Identity() log.info("MultiTalk model detected, patching model...") @@ -3183,7 +3184,8 @@ class WanVideoSampler: #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, - control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None, add_cond=None, cache_state=None): + control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None, + add_cond=None, cache_state=None, context_window=None): z = z.to(dtype) with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])): @@ -3283,18 +3285,29 @@ class WanVideoSampler: audio_embs = [] indices = (torch.arange(4 + 1) - 2) * 1 # split audio with window size - for human_idx in range(1): - center_indices = torch.arange( - 0, - latent_video_length * 4 + 1 if add_cond is not None else (latent_video_length-1) * 4 + 1, - 1, - ).unsqueeze( - 1 - ) + indices.unsqueeze(0) - center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0] - 1) - audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device) - audio_embs.append(audio_emb) + if context_window is None: + for human_idx in range(1): + center_indices = torch.arange( + 0, + latent_video_length * 4 + 1 if add_cond is not None else (latent_video_length-1) * 4 + 1, + 1, + ).unsqueeze( + 1 + ) + indices.unsqueeze(0) + center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0] - 1) + audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device) + audio_embs.append(audio_emb) + else: + for human_idx in range(1): + audio_start = context_window[0] * 4 + audio_end = context_window[-1] * 4 + 1 + print("audio_start: ", audio_start, "audio_end: ", audio_end) + center_indices = torch.arange(audio_start, audio_end, 1).unsqueeze(1) + indices.unsqueeze(0) + center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0] - 1) + audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device) + audio_embs.append(audio_emb) audio_embs = torch.concat(audio_embs, dim=0).to(dtype) + base_params = { 'seq_len': seq_len, @@ -3714,8 +3727,7 @@ class WanVideoSampler: cfg[idx], positive, text_embeds["negative_prompt_embeds"], timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj, - partial_control_camera_latents, partial_add_cond, - current_teacache) + partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c) if cache_args is not None: self.window_tracker.cache_states[window_id] = new_teacache diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index c65b959..1a6e615 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1313,7 +1313,7 @@ class WanModel(ModelMixin, ConfigMixin): x = [u + v for u, v in zip(x, fun_camera)] grid_sizes = torch.stack( - [torch.tensor(u.shape[2:], dtype=torch.long) for u in x]) + [torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x]) x = [u.flatten(2).transpose(1, 2) for u in x]