From 60d7ee614a1980ea96416b9af907db0ed8b4495b Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Tue, 21 Jul 2026 23:05:17 +0800 Subject: [PATCH] fix: improve attention kernel numerical stability and test precision checks - use fmaf() for V-accumulation in scalar decode paths to reduce rounding - delay scale multiplication to after dot-product in scalar prefill - unify __expf/expf across MMA and scalar paths for consistent numerics - harmonize divide-by-zero guards to 1e-20f - add both absolute and relative error checks in standalone tests (atol=0.01, rtol=0.01) --- csrc/kernels/attn_decode_split_kv.cuh | 12 ++++++------ csrc/kernels/attn_paged_decode_split_kv.cuh | 14 +++++++------- csrc/kernels/attn_prefill_split_q.cuh | 6 +++--- csrc/tests/attn_decode_test.cu | 19 ++++++++++++++----- csrc/tests/attn_paged_decode_test.cu | 19 +++++++++++++------ csrc/tests/attn_prefill_test.cu | 19 ++++++++++++++----- 6 files changed, 57 insertions(+), 32 deletions(-) diff --git a/csrc/kernels/attn_decode_split_kv.cuh b/csrc/kernels/attn_decode_split_kv.cuh index 6fc39dc..67b0618 100644 --- a/csrc/kernels/attn_decode_split_kv.cuh +++ b/csrc/kernels/attn_decode_split_kv.cuh @@ -84,8 +84,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams p) { int v_off = kv_base + kv_idx * p.kv_stride_l + lane * hd_per_thread * p.kv_stride_d; for (int i = 0; i < hd_per_thread; i++) - acc_reg[i] = acc_reg[i] * alpha - + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta; + acc_reg[i] = fmaf(acc_reg[i], alpha, + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta); m = new_m; } __syncthreads(); @@ -123,10 +123,10 @@ __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); - acc = acc * corr + op[s * p.head_dim + d] * e; - l = l * corr + li * e; + 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 14492cc..62c6a65 100644 --- a/csrc/kernels/attn_paged_decode_split_kv.cuh +++ b/csrc/kernels/attn_paged_decode_split_kv.cuh @@ -98,12 +98,12 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams p) + (int64_t)kv_head * p.head_dim; #pragma unroll for (int i = 0; i < hd_per_thread; i++) - acc_reg[i] = acc_reg[i] * alpha - + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta; + acc_reg[i] = fmaf(acc_reg[i], alpha, + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta); } else { #pragma unroll for (int i = 0; i < hd_per_thread; i++) - acc_reg[i] = acc_reg[i] * alpha + 0.0f * beta; + acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f); } m = new_m; } @@ -140,10 +140,10 @@ __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); - acc = acc * corr + op[s * p.head_dim + d] * e; - l = l * corr + li * e; + 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_prefill_split_q.cuh b/csrc/kernels/attn_prefill_split_q.cuh index 8f43304..edbef5b 100644 --- a/csrc/kernels/attn_prefill_split_q.cuh +++ b/csrc/kernels/attn_prefill_split_q.cuh @@ -53,7 +53,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d; #pragma unroll for (int i = 0; i < DPT; i++) - qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]) * p.scale; + qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]); } float m = -FLT_MAX, l = 0.0f; @@ -111,7 +111,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { for (int j = 0; j < 8; j++) part = fmaf(qreg[i + j], k8[j], part); } - float dot = group_reduce_sum(part, gmask); + float dot = group_reduce_sum(part, gmask) * p.scale; int kv_idx = kv0 + s; if constexpr (HasMask) { @@ -141,7 +141,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams p) { if (q_row < p.q_len) { int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + q_row * p.q_stride_l + gpos * DPT * p.q_stride_d; - float rl = (l > 1e-10f) ? (1.0f / l) : 0.0f; + float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f; #pragma unroll for (int i = 0; i < DPT; i++) p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl); diff --git a/csrc/tests/attn_decode_test.cu b/csrc/tests/attn_decode_test.cu index 327f289..e31a4cb 100644 --- a/csrc/tests/attn_decode_test.cu +++ b/csrc/tests/attn_decode_test.cu @@ -135,18 +135,27 @@ static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) { float* ref=new float[nQ]; cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1); - float max_err=0; + float max_abs_err=0, max_rel_err=0; for (size_t i=0;imax_err) max_err=d; + float err=fabsf(bf2f(hOut[i])-ref[i]); + if(err>max_abs_err) max_abs_err=err; + float rel=err/fmaxf(fabsf(ref[i]), 1e-8f); + if(rel>max_rel_err) max_rel_err=rel; } - printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err); + const float atol=0.01f, rtol=0.01f; + bool pass=true; + for (size_t i=0;i atol + rtol * fabsf(ref[i])) { pass=false; break; } + } + printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n", + kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL"); cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask); free_scratch(sc); delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp; - return (max_err < 0.05f) ? 0 : 1; + return pass ? 0 : 1; } int main() { diff --git a/csrc/tests/attn_paged_decode_test.cu b/csrc/tests/attn_paged_decode_test.cu index 5a97559..7794f79 100644 --- a/csrc/tests/attn_paged_decode_test.cu +++ b/csrc/tests/attn_paged_decode_test.cu @@ -135,23 +135,30 @@ static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causa for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_paged[i] = __bfloat162float(h_o_bf16[i]); - float max_err = 0.0f; + float max_abs_err = 0.0f, max_rel_err = 0.0f; int bad_idx = -1; for (int i = 0; i < B * Hq * HEAD_DIM; i++) { float e = fabsf(h_o_paged[i] - h_o_ref[i]); - if (e > max_err) { max_err = e; bad_idx = i; } + if (e > max_abs_err) { max_abs_err = e; bad_idx = i; } + float rel = e / fmaxf(fabsf(h_o_ref[i]), 1e-8f); + if (rel > max_rel_err) max_rel_err = rel; } - bool pass = max_err < 0.02f; + const float atol = 0.01f, rtol = 0.01f; + bool pass = true; + for (int i = 0; i < B * Hq * HEAD_DIM; i++) { + float e = fabsf(h_o_paged[i] - h_o_ref[i]); + if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; } + } if (pass) { - printf("PASS (max_abs_err=%.4e)\n", max_err); + printf("PASS (max_abs_err=%.4e max_rel_err=%.4e)\n", max_abs_err, max_rel_err); } else { int b = bad_idx / (Hq * HEAD_DIM); int h = (bad_idx / HEAD_DIM) % Hq; int d = bad_idx % HEAD_DIM; - printf("FAIL (max_abs_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n", - max_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]); + printf("FAIL (max_abs_err=%.4e max_rel_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n", + max_abs_err, max_rel_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]); printf(" ref[0..7]:"); for (int i = 0; i < 8 && i < HEAD_DIM; i++) printf(" %.4f", h_o_ref[i]); diff --git a/csrc/tests/attn_prefill_test.cu b/csrc/tests/attn_prefill_test.cu index 78aebc4..ec137b4 100644 --- a/csrc/tests/attn_prefill_test.cu +++ b/csrc/tests/attn_prefill_test.cu @@ -119,17 +119,26 @@ static int run_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) { float* ref=new float[nQ]; cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1); - float max_err=0; + float max_abs_err=0, max_rel_err=0; for (size_t i=0;imax_err) max_err=d; + float err=fabsf(bf2f(hOut[i])-ref[i]); + if(err>max_abs_err) max_abs_err=err; + float rel=err/fmaxf(fabsf(ref[i]), 1e-8f); + if(rel>max_rel_err) max_rel_err=rel; } - printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err); + const float atol=0.01f, rtol=0.01f; + bool pass=true; + for (size_t i=0;i atol + rtol * fabsf(ref[i])) { pass=false; break; } + } + printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n", + kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL"); cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO); delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp; - return (max_err < 0.05f) ? 0 : 1; + return pass ? 0 : 1; } int main() {