refactor: tidy attention params and launcher interfaces

- rename output pointer field o to o_ptr for consistency with q_ptr/k_ptr/v_ptr
- regroup AttentionParams fields by responsibility and fix misleading comments
- drop unused max_seq_len/total_q fields and paged decode max_seq_len arg
- drop redundant group_size param from decode launchers (computed from p)
This commit is contained in:
2026-08-09 20:52:06 +08:00
parent a5a3cc1fc2
commit cd31f1f62f
14 changed files with 69 additions and 84 deletions
+19 -19
View File
@@ -218,9 +218,9 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
// Kernel launch
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
p.max_context_len = max_ctx;
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0;
p.mask_h_stride = 0; p.mask_l_stride = 0;
@@ -228,7 +228,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
p.o_ptr = d_o; p.o_part = d_op; p.ml_part = d_ml;
dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p});
cudaDeviceSynchronize();
@@ -353,9 +353,9 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
p.max_context_len = max_ctx;
p.causal_offset = -1; p.use_mask = 1;
p.mask = d_mask; p.mask_b_stride = max_sl;
p.mask_h_stride = 0; p.mask_l_stride = 0;
@@ -363,7 +363,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
p.o_ptr = d_o; p.o_part = d_op; p.ml_part = d_ml;
dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p});
cudaDeviceSynchronize();
@@ -486,9 +486,9 @@ static int run_prefill_test(int B, int Hq, int Hkv,
// Kernel launch
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
p.max_context_len = max_ctx;
int max_ql = 0;
for (int b = 0; b < B; b++) max_ql = max(max_ql, q_lens[b]);
p.max_q_len = max_ql;
@@ -499,7 +499,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr;
dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p});
cudaDeviceSynchronize();
@@ -623,9 +623,9 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = q_len;
p.max_context_len = max_ctx;
p.max_q_len = q_len;
p.causal_offset = -1; p.use_mask = 1;
p.mask = d_mask; p.mask_b_stride = q_len * q_len;
@@ -634,7 +634,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr;
dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p});
cudaDeviceSynchronize();
@@ -713,16 +713,16 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = seq_len;
p.max_context_len = max_ctx;
p.causal_offset = 0; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
p.o_ptr = d_o; p.o_part = d_op; p.ml_part = d_ml;
auto launch = [&]() {
dispatch_by_head_dim(HEAD_DIM, PagedDecodeDispatch{p});
@@ -790,17 +790,17 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
AttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
p.max_context_len = max_ctx; p.max_seq_len = kv_len;
p.total_q = total_q; p.max_q_len = q_len;
p.max_context_len = max_ctx;
p.max_q_len = q_len;
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q_ptr = d_q; p.k_ptr = d_k_pool; p.v_ptr = d_v_pool;
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
p.o_ptr = d_o; p.o_part = nullptr; p.ml_part = nullptr;
auto launch = [&]() {
dispatch_by_head_dim(HEAD_DIM, PagedPrefillDispatch{p});
+4 -4
View File
@@ -61,7 +61,7 @@ static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) {
p.use_mask=0; p.causal_offset=causal?0:-1;
p.scale=1.0f/sqrtf((float)D);
set_default_strides(p);
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=dO;
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o_ptr=dO;
DecodeScratch sc;
setup_scratch(p, sc);
@@ -141,7 +141,7 @@ static void bench_decode() {
p.head_dim = D; p.use_mask = 0; p.causal_offset = -1;
p.scale = 1.0f / sqrtf((float)D);
set_default_strides(p);
p.q_ptr = dQ; p.k_ptr = dK; p.v_ptr = dV; p.mask = nullptr; p.o = dO;
p.q_ptr = dQ; p.k_ptr = dK; p.v_ptr = dV; p.mask = nullptr; p.o_ptr = dO;
DecodeScratch sc;
setup_scratch(p, sc);
@@ -188,7 +188,7 @@ static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int ca
p.use_mask=0; p.causal_offset=causal?0:-1;
set_default_strides(p);
p.scale=1.0f/sqrtf((float)D);
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=dO;
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o_ptr=dO;
double t0=now_ms();
dispatch_by_head_dim(D, PrefillDispatch{p});
@@ -262,7 +262,7 @@ static void bench_prefill() {
p.use_mask=0; p.causal_offset=causal?0:-1;
set_default_strides(p);
p.scale=1.0f/sqrtf((float)D);
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o=dO;
p.q_ptr=dQ; p.k_ptr=dK; p.v_ptr=dV; p.mask=nullptr; p.o_ptr=dO;
auto launch = [&]() { dispatch_by_head_dim(D, PrefillDispatch{p}); };
for (int i=0;i<WARMUP;i++) launch();