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
|
int v_off = kv_base + kv_idx * p.kv_stride_l
|
||||||
+ lane * hd_per_thread * p.kv_stride_d;
|
+ lane * hd_per_thread * p.kv_stride_d;
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
acc_reg[i] = acc_reg[i] * alpha
|
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
||||||
+ __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta;
|
__bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta);
|
||||||
m = new_m;
|
m = new_m;
|
||||||
}
|
}
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
@@ -123,10 +123,10 @@ __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 = acc * corr + op[s * p.head_dim + d] * e;
|
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
||||||
l = l * corr + li * e;
|
l = fmaf(l, corr, li * e);
|
||||||
m = nm;
|
m = nm;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -98,12 +98,12 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
|||||||
+ (int64_t)kv_head * p.head_dim;
|
+ (int64_t)kv_head * p.head_dim;
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
acc_reg[i] = acc_reg[i] * alpha
|
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
||||||
+ __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta;
|
__bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta);
|
||||||
} else {
|
} else {
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
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;
|
m = new_m;
|
||||||
}
|
}
|
||||||
@@ -140,10 +140,10 @@ __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 = acc * corr + op[s * p.head_dim + d] * e;
|
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
||||||
l = l * corr + li * e;
|
l = fmaf(l, corr, li * e);
|
||||||
m = nm;
|
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;
|
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < DPT; i++)
|
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;
|
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++)
|
for (int j = 0; j < 8; j++)
|
||||||
part = fmaf(qreg[i + j], k8[j], part);
|
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;
|
int kv_idx = kv0 + s;
|
||||||
if constexpr (HasMask) {
|
if constexpr (HasMask) {
|
||||||
@@ -141,7 +141,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
|||||||
if (q_row < p.q_len) {
|
if (q_row < p.q_len) {
|
||||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
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;
|
+ 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
|
#pragma unroll
|
||||||
for (int i = 0; i < DPT; i++)
|
for (int i = 0; i < DPT; i++)
|
||||||
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
|
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
|
||||||
|
|||||||
@@ -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];
|
float* ref=new float[nQ];
|
||||||
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1);
|
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;i<nQ;i++){
|
for (size_t i=0;i<nQ;i++){
|
||||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
if(d>max_err) max_err=d;
|
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<nQ;i++){
|
||||||
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if (err > 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);
|
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
|
||||||
free_scratch(sc);
|
free_scratch(sc);
|
||||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
|
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() {
|
int main() {
|
||||||
|
|||||||
@@ -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++)
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
||||||
h_o_paged[i] = __bfloat162float(h_o_bf16[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;
|
int bad_idx = -1;
|
||||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
||||||
float e = fabsf(h_o_paged[i] - h_o_ref[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) {
|
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 {
|
} else {
|
||||||
int b = bad_idx / (Hq * HEAD_DIM);
|
int b = bad_idx / (Hq * HEAD_DIM);
|
||||||
int h = (bad_idx / HEAD_DIM) % Hq;
|
int h = (bad_idx / HEAD_DIM) % Hq;
|
||||||
int d = bad_idx % HEAD_DIM;
|
int d = bad_idx % HEAD_DIM;
|
||||||
printf("FAIL (max_abs_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
|
printf("FAIL (max_abs_err=%.4e max_rel_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]);
|
max_abs_err, max_rel_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
|
||||||
printf(" ref[0..7]:");
|
printf(" ref[0..7]:");
|
||||||
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
||||||
printf(" %.4f", h_o_ref[i]);
|
printf(" %.4f", h_o_ref[i]);
|
||||||
|
|||||||
@@ -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];
|
float* ref=new float[nQ];
|
||||||
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
|
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;i<nQ;i++) {
|
for (size_t i=0;i<nQ;i++) {
|
||||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
if(d>max_err) max_err=d;
|
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<nQ;i++) {
|
||||||
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if (err > 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(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
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() {
|
int main() {
|
||||||
|
|||||||
Reference in New Issue
Block a user