fix: zero-init AttentionParams in pure C tests

- paged decode test left new_k_ptr/new_v_ptr as stack garbage; PagedKV::decode_addr then took the new-KV path on wild pointers (illegal access or wrong last-token K/V)
- value-init the POD (= {}) at every construction site
This commit is contained in:
2026-08-24 14:56:19 +08:00
parent 31ca357c61
commit f6db546578
2 changed files with 10 additions and 10 deletions
+6 -6
View File
@@ -238,7 +238,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref); B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
// Kernel launch // Kernel launch
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -373,7 +373,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
h_mask, max_sl, h_mask, max_sl,
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref); B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -509,7 +509,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti); int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti);
// Kernel launch // Kernel launch
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -651,7 +651,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
int *d_qtb, *d_qti; int *d_qtb, *d_qti;
int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti); int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -744,7 +744,7 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len; for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len;
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice); cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -825,7 +825,7 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
int *d_qtb, *d_qti; int *d_qtb, *d_qti;
int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti); int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1; p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
+4 -4
View File
@@ -58,7 +58,7 @@ static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) {
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice); cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice); cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D; p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
p.use_mask=0; p.causal_offset=causal?0:-1; p.use_mask=0; p.causal_offset=causal?0:-1;
p.scale=1.0f/sqrtf((float)D); p.scale=1.0f/sqrtf((float)D);
@@ -139,7 +139,7 @@ static void bench_decode() {
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice); cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
delete[] tmp; delete[] tmp;
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl; p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl;
p.head_dim = D; p.use_mask = 0; p.causal_offset = -1; p.head_dim = D; p.use_mask = 0; p.causal_offset = -1;
p.scale = 1.0f / sqrtf((float)D); p.scale = 1.0f / sqrtf((float)D);
@@ -186,7 +186,7 @@ static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int ca
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]); for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice); cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D; p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
p.use_mask=0; p.causal_offset=causal?0:-1; p.use_mask=0; p.causal_offset=causal?0:-1;
set_default_strides(p); set_default_strides(p);
@@ -266,7 +266,7 @@ static void bench_prefill() {
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf()); for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice); cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
AttentionParams<bf16> p; AttentionParams<bf16> p = {};
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D; p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
p.use_mask=0; p.causal_offset=causal?0:-1; p.use_mask=0; p.causal_offset=causal?0:-1;
set_default_strides(p); set_default_strides(p);