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)
This commit is contained in:
2026-07-21 23:05:17 +08:00
parent f7a16efc9d
commit 60d7ee614a
6 changed files with 57 additions and 32 deletions
+6 -6
View File
@@ -84,8 +84,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> 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<bf16> 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;
}