compile fixes and cleanup
This commit is contained in:
+26
-101
@@ -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
|
||||
+2
-2
@@ -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 = {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 = [
|
||||
|
||||
Reference in New Issue
Block a user