perf: flatten paged prefill tile dispatch

- remove the host-provided max_q_len argument
- dispatch only the ragged prefill tile upper bound
- validate the rebuilt CUDA backend end to end
This commit is contained in:
2026-08-09 23:12:53 +08:00
parent cd31f1f62f
commit c5fba9c238
10 changed files with 69 additions and 33 deletions
+3 -5
View File
@@ -489,9 +489,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
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;
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;
p.q_len = total_q;
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;
@@ -626,7 +624,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
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_q_len = q_len;
p.q_len = B * q_len;
p.causal_offset = -1; p.use_mask = 1;
p.mask = d_mask; p.mask_b_stride = q_len * q_len;
p.mask_h_stride = 0; p.mask_l_stride = q_len;
@@ -793,7 +791,7 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
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_q_len = q_len;
p.q_len = B * 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);