support context windows, torch compile fixes
This commit is contained in:
+20
-8
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user