From 28d1bd07cf291c5de68dd4a3463f4b3c76b86b6b Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Fri, 31 Jul 2026 00:19:18 +0800 Subject: [PATCH] style: unify decode expf to __expf - attn_decode_split_kv.cuh: 4 expf -> __expf - attn_paged_decode_split_kv.cuh: 4 expf -> __expf - --use_fast_math makes expf emit __expf anyway, so no behavior change - aligns decode with prefill/mma kernels that already use __expf --- csrc/kernels/attn_decode_split_kv.cuh | 8 ++++---- csrc/kernels/attn_paged_decode_split_kv.cuh | 8 ++++---- 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/csrc/kernels/attn_decode_split_kv.cuh b/csrc/kernels/attn_decode_split_kv.cuh index 30ab4e5..a5ad21d 100644 --- a/csrc/kernels/attn_decode_split_kv.cuh +++ b/csrc/kernels/attn_decode_split_kv.cuh @@ -70,8 +70,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { } float new_m = fmaxf(m, partial); - float alpha = expf(m - new_m); - float beta = expf(partial - new_m); + float alpha = __expf(m - new_m); + float beta = __expf(partial - new_m); d = d * alpha + beta; int v_off = kv_base + kv_idx * p.kv_stride_l @@ -116,8 +116,8 @@ __global__ void attn_decode_combine_kernel(AttentionParams p) { if (mi <= -FLT_MAX) continue; float li = mlp[s * 2 + 1]; float nm = fmaxf(m, mi); - float corr = expf(m - nm); - float e = expf(mi - nm); + float corr = __expf(m - nm); + float e = __expf(mi - nm); acc = fmaf(acc, corr, op[s * p.head_dim + d] * e); l = fmaf(l, corr, li * e); m = nm; diff --git a/csrc/kernels/attn_paged_decode_split_kv.cuh b/csrc/kernels/attn_paged_decode_split_kv.cuh index 42268f2..b3ee857 100644 --- a/csrc/kernels/attn_paged_decode_split_kv.cuh +++ b/csrc/kernels/attn_paged_decode_split_kv.cuh @@ -77,8 +77,8 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams p) } float new_m = fmaxf(m, partial); - float alpha = expf(m - new_m); - float beta = expf(partial - new_m); + float alpha = __expf(m - new_m); + float beta = __expf(partial - new_m); d = d * alpha + beta; int pos = chunk_start + s; @@ -133,8 +133,8 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams p) { if (mi <= -FLT_MAX) continue; float li = mlp[s * 2 + 1]; float nm = fmaxf(m, mi); - float corr = expf(m - nm); - float e = expf(mi - nm); + float corr = __expf(m - nm); + float e = __expf(mi - nm); acc = fmaf(acc, corr, op[s * p.head_dim + d] * e); l = fmaf(l, corr, li * e); m = nm;