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
This commit is contained in:
2026-07-31 00:19:18 +08:00
parent 02625739fe
commit 28d1bd07cf
2 changed files with 8 additions and 8 deletions
+4 -4
View File
@@ -70,8 +70,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
} }
float new_m = fmaxf(m, partial); float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m); float alpha = __expf(m - new_m);
float beta = expf(partial - new_m); float beta = __expf(partial - new_m);
d = d * alpha + beta; d = d * alpha + beta;
int v_off = kv_base + kv_idx * p.kv_stride_l int v_off = kv_base + kv_idx * p.kv_stride_l
@@ -116,8 +116,8 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
if (mi <= -FLT_MAX) continue; if (mi <= -FLT_MAX) continue;
float li = mlp[s * 2 + 1]; float li = mlp[s * 2 + 1];
float nm = fmaxf(m, mi); float nm = fmaxf(m, mi);
float corr = expf(m - nm); float corr = __expf(m - nm);
float e = expf(mi - nm); float e = __expf(mi - nm);
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e); acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
l = fmaf(l, corr, li * e); l = fmaf(l, corr, li * e);
m = nm; m = nm;
+4 -4
View File
@@ -77,8 +77,8 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
} }
float new_m = fmaxf(m, partial); float new_m = fmaxf(m, partial);
float alpha = expf(m - new_m); float alpha = __expf(m - new_m);
float beta = expf(partial - new_m); float beta = __expf(partial - new_m);
d = d * alpha + beta; d = d * alpha + beta;
int pos = chunk_start + s; int pos = chunk_start + s;
@@ -133,8 +133,8 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
if (mi <= -FLT_MAX) continue; if (mi <= -FLT_MAX) continue;
float li = mlp[s * 2 + 1]; float li = mlp[s * 2 + 1];
float nm = fmaxf(m, mi); float nm = fmaxf(m, mi);
float corr = expf(m - nm); float corr = __expf(m - nm);
float e = expf(mi - nm); float e = __expf(mi - nm);
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e); acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
l = fmaf(l, corr, li * e); l = fmaf(l, corr, li * e);
m = nm; m = nm;