diff --git a/lynx/modules.py b/lynx/modules.py index 15087bf..89b922f 100644 --- a/lynx/modules.py +++ b/lynx/modules.py @@ -14,12 +14,11 @@ def merge_token_lists(list1, list2, dim): assert(len(list1) == len(list2)) return [torch.cat((t1, t2), dim) for t1, t2 in zip(list1, list2)] -try: +try: from sageattention import sageattn_varlen except ImportError: sageattn_varlen = None - class WanLynxIPCrossAttention(nn.Module): def __init__(self, cross_attention_dim=5120, dim=5120, n_registers=16, bias=True): super().__init__() @@ -29,7 +28,7 @@ class WanLynxIPCrossAttention(nn.Module): self.registers = nn.Parameter(torch.randn(1, n_registers, cross_attention_dim) / dim**0.5) else: self.registers = None - + def forward(self, block, q, x, ip_x): b, n, d = x.size(0), block.num_heads, block.head_dim s = q.shape[1] @@ -50,99 +49,37 @@ class WanLynxIPCrossAttention(nn.Module): ip_key = ip_key * torch.rsqrt(ip_key.pow(2).mean(dim=-1, keepdim=True) + 1e-5).to(ip_key.dtype) else: # full model ip_key = block.norm_k(ip_key) - if sageattn_varlen is not None: - q_lens = [s] * b - k_lens = ip_lens - - cu_seqlens_q = torch.tensor([0] + list(torch.cumsum(torch.tensor(q_lens), dim=0)), device=q.device, dtype=torch.int32) - cu_seqlens_k = torch.tensor([0] + list(torch.cumsum(torch.tensor(k_lens), dim=0)), device=q.device, dtype=torch.int32) - - return sageattn_varlen( - q.view(-1, n, d), - ip_key.view(-1, n, d), - ip_value.view(-1, n, d), - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k=cu_seqlens_k, - max_seqlen_k=max(ip_lens), - max_seqlen_q=s, - ).reshape(b, -1, n * d) return attention( - q, - ip_key.view(b, -1, n, d), + q, + ip_key.view(b, -1, n, d), ip_value.view(b, -1, n, d) ).reshape(b, -1, n * d) - -# class WanLynxIPCrossAttention(nn.Module): -# def __init__(self, cross_attention_dim=5120, dim=5120, n_registers=16, bias=True): -# super().__init__() -# self.to_k_ip = nn.Linear(cross_attention_dim, dim, bias=bias) -# self.to_v_ip = nn.Linear(cross_attention_dim, dim, bias=bias) -# if n_registers > 0: -# self.registers = nn.Parameter(torch.randn(1, n_registers, cross_attention_dim) / dim**0.5) -# else: -# self.registers = None - -# def forward(self, block, q, x, ip_x): -# b, n, d = x.size(0), block.num_heads, block.head_dim -# s = q.shape[1] - -# if self.registers is not None: -# #print("self.registers.shape", self.registers.shape) #torch.Size([1, 16, 5120]) -# #print("ip_x.shape", ip_x.shape) #torch.Size([1, 16, 5120]) - -# ip_lens = [ip_x.shape[1]] -# ip_x_list = vector_to_list(ip_x, ip_lens, 1) -# ip_x_list = merge_token_lists(ip_x_list, [self.registers] * len(ip_x_list), 1) -# ip_x, ip_lens = list_to_vector(ip_x_list, 1) - -# ip_key = self.to_k_ip(ip_x) -# ip_value = self.to_v_ip(ip_x) - -# if self.registers is None: # lite model normalization -# ip_key = ip_key * torch.rsqrt(ip_key.pow(2).mean(dim=-1, keepdim=True) + 1e-5).to(ip_key.dtype) -# else: # full model normalization -# ip_key = block.norm_k(ip_key) - -# q_lens = [s] * b -# k_lens = ip_lens - -# cu_seqlens_q = torch.tensor([0] + list(torch.cumsum(torch.tensor(q_lens), dim=0)), device=q.device, dtype=torch.int32) -# cu_seqlens_k = torch.tensor([0] + list(torch.cumsum(torch.tensor(k_lens), dim=0)), device=q.device, dtype=torch.int32) - -# ip_x = sageattn_varlen( -# q.view(-1, n, d), -# ip_key.view(-1, n, d), -# ip_value.view(-1, n, d), -# cu_seqlens_q=cu_seqlens_q, -# cu_seqlens_k=cu_seqlens_k, -# max_seqlen_k=max(ip_lens), -# max_seqlen_q=s, - -# ).reshape(b, -1, n * d) - -# return ip_x +@torch.compiler.disable() class WanLynxRefAttention(nn.Module): def __init__(self, dim=5120, bias=True, attention_mode="sdpa"): super().__init__() self.to_k_ref = nn.Linear(dim, dim, bias=bias) self.to_v_ref = nn.Linear(dim, dim, bias=bias) self.attention_mode = attention_mode - + # Pre-compute attention mode flags to avoid string operations in forward + self.use_flash_attn = "flash_attn" in attention_mode + self.use_sageattn = sageattn_varlen is not None + def forward(self, block, q, ref_feature): - b, s, n, d = q.shape + b, s, n, d = q.shape ref_key = self.to_k_ref(ref_feature) ref_value = self.to_v_ref(ref_feature) ref_key = block.norm_k(ref_key) - attn_mask = None - if not "flash_attn" in self.attention_mode and sageattn_varlen is None: + # Use pre-computed flags instead of runtime string checks + if not self.use_flash_attn and not self.use_sageattn: # Pad ref_key and ref_value to match q's sequence length (s) seq_len = ref_key.shape[1] - pad_len = s - ref_key.shape[1] + pad_len = s - seq_len if pad_len > 0: # Pad on the sequence dimension (dim=1) ref_key = torch.nn.functional.pad(ref_key, (0, 0, 0, pad_len)) @@ -156,37 +93,25 @@ class WanLynxRefAttention(nn.Module): ref_value = ref_value.view(b, s, n, d) ref_x = attention( - q, - ref_key, + q, + ref_key, ref_value, attention_mode="sdpa", attn_mask=attn_mask, ) - elif sageattn_varlen is not None: + else: q_lens = [s] * b k_lens = [ref_key.shape[1]] * b - cu_seqlens_q = torch.tensor([0] + list(torch.cumsum(torch.tensor(q_lens), dim=0)), device=q.device, dtype=torch.int32) - cu_seqlens_k = torch.tensor([0] + list(torch.cumsum(torch.tensor(k_lens), dim=0)), device=q.device, dtype=torch.int32) - - ref_x = sageattn_varlen( - q.view(-1, n, d), - ref_key.view(-1, n, d), - ref_value.view(-1, n, d), - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k=cu_seqlens_k, - max_seqlen_k=max(k_lens), - max_seqlen_q=s - ) - else: - ref_key = ref_key.view(-1, n, d) - ref_value = ref_value.view(-1, n, d) + ref_x = attention( - q, - ref_key, - ref_value, - q_lens=torch.tensor([s]*b, device=q.device), - k_lens=torch.tensor([ref_key.shape[1]]*b, device=q.device), - attention_mode=self.attention_mode, - ) + q.view(-1, n, d), + ref_key.view(-1, n, d), + ref_value.view(-1, n, d), + q_lens=q_lens, + k_lens=k_lens, + max_seqlen_k=ref_key.shape[1], + max_seqlen_q=s, + attention_mode='sageattn_varlen' if self.use_sageattn else self.attention_mode, + ) return ref_x \ No newline at end of file diff --git a/lynx/nodes.py b/lynx/nodes.py index c6ec994..1e25755 100644 --- a/lynx/nodes.py +++ b/lynx/nodes.py @@ -221,8 +221,8 @@ class WanVideoAddLynxEmbeds: if ref_image is not None: vae.to(device) ref_image_in = (ref_image[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype) - ref_latent = vae.encode([ref_image_in], device, tiled=False) - ref_latent_uncond = vae.encode([torch.zeros_like(ref_image_in)], device, tiled=False) + ref_latent = vae.encode([ref_image_in], device, tiled=False, sample=True) + ref_latent_uncond = vae.encode([torch.zeros_like(ref_image_in)], device, tiled=False, sample=True) vae.to(offload_device) new_entry = { diff --git a/nodes_sampler.py b/nodes_sampler.py index 7824e4b..d6064ae 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1035,6 +1035,8 @@ class WanVideoSampler: ) log.info(f"Extracted {len(lynx_ref_buffer_uncond)} uncond ref buffers") + lynx_embeds["ip_x"] = lynx_embeds["ip_x"].to(device, dtype) + lynx_embeds["ip_x_uncond"] = lynx_embeds["ip_x_uncond"].to(device, dtype) lynx_embeds["ref_feature_extractor"] = False lynx_embeds["ref_latent"] = lynx_embeds["ref_text_embed"] = None lynx_embeds["ref_buffer"] = lynx_ref_buffer diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index a6a2b32..c6be1bc 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -42,6 +42,20 @@ except: except: SAGE3_AVAILABLE = False +try: + from sageattention import sageattn_varlen + @torch.compiler.disable() + def sageattn_varlen_func(q, k, v, q_lens, k_lens, max_seqlen_q, max_seqlen_k, dropout_p=0, is_causal=False): + cu_seqlens_q = torch.tensor([0] + list(torch.cumsum(torch.tensor(q_lens), dim=0)), device=q.device, dtype=torch.int32) + cu_seqlens_k = torch.tensor([0] + list(torch.cumsum(torch.tensor(k_lens), dim=0)), device=q.device, dtype=torch.int32) + if not (q.dtype == k.dtype == v.dtype): + return sageattn_varlen(q, k.to(q.dtype), v.to(q.dtype), cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal) + elif q.dtype == torch.float32: + return sageattn_varlen(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal).to(torch.float32) + else: + return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal) +except: + sageattn_varlen_func = None __all__ = [ 'flash_attention', @@ -162,6 +176,8 @@ def attention( v, q_lens=None, k_lens=None, + max_seqlen_q=None, + max_seqlen_k=None, dropout_p=0., softmax_scale=None, q_scale=None, @@ -203,5 +219,13 @@ def attention( v.transpose(1,2), per_block_mean=False #seems necessary for reasonable VRAM usage, not sure of other implications ).transpose(1,2).contiguous() + elif attention_mode == 'sageattn_varlen': + return sageattn_varlen_func( + q,k,v, + q_lens=q_lens, + k_lens=k_lens, + max_seqlen_k=max_seqlen_k, + max_seqlen_q=max_seqlen_q + ) else: return sageattn_func(q, k, v, tensor_layout="NHD").contiguous() diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 557a115..3c388b6 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -442,10 +442,13 @@ class WanSelfAttention(nn.Module): if attention_mode_override is not None: attention_mode = attention_mode_override + if self.ref_adapter is not None and lynx_ref_feature is not None: + ref_x = self.ref_adapter(self, q, lynx_ref_feature) + x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode) if self.ref_adapter is not None and lynx_ref_feature is not None: - x = x.add(self.ref_adapter(self, q, lynx_ref_feature), alpha=lynx_ref_scale) + x = x.add(ref_x, alpha=lynx_ref_scale) # output return self.o(x.flatten(2)) @@ -2150,7 +2153,7 @@ class WanModel(torch.nn.Module): hidden_states = torch.cat([hidden_states, torch.zeros_like(hidden_states[:, :4])], dim=1) render_latent = torch.cat([hidden_states[:, :20], render_latent], dim=1) - # embeddings + # patch embed if control_lora_enabled: self.expanded_patch_embedding.to(device) x = [