[Diffusion][CPU] Adding AMX optimizations for CPU platform (#28527)

Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
jianan-gu
2026-07-09 10:26:22 +08:00
committed by GitHub
co-authored by Ma Mingfei
parent 395a2201e4
commit 177c048c68
12 changed files with 166 additions and 13 deletions
+5 -4
View File
@@ -443,7 +443,8 @@ at::Tensor flash_attn_varlen_func(
const at::Tensor& cu_seqlens_k,
int64_t max_seqlen_q,
int64_t max_seqlen_k,
bool causal) {
bool causal,
const std::optional<double>& sm_scale) {
CHECK_LAST_DIM_CONTIGUOUS_INPUT(q);
CHECK_LAST_DIM_CONTIGUOUS_INPUT(k);
CHECK_LAST_DIM_CONTIGUOUS_INPUT(v);
@@ -480,7 +481,7 @@ at::Tensor flash_attn_varlen_func(
TORCH_CHECK(head_size_v % 2 == 0, "invalid head_size_v ", head_size_v);
// softmax scale
double sm_scale = 1.0 / std::sqrt(static_cast<double>(head_size));
double _sm_scale = sm_scale.has_value() ? sm_scale.value() : 1.0 / std::sqrt(static_cast<double>(head_size));
// check whether the batch has variant lengths
const bool is_varlen =
@@ -522,7 +523,7 @@ at::Tensor flash_attn_varlen_func(
k_strideH,
v_strideN,
v_strideH,
sm_scale,
_sm_scale,
sz,
causal);
} else {
@@ -545,7 +546,7 @@ at::Tensor flash_attn_varlen_func(
k_strideH,
v_strideN,
v_strideH,
sm_scale,
_sm_scale,
sz,
causal);
}
+3 -2
View File
@@ -153,7 +153,8 @@ at::Tensor flash_attn_varlen_func(
const at::Tensor& cu_seqlens_k,
int64_t max_seqlen_q,
int64_t max_seqlen_k,
bool causal);
bool causal,
const std::optional<double>& sm_scale);
// linear attention
std::tuple<at::Tensor, at::Tensor> chunk_gated_delta_rule_cpu(
@@ -533,7 +534,7 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
// flash attn
m.def(
"flash_attn_varlen_func(Tensor q, Tensor k, Tensor v, Tensor cu_seqlens_q, Tensor cu_seqlens_k, "
"int max_seqlen_q, int max_seqlen_k, bool causal) -> Tensor");
"int max_seqlen_q, int max_seqlen_k, bool causal, float? sm_scale) -> Tensor");
m.impl("flash_attn_varlen_func", torch::kCPU, &flash_attn_varlen_func);
// linear attn