support context windows, torch compile fixes

This commit is contained in:
kijai
2025-06-19 20:29:32 +03:00
parent f3614e6720
commit 09b4c3a865
4 changed files with 49 additions and 25 deletions
+20 -8
View File
@@ -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,
+1 -1
View File
@@ -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)
+27 -15
View File
@@ -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,19 +3285,30 @@ 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,
'device': device,
@@ -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
+1 -1
View File
@@ -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]