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:
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
@@ -98,12 +98,12 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> 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<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;
|
||||
}
|
||||
|
||||
|
||||
@@ -53,7 +53,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> 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<bf16> p) {
|
||||
for (int j = 0; j < 8; j++)
|
||||
part = fmaf(qreg[i + j], k8[j], part);
|
||||
}
|
||||
float dot = group_reduce_sum<G>(part, gmask);
|
||||
float dot = group_reduce_sum<G>(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<bf16> 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);
|
||||
|
||||
Reference in New Issue
Block a user