Add UltraViCo -sage attention mode and refactor some attention code
https://github.com/thu-ml/DiT-Extrapolation/
This commit is contained in:
@@ -1010,6 +1010,7 @@ class WanVideoModelLoader:
|
||||
"sageattn_3",
|
||||
"radial_sage_attention",
|
||||
"sageattn_compiled",
|
||||
"sageattn_ultravico",
|
||||
], {"default": "sdpa"}),
|
||||
"compile_args": ("WANCOMPILEARGS", ),
|
||||
"block_swap_args": ("BLOCKSWAPARGS", ),
|
||||
|
||||
@@ -0,0 +1,230 @@
|
||||
# https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/attn_qk_int8_per_block.py
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag,
|
||||
K_ptrs, K_scale_ptr, V_ptrs, stride_kn, stride_vn,
|
||||
Block_bias_ptrs, stride_bbz, stride_bbh, stride_bm, stride_bn,
|
||||
Decay_mask_ptrs, stride_dmz, stride_dmh, stride_dm, stride_dn,
|
||||
start_m,
|
||||
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr, BLOCK_N: tl.constexpr,
|
||||
STAGE: tl.constexpr, offs_m: tl.constexpr, offs_n: tl.constexpr,
|
||||
xpos_xi: tl.constexpr = 0.9999934149894527,
|
||||
frame_tokens: tl.constexpr = 1560,
|
||||
sigmoid_a: tl.constexpr = 1.0,
|
||||
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
|
||||
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
|
||||
sink_width: tl.constexpr = 4,
|
||||
window_width: tl.constexpr = 16,
|
||||
multi_factor: tl.constexpr = None,
|
||||
entropy_factor: tl.constexpr = None,
|
||||
):
|
||||
|
||||
|
||||
lo, hi = 0, kv_len
|
||||
for start_n in range(lo, hi, BLOCK_N):
|
||||
start_n = tl.multiple_of(start_n, BLOCK_N)
|
||||
k_mask = offs_n[None, :] < (kv_len - start_n)
|
||||
k = tl.load(K_ptrs, mask = k_mask)
|
||||
k_scale = tl.load(K_scale_ptr)
|
||||
|
||||
|
||||
m = offs_m[:, None]
|
||||
n = start_n + offs_n
|
||||
|
||||
qk = tl.dot(q, k).to(tl.float32) * q_scale * k_scale
|
||||
|
||||
window_th = 1560 * 21 / 2
|
||||
dist2 = tl.abs(m - n).to(tl.int32)
|
||||
dist_mask = dist2 <= window_th
|
||||
|
||||
negative_mask = (qk<0)
|
||||
|
||||
qk = tl.where(dist_mask | negative_mask, qk, qk*multi_factor)
|
||||
|
||||
window3 = (m <= frame_tokens) & (n > 21*frame_tokens)
|
||||
qk = tl.where(window3, -1e4, qk)
|
||||
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, 1))
|
||||
qk = qk - m_ij[:, None]
|
||||
p = tl.math.exp2(qk)
|
||||
l_ij = tl.sum(p, 1)
|
||||
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
|
||||
acc = acc * alpha[:, None]
|
||||
|
||||
v = tl.load(V_ptrs, mask = offs_n[:, None] < (kv_len - start_n))
|
||||
p = p.to(tl.float16)
|
||||
|
||||
acc += tl.dot(p, v, out_dtype=tl.float16)
|
||||
m_i = m_ij
|
||||
K_ptrs += BLOCK_N * stride_kn
|
||||
K_scale_ptr += 1
|
||||
V_ptrs += BLOCK_N * stride_vn
|
||||
return acc, l_i
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd(Q, K, V, Q_scale, K_scale, Out,
|
||||
Block_bias, Decay_mask,
|
||||
flags, stride_f_b, stride_f_h,
|
||||
stride_qz, stride_qh, stride_qn,
|
||||
stride_kz, stride_kh, stride_kn,
|
||||
stride_vz, stride_vh, stride_vn,
|
||||
stride_oz, stride_oh, stride_on,
|
||||
stride_bbz, stride_bbh, stride_bm, stride_bn,
|
||||
stride_dmz, stride_dmh, stride_dm, stride_dn,
|
||||
qo_len, kv_len, H: tl.constexpr, num_kv_groups: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
STAGE: tl.constexpr,
|
||||
xpos_xi: tl.constexpr = 0.9999934149894527,
|
||||
frame_tokens: tl.constexpr = 1560,
|
||||
sigmoid_a: tl.constexpr = 1.0,
|
||||
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
|
||||
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
|
||||
sink_width: tl.constexpr = 4,
|
||||
window_width: tl.constexpr = 16,
|
||||
multi_factor: tl.constexpr = None,
|
||||
entropy_factor: tl.constexpr = None,
|
||||
):
|
||||
start_m = tl.program_id(0)
|
||||
|
||||
off_z = tl.program_id(2).to(tl.int64)
|
||||
off_h = tl.program_id(1).to(tl.int64)
|
||||
|
||||
q_scale_offset = (off_z * H + off_h) * tl.cdiv(qo_len, BLOCK_M)
|
||||
k_scale_offset = (off_z * (H // num_kv_groups) + off_h // num_kv_groups) * tl.cdiv(kv_len, BLOCK_N)
|
||||
|
||||
flag_ptr = flags + off_z * stride_f_b + off_h * stride_f_h
|
||||
current_flag = tl.load(flag_ptr)
|
||||
|
||||
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
offs_n = tl.arange(0, BLOCK_N)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
Q_ptrs = Q + (off_z * stride_qz + off_h * stride_qh) + offs_m[:, None] * stride_qn + offs_k[None, :]
|
||||
Q_scale_ptr = Q_scale + q_scale_offset + start_m
|
||||
K_ptrs = K + (off_z * stride_kz + (off_h // num_kv_groups) * stride_kh) + offs_n[None, :] * stride_kn + offs_k[:, None]
|
||||
K_scale_ptr = K_scale + k_scale_offset
|
||||
V_ptrs = V + (off_z * stride_vz + (off_h // num_kv_groups) * stride_vh) + offs_n[:, None] * stride_vn + offs_k[None, :]
|
||||
O_block_ptr = Out + (off_z * stride_oz + off_h * stride_oh) + offs_m[:, None] * stride_on + offs_k[None, :]
|
||||
|
||||
# # 计算block_bias指针
|
||||
Block_bias_ptrs = Block_bias + off_z * stride_bbz + off_h * stride_bbh
|
||||
|
||||
# 计算decay_mask指针
|
||||
Decay_mask_ptrs = Decay_mask + off_z * stride_dmz + off_h * stride_dmh
|
||||
|
||||
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
q = tl.load(Q_ptrs, mask = offs_m[:, None] < qo_len)
|
||||
q_scale = tl.load(Q_scale_ptr)
|
||||
acc, l_i = _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag, K_ptrs, K_scale_ptr, V_ptrs,
|
||||
stride_kn, stride_vn,
|
||||
Block_bias_ptrs, stride_bbz, stride_bbh, stride_bm, stride_bn,
|
||||
Decay_mask_ptrs, stride_dmz, stride_dmh, stride_dm, stride_dn,
|
||||
start_m,
|
||||
BLOCK_M, HEAD_DIM, BLOCK_N,
|
||||
4 - STAGE, offs_m, offs_n,
|
||||
xpos_xi=xpos_xi,
|
||||
frame_tokens=frame_tokens,
|
||||
sigmoid_a=sigmoid_a,
|
||||
alpha_xpos_xi=alpha_xpos_xi,
|
||||
beta_xpos_xi=beta_xpos_xi,
|
||||
sink_width=sink_width,
|
||||
window_width=window_width,
|
||||
multi_factor=multi_factor,
|
||||
entropy_factor=entropy_factor,
|
||||
)
|
||||
acc = acc / l_i[:, None]
|
||||
tl.store(O_block_ptr, acc.to(Out.type.element_ty), mask = (offs_m[:, None] < qo_len))
|
||||
|
||||
def forward(q, k, v, flags, block_bias, decay_mask, q_scale, k_scale, tensor_layout="HND", output_dtype=torch.float16,
|
||||
xpos_xi: tl.constexpr = 0.9999934149894527,
|
||||
frame_tokens: tl.constexpr = 1560,
|
||||
sigmoid_a: tl.constexpr = 1.0,
|
||||
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
|
||||
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
|
||||
BLOCK_M: tl.constexpr = 128,
|
||||
BLOCK_N: tl.constexpr = 128,
|
||||
sink_width: tl.constexpr = 4,
|
||||
window_width: tl.constexpr = 16,
|
||||
multi_factor: tl.constexpr = None,
|
||||
entropy_factor: tl.constexpr = None,
|
||||
):
|
||||
stage = 1
|
||||
|
||||
o = torch.empty(q.shape, dtype=output_dtype, device=q.device)
|
||||
|
||||
b, h_qo, qo_len, head_dim = q.shape
|
||||
if block_bias is None:
|
||||
block_bias = torch.zeros((b, h_qo, (qo_len + BLOCK_M - 1) // BLOCK_M, (qo_len + BLOCK_N - 1) // BLOCK_N), dtype=torch.float16, device=q.device)
|
||||
|
||||
if decay_mask is None:
|
||||
decay_mask = torch.zeros((b, h_qo, (qo_len + BLOCK_M - 1) // BLOCK_M, (qo_len + BLOCK_N - 1) // BLOCK_N), dtype=torch.bool, device=q.device)
|
||||
|
||||
if tensor_layout == "HND":
|
||||
b, h_qo, qo_len, head_dim = q.shape
|
||||
_, h_kv, kv_len, _ = k.shape
|
||||
|
||||
stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(1), q.stride(2)
|
||||
stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(1), k.stride(2)
|
||||
stride_bz_v, stride_h_v, stride_seq_v = v.stride(0), v.stride(1), v.stride(2)
|
||||
stride_bz_o, stride_h_o, stride_seq_o = o.stride(0), o.stride(1), o.stride(2)
|
||||
stride_bbz, stride_bbh, stride_bm, stride_bn = block_bias.stride()
|
||||
stride_dmz, stride_dmh, stride_dm, stride_dn = decay_mask.stride()
|
||||
# elif tensor_layout == "NHD":
|
||||
# b, qo_len, h_qo, head_dim = q.shape
|
||||
# _, kv_len, h_kv, _ = k.shape
|
||||
|
||||
# stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(2), q.stride(1)
|
||||
# stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(2), k.stride(1)
|
||||
# stride_bz_v, stride_h_v, stride_seq_v = v.stride(0), v.stride(2), v.stride(1)
|
||||
# stride_bz_o, stride_h_o, stride_seq_o = o.stride(0), o.stride(2), o.stride(1)
|
||||
# stride_bbz, stride_bbh, stride_bm, stride_bn = block_bias.stride(0), block_bias.stride(2), block_bias.stride(1), block_bias.stride(3)
|
||||
else:
|
||||
raise ValueError(f"tensor_layout {tensor_layout} not supported")
|
||||
|
||||
stride_f_b, stride_f_h = flags.stride()
|
||||
|
||||
HEAD_DIM_K = head_dim
|
||||
num_kv_groups = h_qo // h_kv
|
||||
|
||||
grid = (triton.cdiv(qo_len, BLOCK_M), h_qo, b)
|
||||
_attn_fwd[grid](
|
||||
q, k, v, q_scale, k_scale, o,
|
||||
block_bias, decay_mask,
|
||||
flags,
|
||||
stride_f_b, stride_f_h,
|
||||
stride_bz_q, stride_h_q, stride_seq_q,
|
||||
stride_bz_k, stride_h_k, stride_seq_k,
|
||||
stride_bz_v, stride_h_v, stride_seq_v,
|
||||
stride_bz_o, stride_h_o, stride_seq_o,
|
||||
stride_bbz, stride_bbh, stride_bm, stride_bn,
|
||||
stride_dmz, stride_dmh, stride_dm, stride_dn,
|
||||
qo_len, kv_len,
|
||||
h_qo, num_kv_groups,
|
||||
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, HEAD_DIM=HEAD_DIM_K,
|
||||
STAGE=stage,
|
||||
num_warps=4 if head_dim == 64 else 8,
|
||||
num_stages=3 if head_dim == 64 else 4,
|
||||
xpos_xi=xpos_xi,
|
||||
frame_tokens=frame_tokens,
|
||||
sigmoid_a=sigmoid_a,
|
||||
alpha_xpos_xi=alpha_xpos_xi,
|
||||
beta_xpos_xi=beta_xpos_xi,
|
||||
sink_width=sink_width,
|
||||
window_width=window_width,
|
||||
multi_factor=multi_factor,
|
||||
entropy_factor=entropy_factor,
|
||||
)
|
||||
return o
|
||||
@@ -0,0 +1,64 @@
|
||||
# source https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/core.py
|
||||
|
||||
import torch
|
||||
import triton.language as tl
|
||||
|
||||
from .quant_per_block import per_block_int8
|
||||
from .attn_qk_int8_per_block import forward as attn_false
|
||||
|
||||
from typing import Optional
|
||||
|
||||
def sage_attention(
|
||||
qkv: list[torch.Tensor],
|
||||
tensor_layout: str ="HND",
|
||||
is_causal=False,
|
||||
sm_scale: Optional[float] = None,
|
||||
smooth_k: bool =True,
|
||||
xpos_xi: tl.constexpr = 0.9999934149894527,
|
||||
flags = None,
|
||||
block_bias = None,
|
||||
sigmoid_a: float = 1.0,
|
||||
alpha_xpos_xi: float = 0.97,
|
||||
beta_xpos_xi: float = 0.8,
|
||||
decay_mask = None,
|
||||
sink_width: int = 4,
|
||||
window_width: int = 21,
|
||||
multi_factor: Optional[float] = None,
|
||||
entropy_factor: Optional[float] = None,
|
||||
block_size : int = 64,
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
dtype = qkv[0].dtype
|
||||
q, k, v = qkv[0].transpose(1, 2), qkv[1].transpose(1, 2), qkv[2].transpose(1, 2) # to HND
|
||||
|
||||
if flags == None:
|
||||
flags = torch.zeros([q.shape[0],q.shape[1]], dtype=torch.int32, device=q.device)
|
||||
|
||||
seq_dim = 2
|
||||
|
||||
if smooth_k:
|
||||
km = k.mean(dim=seq_dim, keepdim=True)
|
||||
k -= km
|
||||
else:
|
||||
km = None
|
||||
|
||||
if dtype == torch.bfloat16 or dtype == torch.float32:
|
||||
v = v.to(torch.float16)
|
||||
|
||||
if q.dtype != k.dtype or q.dtype != v.dtype:
|
||||
k, v = k.to(q.dtype), v.to(q.dtype)
|
||||
|
||||
q_int8, q_scale, k_int8, k_scale = per_block_int8(q, k, sm_scale=sm_scale, tensor_layout=tensor_layout, BLKQ=block_size, BLKK=block_size)
|
||||
del q, k
|
||||
|
||||
o = attn_false(q_int8, k_int8, v, flags, block_bias, decay_mask, q_scale, k_scale,
|
||||
tensor_layout=tensor_layout, output_dtype=dtype, xpos_xi=xpos_xi, sigmoid_a=sigmoid_a,
|
||||
alpha_xpos_xi=alpha_xpos_xi, beta_xpos_xi=beta_xpos_xi,
|
||||
BLOCK_M=block_size, BLOCK_N=block_size,
|
||||
sink_width=sink_width,
|
||||
window_width=window_width,
|
||||
multi_factor=multi_factor,
|
||||
entropy_factor=entropy_factor,
|
||||
)
|
||||
|
||||
return o.transpose(1, 2).contiguous()
|
||||
@@ -0,0 +1,84 @@
|
||||
# https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/quant_per_block.py
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
@triton.jit
|
||||
def quant_per_block_int8_kernel(Input, Output, Scale, L,
|
||||
stride_iz, stride_ih, stride_in,
|
||||
stride_oz, stride_oh, stride_on,
|
||||
stride_sz, stride_sh,
|
||||
sm_scale,
|
||||
C: tl.constexpr, BLK: tl.constexpr):
|
||||
off_blk = tl.program_id(0)
|
||||
off_h = tl.program_id(1)
|
||||
off_b = tl.program_id(2)
|
||||
|
||||
offs_n = off_blk * BLK + tl.arange(0, BLK)
|
||||
offs_k = tl.arange(0, C)
|
||||
|
||||
input_ptrs = Input + off_b * stride_iz + off_h * stride_ih + offs_n[:, None] * stride_in + offs_k[None, :]
|
||||
output_ptrs = Output + off_b * stride_oz + off_h * stride_oh + offs_n[:, None] * stride_on + offs_k[None, :]
|
||||
scale_ptrs = Scale + off_b * stride_sz + off_h * stride_sh + off_blk
|
||||
|
||||
x = tl.load(input_ptrs, mask=offs_n[:, None] < L)
|
||||
x = x.to(tl.float32)
|
||||
x *= sm_scale
|
||||
scale = tl.max(tl.abs(x)) / 127.
|
||||
x_int8 = x / scale
|
||||
x_int8 += 0.5 * tl.where(x_int8 >= 0, 1, -1)
|
||||
x_int8 = x_int8.to(tl.int8)
|
||||
tl.store(output_ptrs, x_int8, mask=offs_n[:, None] < L)
|
||||
tl.store(scale_ptrs, scale)
|
||||
|
||||
def per_block_int8(q, k, BLKQ=128, BLKK=64, sm_scale=None, tensor_layout="HND"):
|
||||
q_int8 = torch.empty(q.shape, dtype=torch.int8, device=q.device)
|
||||
k_int8 = torch.empty(k.shape, dtype=torch.int8, device=k.device)
|
||||
|
||||
if tensor_layout == "HND":
|
||||
b, h_qo, qo_len, head_dim = q.shape
|
||||
_, h_kv, kv_len, _ = k.shape
|
||||
|
||||
stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(1), q.stride(2)
|
||||
stride_bz_qo, stride_h_qo, stride_seq_qo = q_int8.stride(0), q_int8.stride(1), q_int8.stride(2)
|
||||
stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(1), k.stride(2)
|
||||
stride_bz_ko, stride_h_ko, stride_seq_ko = k_int8.stride(0), k_int8.stride(1), k_int8.stride(2)
|
||||
# elif tensor_layout == "NHD":
|
||||
# b, qo_len, h_qo, head_dim = q.shape
|
||||
# _, kv_len, h_kv, _ = k.shape
|
||||
|
||||
# stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(2), q.stride(1)
|
||||
# stride_bz_qo, stride_h_qo, stride_seq_qo = q_int8.stride(0), q_int8.stride(2), q_int8.stride(1)
|
||||
# stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(2), k.stride(1)
|
||||
# stride_bz_ko, stride_h_ko, stride_seq_ko = k_int8.stride(0), k_int8.stride(2), k_int8.stride(1)
|
||||
else:
|
||||
raise ValueError(f"Unknown tensor layout: {tensor_layout}")
|
||||
|
||||
q_scale = torch.empty((b, h_qo, (qo_len + BLKQ - 1) // BLKQ, 1), device=q.device, dtype=torch.float32)
|
||||
k_scale = torch.empty((b, h_kv, (kv_len + BLKK - 1) // BLKK, 1), device=q.device, dtype=torch.float32)
|
||||
|
||||
if sm_scale is None:
|
||||
sm_scale = head_dim**-0.5
|
||||
|
||||
grid = ((qo_len + BLKQ - 1) // BLKQ, h_qo, b)
|
||||
quant_per_block_int8_kernel[grid](
|
||||
q, q_int8, q_scale, qo_len,
|
||||
stride_bz_q, stride_h_q, stride_seq_q,
|
||||
stride_bz_qo, stride_h_qo, stride_seq_qo,
|
||||
q_scale.stride(0), q_scale.stride(1),
|
||||
sm_scale=(sm_scale * 1.44269504),
|
||||
C=head_dim, BLK=BLKQ
|
||||
)
|
||||
|
||||
grid = ((kv_len + BLKK - 1) // BLKK, h_kv, b)
|
||||
quant_per_block_int8_kernel[grid](
|
||||
k, k_int8, k_scale, kv_len,
|
||||
stride_bz_k, stride_h_k, stride_seq_k,
|
||||
stride_bz_ko, stride_h_ko, stride_seq_ko,
|
||||
k_scale.stride(0), k_scale.stride(1),
|
||||
sm_scale=1.0,
|
||||
C=head_dim, BLK=BLKK
|
||||
)
|
||||
|
||||
return q_int8, q_scale, k_int8, k_scale
|
||||
+36
-183
@@ -1,21 +1,13 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import torch
|
||||
from ...utils import log
|
||||
|
||||
# Flash Attention imports
|
||||
try:
|
||||
import flash_attn_interface
|
||||
FLASH_ATTN_3_AVAILABLE = True
|
||||
except Exception as e:
|
||||
FLASH_ATTN_3_AVAILABLE = False
|
||||
def attention_func_error(*args, **kwargs):
|
||||
raise ImportError("Selected attention mode not available. Please ensure required packages are installed correctly.")
|
||||
|
||||
try:
|
||||
import flash_attn
|
||||
FLASH_ATTN_2_AVAILABLE = True
|
||||
except Exception as e:
|
||||
FLASH_ATTN_2_AVAILABLE = False
|
||||
from .attention_flash import flash_attention
|
||||
|
||||
# Sage Attention imports
|
||||
# using custom ops to avoid graph breaks with torch.compile
|
||||
try:
|
||||
from sageattention import sageattn
|
||||
|
||||
@@ -69,188 +61,49 @@ try:
|
||||
# Return tensor with same shape as q
|
||||
return q.clone()
|
||||
except Exception as e:
|
||||
sageattn_varlen_func = None
|
||||
sageattn_varlen_func = attention_func_error
|
||||
|
||||
# sage3
|
||||
try:
|
||||
from sageattn3 import sageattn3_blackwell as sageattn_blackwell
|
||||
except:
|
||||
try:
|
||||
from sageattn import sageattn_blackwell
|
||||
except:
|
||||
SAGE3_AVAILABLE = False
|
||||
sageattn_blackwell = attention_func_error
|
||||
|
||||
try:
|
||||
from ...ultravico.sageattn.core import sage_attention as sageattn_ultravico
|
||||
@torch.library.custom_op("wanvideo::sageattn_ultravico", mutates_args=())
|
||||
def sageattn_func_ultravico(qkv: List[torch.Tensor], attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, multi_factor: float = 0.9
|
||||
) -> torch.Tensor:
|
||||
return sageattn_ultravico(qkv, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, multi_factor=multi_factor)
|
||||
|
||||
@sageattn_func_ultravico.register_fake
|
||||
def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9):
|
||||
torch.empty_like(qkv[0]).contiguous()
|
||||
except:
|
||||
sageattn_ultravico = attention_func_error
|
||||
|
||||
|
||||
__all__ = [
|
||||
'flash_attention',
|
||||
'attention',
|
||||
]
|
||||
|
||||
|
||||
def flash_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
q_lens=None,
|
||||
k_lens=None,
|
||||
dropout_p=0.,
|
||||
softmax_scale=None,
|
||||
q_scale=None,
|
||||
causal=False,
|
||||
window_size=(-1, -1),
|
||||
deterministic=False,
|
||||
dtype=torch.bfloat16,
|
||||
version=None,
|
||||
):
|
||||
"""
|
||||
q: [B, Lq, Nq, C1].
|
||||
k: [B, Lk, Nk, C1].
|
||||
v: [B, Lk, Nk, C2]. Nq must be divisible by Nk.
|
||||
q_lens: [B].
|
||||
k_lens: [B].
|
||||
dropout_p: float. Dropout probability.
|
||||
softmax_scale: float. The scaling of QK^T before applying softmax.
|
||||
causal: bool. Whether to apply causal attention mask.
|
||||
window_size: (left right). If not (-1, -1), apply sliding window local attention.
|
||||
deterministic: bool. If True, slightly slower and uses more memory.
|
||||
dtype: torch.dtype. Apply when dtype of q/k/v is not float16/bfloat16.
|
||||
"""
|
||||
half_dtypes = (torch.float16, torch.bfloat16)
|
||||
#assert dtype in half_dtypes
|
||||
#assert q.device.type == 'cuda' and q.size(-1) <= 256
|
||||
|
||||
# params
|
||||
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
|
||||
|
||||
def half(x):
|
||||
return x if x.dtype in half_dtypes else x.to(dtype)
|
||||
|
||||
# preprocess query
|
||||
if q_lens is None:
|
||||
q = half(q.flatten(0, 1))
|
||||
q_lens = torch.tensor(
|
||||
[lq] * b, dtype=torch.int32).to(
|
||||
device=q.device, non_blocking=True)
|
||||
else:
|
||||
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
|
||||
|
||||
# preprocess key, value
|
||||
if k_lens is None:
|
||||
k = half(k.flatten(0, 1))
|
||||
v = half(v.flatten(0, 1))
|
||||
k_lens = torch.tensor(
|
||||
[lk] * b, dtype=torch.int32).to(
|
||||
device=k.device, non_blocking=True)
|
||||
else:
|
||||
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
|
||||
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
|
||||
|
||||
q = q.to(v.dtype)
|
||||
k = k.to(v.dtype)
|
||||
|
||||
if q_scale is not None:
|
||||
q = q * q_scale
|
||||
|
||||
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
|
||||
log.warning('Flash attention 3 is not available, use flash attention 2 instead.')
|
||||
|
||||
# apply attention
|
||||
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
|
||||
# Note: dropout_p, window_size are not supported in FA3 now.
|
||||
x = flash_attn_interface.flash_attn_varlen_func(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
seqused_q=None,
|
||||
seqused_k=None,
|
||||
max_seqlen_q=lq,
|
||||
max_seqlen_k=lk,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic)[0].unflatten(0, (b, lq))
|
||||
else:
|
||||
assert FLASH_ATTN_2_AVAILABLE
|
||||
x = flash_attn.flash_attn_varlen_func(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
max_seqlen_q=lq,
|
||||
max_seqlen_k=lk,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
window_size=window_size,
|
||||
deterministic=deterministic).unflatten(0, (b, lq))
|
||||
|
||||
# output
|
||||
return x.type(out_dtype)
|
||||
|
||||
|
||||
def attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
q_lens=None,
|
||||
k_lens=None,
|
||||
max_seqlen_q=None,
|
||||
max_seqlen_k=None,
|
||||
dropout_p=0.,
|
||||
softmax_scale=None,
|
||||
q_scale=None,
|
||||
causal=False,
|
||||
window_size=(-1, -1),
|
||||
deterministic=False,
|
||||
dtype=torch.bfloat16,
|
||||
attention_mode='sdpa',
|
||||
attn_mask=None,
|
||||
):
|
||||
def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0.,
|
||||
softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16,
|
||||
attention_mode='sdpa', attn_mask=None, multi_factor=0.9):
|
||||
if "flash" in attention_mode:
|
||||
if attention_mode == 'flash_attn_2':
|
||||
fa_version = 2
|
||||
elif attention_mode == 'flash_attn_3':
|
||||
fa_version = 3
|
||||
return flash_attention(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
q_lens=q_lens,
|
||||
k_lens=k_lens,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
q_scale=q_scale,
|
||||
causal=causal,
|
||||
window_size=window_size,
|
||||
deterministic=deterministic,
|
||||
dtype=dtype,
|
||||
version=fa_version,
|
||||
return flash_attention(q, k, v, q_lens=q_lens, k_lens=k_lens, dropout_p=dropout_p, softmax_scale=softmax_scale,
|
||||
q_scale=q_scale, causal=causal, window_size=window_size, deterministic=deterministic, dtype=dtype, version=2 if attention_mode == 'flash_attn_2' else 3,
|
||||
)
|
||||
elif attention_mode == 'sdpa':
|
||||
elif attention_mode == 'sageattn_3':
|
||||
return sageattn_blackwell(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), per_block_mean=False).transpose(1,2).contiguous()
|
||||
elif attention_mode == 'sageattn_varlen':
|
||||
return torch.ops.wanvideo.sageattn_varlen(q,k,v, q_lens=q_lens, k_lens=k_lens, max_seqlen_k=max_seqlen_k, max_seqlen_q=max_seqlen_q)
|
||||
elif attention_mode == 'sageattn_compiled': # for sage versions that allow torch.compile, may be redundant now as other sageattn ops are wrapper in custom ops
|
||||
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
|
||||
elif attention_mode == 'sageattn':
|
||||
return torch.ops.wanvideo.sageattn(q, k, v, tensor_layout="NHD").contiguous()
|
||||
elif attention_mode == 'sageattn_ultravico':
|
||||
return torch.ops.wanvideo.sageattn_ultravico([q, k, v], multi_factor=multi_factor).contiguous()
|
||||
else: # sdpa
|
||||
if not (q.dtype == k.dtype == v.dtype):
|
||||
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2).to(q.dtype), v.transpose(1, 2).to(q.dtype), attn_mask=attn_mask).transpose(1, 2).contiguous()
|
||||
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), attn_mask=attn_mask).transpose(1, 2).contiguous()
|
||||
elif attention_mode == 'sageattn_3':
|
||||
return sageattn_blackwell(
|
||||
q.transpose(1,2),
|
||||
k.transpose(1,2),
|
||||
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 torch.ops.wanvideo.sageattn_varlen(
|
||||
q,k,v,
|
||||
q_lens=q_lens,
|
||||
k_lens=k_lens,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
max_seqlen_q=max_seqlen_q
|
||||
)
|
||||
elif attention_mode == 'sageattn_compiled':
|
||||
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
|
||||
else:
|
||||
return torch.ops.wanvideo.sageattn(q, k, v, tensor_layout="NHD").contiguous()
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
import torch
|
||||
from ...utils import log
|
||||
|
||||
def attention_func_error(*args, **kwargs):
|
||||
raise ImportError("Selected attention mode not available. Please ensure required packages are installed correctly.")
|
||||
|
||||
try:
|
||||
import flash_attn_interface
|
||||
FLASH_ATTN_3_AVAILABLE = True
|
||||
except Exception as e:
|
||||
FLASH_ATTN_3_AVAILABLE = False
|
||||
|
||||
try:
|
||||
import flash_attn
|
||||
FLASH_ATTN_2_AVAILABLE = True
|
||||
except Exception as e:
|
||||
FLASH_ATTN_2_AVAILABLE = False
|
||||
|
||||
if not FLASH_ATTN_2_AVAILABLE and not FLASH_ATTN_3_AVAILABLE:
|
||||
flash_attention = attention_func_error
|
||||
else:
|
||||
def flash_attention(q, k, v, q_lens=None, k_lens=None, dropout_p=0., softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16, version=None):
|
||||
half_dtypes = (torch.float16, torch.bfloat16)
|
||||
|
||||
# params
|
||||
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
|
||||
|
||||
def half(x):
|
||||
return x if x.dtype in half_dtypes else x.to(dtype)
|
||||
|
||||
# preprocess query
|
||||
if q_lens is None:
|
||||
q = half(q.flatten(0, 1))
|
||||
q_lens = torch.tensor(
|
||||
[lq] * b, dtype=torch.int32).to(
|
||||
device=q.device, non_blocking=True)
|
||||
else:
|
||||
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
|
||||
|
||||
# preprocess key, value
|
||||
if k_lens is None:
|
||||
k = half(k.flatten(0, 1))
|
||||
v = half(v.flatten(0, 1))
|
||||
k_lens = torch.tensor(
|
||||
[lk] * b, dtype=torch.int32).to(
|
||||
device=k.device, non_blocking=True)
|
||||
else:
|
||||
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
|
||||
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
|
||||
|
||||
q = q.to(v.dtype)
|
||||
k = k.to(v.dtype)
|
||||
|
||||
if q_scale is not None:
|
||||
q = q * q_scale
|
||||
|
||||
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
|
||||
log.warning('Flash attention 3 is not available, use flash attention 2 instead.')
|
||||
|
||||
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
|
||||
# Note: dropout_p, window_size are not supported in FA3 now.
|
||||
x = flash_attn_interface.flash_attn_varlen_func(q=q, k=k, v=v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
seqused_q=None, seqused_k=None, max_seqlen_q=lq, max_seqlen_k=lk,
|
||||
softmax_scale=softmax_scale, causal=causal,
|
||||
deterministic=deterministic)[0].unflatten(0, (b, lq))
|
||||
else:
|
||||
assert FLASH_ATTN_2_AVAILABLE
|
||||
x = flash_attn.flash_attn_varlen_func(q=q, k=k, v=v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
max_seqlen_q=lq, max_seqlen_k=lk, dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale, causal=causal, window_size=window_size,
|
||||
deterministic=deterministic).unflatten(0, (b, lq))
|
||||
return x.type(out_dtype)
|
||||
Reference in New Issue
Block a user