[Diffusion][CPU] Adding AMX optimizations for CPU platform (#28527)
Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
@@ -23,6 +23,7 @@ def flash_attn_varlen_ref(
|
||||
cu_seqlens_k,
|
||||
is_causal,
|
||||
enable_gqa,
|
||||
softmax_scale=None,
|
||||
):
|
||||
cu_q = cu_seqlens_q.tolist()
|
||||
cu_k = cu_seqlens_k.tolist()
|
||||
@@ -43,6 +44,7 @@ def flash_attn_varlen_ref(
|
||||
v[:, :, start_k:end_k, :],
|
||||
is_causal=is_causal,
|
||||
enable_gqa=enable_gqa,
|
||||
scale=softmax_scale,
|
||||
)
|
||||
|
||||
# [1, H, T, D] -> [T, H, D]
|
||||
@@ -58,6 +60,7 @@ def flash_attn_non_varlen_ref(
|
||||
cu_seqlens_k,
|
||||
is_causal,
|
||||
enable_gqa,
|
||||
softmax_scale=None,
|
||||
):
|
||||
cu_q = cu_seqlens_q.tolist()
|
||||
cu_k = cu_seqlens_k.tolist()
|
||||
@@ -75,6 +78,7 @@ def flash_attn_non_varlen_ref(
|
||||
v,
|
||||
is_causal=is_causal,
|
||||
enable_gqa=enable_gqa,
|
||||
scale=softmax_scale,
|
||||
)
|
||||
# [B, H, T, D] -> [B * T, H, D]
|
||||
return out.transpose(1, 2).reshape(batch * T, H, D)
|
||||
@@ -91,6 +95,7 @@ class TestFlashAttn(CustomTestCase):
|
||||
head_dim=[32, 48], # test when D is not 32x
|
||||
head_dim_v=[32],
|
||||
is_causal=[True, False],
|
||||
softmax_scale=[None, 0.2],
|
||||
)
|
||||
def test_flash_attn_varlen(
|
||||
self,
|
||||
@@ -102,6 +107,7 @@ class TestFlashAttn(CustomTestCase):
|
||||
head_dim,
|
||||
head_dim_v,
|
||||
is_causal,
|
||||
softmax_scale,
|
||||
):
|
||||
dtype = torch.bfloat16
|
||||
|
||||
@@ -127,6 +133,7 @@ class TestFlashAttn(CustomTestCase):
|
||||
cu_seqlens_k,
|
||||
is_causal=is_causal,
|
||||
enable_gqa=num_heads != num_heads_kv,
|
||||
softmax_scale=softmax_scale,
|
||||
)
|
||||
|
||||
out = flash_attn_varlen_func(
|
||||
@@ -138,6 +145,7 @@ class TestFlashAttn(CustomTestCase):
|
||||
seqlens_q.max().item(),
|
||||
seqlens_k.max().item(),
|
||||
is_causal,
|
||||
softmax_scale,
|
||||
)
|
||||
|
||||
atol = rtol = precision[dtype]
|
||||
@@ -153,6 +161,7 @@ class TestFlashAttn(CustomTestCase):
|
||||
head_dim=[32],
|
||||
head_dim_v=[32],
|
||||
is_causal=[False],
|
||||
softmax_scale=[None, 0.2],
|
||||
)
|
||||
def test_flash_attn_large_size(
|
||||
self,
|
||||
@@ -164,6 +173,7 @@ class TestFlashAttn(CustomTestCase):
|
||||
head_dim,
|
||||
head_dim_v,
|
||||
is_causal,
|
||||
softmax_scale,
|
||||
):
|
||||
dtype = torch.bfloat16
|
||||
|
||||
@@ -190,6 +200,7 @@ class TestFlashAttn(CustomTestCase):
|
||||
cu_seqlens_k,
|
||||
is_causal=is_causal,
|
||||
enable_gqa=num_heads != num_heads_kv,
|
||||
softmax_scale=softmax_scale,
|
||||
)
|
||||
|
||||
out = flash_attn_varlen_func(
|
||||
@@ -201,6 +212,7 @@ class TestFlashAttn(CustomTestCase):
|
||||
seqlens_q.max().item(),
|
||||
seqlens_k.max().item(),
|
||||
is_causal,
|
||||
softmax_scale,
|
||||
)
|
||||
|
||||
atol = rtol = precision[dtype]
|
||||
@@ -226,7 +238,7 @@ class TestFlashAttn(CustomTestCase):
|
||||
q, k, v, cu_seqlens, cu_seqlens, is_causal=True, enable_gqa=True
|
||||
)
|
||||
out = flash_attn_varlen_func(
|
||||
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, True
|
||||
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, True, None
|
||||
)
|
||||
|
||||
atol = rtol = precision[dtype]
|
||||
|
||||
Reference in New Issue
Block a user