feat: SGLang-style paged attention kernels replace page-table path

- PagedAttentionParams uses flat KV pool + req_to_token + kv_indptr/qo_indptr instead of page_table
- MMA split-KV decode and split-Q prefill kernels with indirect ragged-batch addressing
- Prefill kernel accepts 4D mask (causal-aware); decode kernel supports 2D mask
- CudaBackend is inference-only: kv_cache=None raises, no torch fallback
- benchmark.py: required --ckpt, --backend/--compare options
- Parallel build isolates build-temp/build-lib per subprocess
- Standalone test covers decode/prefill with mask, 27 cases pass
This commit is contained in:
2026-08-01 15:41:25 +08:00
parent 9960f79920
commit 41dcf0feb9
17 changed files with 1683 additions and 535 deletions
-308
View File
@@ -1,308 +0,0 @@
// Compile:
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
// --extra-device-vectorization csrc/tests/attn_paged_decode_test.cu \
// -o /tmp/test_paged && /tmp/test_paged
#include <cstring>
#include "test_utils.cuh"
#include "../kernels/attn_dispatchers.cuh"
static void gather_kv_cpu(
const bf16* h_k_pool, const bf16* h_v_pool,
const int64_t* h_pt, int B, int Hkv, int kv_len,
int page_size, int head_dim,
bf16* h_k, bf16* h_v)
{
int max_pages = (kv_len + page_size - 1) / page_size;
size_t page_stride = (size_t)page_size * Hkv * head_dim;
for (int b = 0; b < B; b++) {
for (int pos = 0; pos < kv_len; pos++) {
int log_pg = pos / page_size;
int pg_off = pos % page_size;
int phys = (int)h_pt[b * max_pages + log_pg];
for (int h = 0; h < Hkv; h++) {
size_t src_base = (size_t)phys * page_stride
+ (size_t)pg_off * Hkv * head_dim
+ h * head_dim;
size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim
+ (size_t)pos * head_dim;
memcpy(h_k + dst_base, h_k_pool + src_base, head_dim * sizeof(bf16));
memcpy(h_v + dst_base, h_v_pool + src_base, head_dim * sizeof(bf16));
}
}
}
}
template <int HEAD_DIM>
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causal, int seed) {
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d causal=%d ... ",
B, Hq, Hkv, kv_len, page_size, HEAD_DIM, causal);
fflush(stdout);
int max_pages = (kv_len + page_size - 1) / page_size;
int n_phys_pages = B * max_pages;
int max_splits = 32;
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
size_t sz_o = sz_q;
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
bf16 *d_q, *d_o_paged;
bf16 *d_k_pool, *d_v_pool;
int64_t* d_pt;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q);
cudaMalloc(&d_o_paged, sz_o);
cudaMalloc(&d_k_pool, sz_kv);
cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_pt, sz_pt);
cudaMalloc(&d_op, sz_op);
cudaMalloc(&d_ml, sz_ml);
srand(seed);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
h_q[i] = __float2bfloat16(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
size_t ps = (size_t)page_size * Hkv * HEAD_DIM;
for (int pg = 0; pg < n_phys_pages; pg++) {
for (int off = 0; off < page_size; off++) {
for (int h = 0; h < Hkv; h++) {
for (int d = 0; d < HEAD_DIM; d++) {
float v = sinf((float)(pg * 7919 + off * 1049 + h * 331 + d));
size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM
+ h * HEAD_DIM + d;
h_k_pool[idx] = __float2bfloat16(v);
h_v_pool[idx] = __float2bfloat16(v * 0.3f);
}
}
}
}
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_pt = (int64_t*)malloc(sz_pt);
int next_pg = 0;
for (int b = 0; b < B; b++)
for (int p = 0; p < max_pages; p++)
h_pt[b * max_pages + p] = next_pg++;
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
bf16* h_k_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
bf16* h_v_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
gather_kv_cpu(h_k_pool, h_v_pool, h_pt, B, Hkv, kv_len, page_size, HEAD_DIM, h_k_cont, h_v_cont);
float* h_q_f = (float*)malloc((size_t)B * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < B * kv_len * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_cont[i]);
h_v_f[i] = bf2f(h_v_cont[i]);
}
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv,
1, kv_len, HEAD_DIM, causal ? 0 : -1);
PagedAttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1;
p.kv_len = kv_len; p.head_dim = HEAD_DIM;
p.use_mask = 0; p.causal_offset = causal ? 0 : -1;
set_default_paged_strides(p);
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.page_size = page_size; p.max_pages = max_pages;
p.page_table = d_pt;
p.k_cache = d_k_pool; p.v_cache = d_v_pool;
p.q = d_q; p.mask = nullptr; p.o = d_o_paged;
p.o_part = d_op; p.ml_part = d_ml;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p); });
cudaDeviceSynchronize();
bf16* h_o_bf16 = (bf16*)malloc(sz_o);
cudaMemcpy(h_o_bf16, d_o_paged, sz_o, cudaMemcpyDeviceToHost);
float* h_o_paged = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
h_o_paged[i] = __bfloat162float(h_o_bf16[i]);
float max_abs_err = 0.0f, max_rel_err = 0.0f;
int bad_idx = -1;
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_paged[i] - h_o_ref[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;
}
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) {
printf("PASS (max_abs_err=%.4e max_rel_err=%.4e)\n", max_abs_err, max_rel_err);
} else {
int b = bad_idx / (Hq * HEAD_DIM);
int h = (bad_idx / HEAD_DIM) % Hq;
int d = bad_idx % HEAD_DIM;
printf("FAIL (max_abs_err=%.4e max_rel_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
max_abs_err, max_rel_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
printf(" ref[0..7]:");
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
printf(" %.4f", h_o_ref[i]);
printf("\n got[0..7]:");
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
printf(" %.4f", h_o_paged[i]);
printf("\n");
}
free(h_q); free(h_k_pool); free(h_v_pool); free(h_pt);
free(h_k_cont); free(h_v_cont);
free(h_q_f); free(h_k_f); free(h_v_f);
free(h_o_ref); free(h_o_bf16); free(h_o_paged);
cudaFree(d_q); cudaFree(d_o_paged);
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
cudaFree(d_op); cudaFree(d_ml);
return pass ? 0 : 1;
}
struct TestCase {
int head_dim;
int B, Hq, Hkv, kv_len, page_size, causal, seed;
};
static const TestCase TESTS[] = {
{128, 1, 1, 1, 8, 128, 0, 1},
{128, 1, 4, 4, 128, 128, 0, 2},
{128, 2, 4, 4, 256, 128, 0, 3},
{128, 1, 4, 1, 64, 64, 0, 4},
{128, 1, 8, 2, 64, 128, 0, 5},
{128, 2, 16, 4, 128, 128, 0, 6},
{64, 1, 4, 2, 32, 128, 0, 7},
{256, 1, 2, 1, 16, 128, 0, 8},
{32, 1, 4, 2, 32, 64, 0, 9},
{128, 3, 8, 2, 256, 128, 0, 10},
{128, 2, 32, 8, 512, 128, 0, 11},
{128, 1, 16, 2, 256, 128, 0, 12},
{128, 2, 32, 4, 512, 128, 0, 13},
{128, 2, 8, 2, 128, 128, 1, 14}, // causal
};
static int dispatch_test(const TestCase& tc) {
int r = 0;
dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.causal, tc.seed);
});
return r;
}
template <int HEAD_DIM>
static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
int max_pages = (kv_len + page_size - 1) / page_size;
int n_phys_pages = B * max_pages;
int max_splits = 32;
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t* d_pt;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_pt, sz_pt);
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_pt = (int64_t*)malloc(sz_pt);
int next_pg = 0;
for (int b = 0; b < B; b++)
for (int p = 0; p < max_pages; p++)
h_pt[b * max_pages + p] = next_pg++;
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
free(h_pt);
PagedAttentionParams<bf16> pa;
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1;
pa.kv_len = kv_len; pa.head_dim = HEAD_DIM;
pa.use_mask = 0; pa.causal_offset = -1;
set_default_paged_strides(pa);
pa.scale = 1.0f / sqrtf((float)HEAD_DIM);
pa.page_size = page_size; pa.max_pages = max_pages;
pa.page_table = d_pt;
pa.k_cache = d_k_pool; pa.v_cache = d_v_pool;
pa.q = d_q; pa.mask = nullptr; pa.o = d_o;
pa.o_part = d_op; pa.ml_part = d_ml;
const int WARMUP = 10, ITERS = 100;
auto launch = [&]() {
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(pa); });
};
double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM;
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
char cfg[64];
snprintf(cfg, sizeof(cfg),
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d page=%3d",
B, Hq, Hkv, 1, kv_len, HEAD_DIM, page_size);
print_bench_row(cfg, r);
free(tmp);
cudaFree(d_q); cudaFree(d_o);
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
cudaFree(d_op); cudaFree(d_ml);
}
static void bench() {
printf("\n===== PAGED DECODE BENCH =====\n");
print_bench_header();
bench_config<128>(1, 32, 4, 512, 128);
bench_config<128>(1, 32, 4, 1024, 128);
bench_config<128>(1, 32, 4, 2048, 128);
bench_config<128>(1, 32, 4, 4096, 128);
bench_config<128>(16, 32, 4, 2048, 128);
bench_config<128>(32, 32, 4, 1024, 128);
}
int main() {
int n = sizeof(TESTS) / sizeof(TESTS[0]);
int fail = 0;
printf("=== Paged Decode vs CPU reference (%d cases) ===\n\n", n);
for (int i = 0; i < n; i++) {
fail += dispatch_test(TESTS[i]);
if (fail) break;
}
if (fail) {
printf("\nFAILED (%d/%d tests failed)\n", fail, n);
return fail;
}
printf("\nAll %d tests passed!\n", n);
bench();
return 0;
}
+939
View File
@@ -0,0 +1,939 @@
// Compile:
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
// --extra-device-vectorization csrc/tests/attn_paged_test.cu \
// -o /tmp/test_paged && /tmp/test_paged
#include <cstring>
#include <vector>
#include "test_utils.cuh"
#include "../kernels/attn_dispatchers.cuh"
// ---- CPU reference: paged decode with variable seq_lens ----
// Q: [B, Hq, D], K/V pool: [pool_size, Hkv, D]
// req_to_token: [num_reqs, max_ctx_len], req_pool_indices: [B]
// kv_indptr: [B+1]. mask: [B, max_seq_len] bool (True=keep) or NULL.
static void cpu_paged_decode_ref(
const float* Q, const float* K_pool, const float* V_pool,
const int64_t* req_to_token, const int64_t* req_pool_indices,
const int* kv_indptr, const bool* mask, int mask_b_stride,
int B, int Hq, int Hkv, int D, int max_ctx_len,
float* O)
{
float scale = 1.0f / sqrtf((float)D);
int n_rep = Hq / Hkv;
for (int b = 0; b < B; b++) {
int seq_len = kv_indptr[b + 1] - kv_indptr[b];
int64_t req_idx = req_pool_indices[b];
for (int h = 0; h < Hq; h++) {
int kv_h = h / n_rep;
float mv = -INFINITY, sv = 0.0f;
float accum[256] = {0.0f};
for (int kj = 0; kj < seq_len; kj++) {
if (mask && !mask[b * mask_b_stride + kj]) continue;
int64_t slot = req_to_token[req_idx * max_ctx_len + kj];
float dot = 0.0f;
for (int d = 0; d < D; d++)
dot += Q[(b * Hq + h) * D + d] *
K_pool[slot * Hkv * D + kv_h * D + d];
dot *= scale;
float nm = fmaxf(mv, dot);
float a = expf(mv - nm);
float be = expf(dot - nm);
sv = sv * a + be;
for (int d = 0; d < D; d++)
accum[d] = accum[d] * a +
V_pool[slot * Hkv * D + kv_h * D + d] * be;
mv = nm;
}
float inv = 1.0f / sv;
for (int d = 0; d < D; d++)
O[(b * Hq + h) * D + d] = accum[d] * inv;
}
}
}
// ---- CPU reference: paged prefill with ragged batch ----
// Q: [total_q, Hq, D], K/V pool: [pool_size, Hkv, D]
// req_to_token: [num_reqs, max_ctx_len], req_pool_indices: [B]
// kv_indptr: [B+1], qo_indptr: [B+1].
// mask: [B, max_q_len, max_seq_len] bool (True=keep, q-local + kv-local
// positions) or NULL. Used only when causal==0 to apply an arbitrary
// attention mask on top of the (unused) causal logic.
static void cpu_paged_prefill_ref(
const float* Q, const float* K_pool, const float* V_pool,
const int64_t* req_to_token, const int64_t* req_pool_indices,
const int* kv_indptr, const int* qo_indptr,
const bool* mask, int mask_q_stride, int mask_kv_stride,
int B, int Hq, int Hkv, int D, int max_ctx_len, int causal,
float* O)
{
float scale = 1.0f / sqrtf((float)D);
int n_rep = Hq / Hkv;
for (int b = 0; b < B; b++) {
int seq_len = kv_indptr[b + 1] - kv_indptr[b];
int q_len = qo_indptr[b + 1] - qo_indptr[b];
int causal_off = seq_len - q_len;
int64_t req_idx = req_pool_indices[b];
for (int h = 0; h < Hq; h++) {
int kv_h = h / n_rep;
for (int qi = 0; qi < q_len; qi++) {
float mv = -INFINITY, sv = 0.0f;
float accum[256] = {0.0f};
int lim = causal ? min(seq_len, causal_off + qi + 1) : seq_len;
for (int kj = 0; kj < lim; kj++) {
if (mask && !mask[b * mask_q_stride * mask_kv_stride
+ qi * mask_kv_stride + kj]) continue;
int64_t slot = req_to_token[req_idx * max_ctx_len + kj];
float dot = 0.0f;
for (int d = 0; d < D; d++)
dot += Q[(qo_indptr[b] + qi) * Hq * D + h * D + d] *
K_pool[slot * Hkv * D + kv_h * D + d];
dot *= scale;
float nm = fmaxf(mv, dot);
float a = expf(mv - nm);
float be = expf(dot - nm);
sv = sv * a + be;
for (int d = 0; d < D; d++)
accum[d] = accum[d] * a +
V_pool[slot * Hkv * D + kv_h * D + d] * be;
mv = nm;
}
float inv = 1.0f / sv;
for (int d = 0; d < D; d++)
O[(qo_indptr[b] + qi) * Hq * D + h * D + d] = accum[d] * inv;
}
}
}
}
// ======================================================================
// DECODE TEST
// ======================================================================
template <int HEAD_DIM>
static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
int causal, int seed) {
// Variable seq_lens per request
srand(seed);
std::vector<int> seq_lens(B);
for (int b = 0; b < B; b++)
seq_lens[b] = 8 + rand() % (max_seq - 8);
int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end());
int max_ctx = max_sl + 16;
int pool_size = B * max_ctx;
int num_reqs = B + 4;
printf("DECODE B=%d Hq=%d Hkv=%d D=%d seqs=[", B, Hq, Hkv, HEAD_DIM);
for (int b = 0; b < B; b++) printf("%d%s", seq_lens[b], b < B-1 ? "," : "");
printf("] causal=%d ... ", causal);
fflush(stdout);
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi;
int *d_kvi;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi);
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
h_k_pool[i] = f2bf(rnd());
h_v_pool[i] = f2bf(rnd());
}
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
// req_to_token: assign unique slots per request (scattered, not contiguous)
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
int next_slot = 0;
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) {
h_rtt[r * max_ctx + p] = next_slot % pool_size;
next_slot++;
}
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
// req_pool_indices: pick B random request rows
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
// kv_indptr: prefix sum of seq_lens
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0;
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_lens[b];
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
// CPU reference
float* h_q_f = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_pool[i]);
h_v_f[i] = bf2f(h_v_pool[i]);
}
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
cpu_paged_decode_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi,
nullptr, 0,
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
// Kernel launch
PagedAttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
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_q_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = 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;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p); });
cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q);
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f;
bool pass = true;
float max_err = 0.0f;
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_got[i] - h_o_ref[i]);
if (e > max_err) max_err = e;
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
}
if (pass) printf("PASS (max_err=%.4e)\n", max_err);
else printf("FAIL (max_err=%.4e)\n", max_err);
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
free(h_kvi); free(h_q_f); free(h_k_f); free(h_v_f);
free(h_o_ref); free(h_o_bf); free(h_o_got);
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_op); cudaFree(d_ml);
return pass ? 0 : 1;
}
// ======================================================================
// DECODE WITH MASK TEST (regression: 2D mask on mixed seq_lens)
// ======================================================================
template <int HEAD_DIM>
static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
int seed) {
srand(seed);
std::vector<int> seq_lens(B);
for (int b = 0; b < B; b++)
seq_lens[b] = 8 + rand() % (max_seq - 8);
int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end());
int max_ctx = max_sl + 16;
int pool_size = B * max_ctx;
int num_reqs = B + 4;
printf("DECODE-MASK B=%d Hq=%d Hkv=%d D=%d max_sl=%d ... ", B, Hq, Hkv, HEAD_DIM, max_sl);
fflush(stdout);
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_mask = (size_t)B * max_sl * sizeof(bool);
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi;
int *d_kvi;
bool *d_mask;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi);
cudaMalloc(&d_mask, sz_mask);
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
h_k_pool[i] = f2bf(rnd());
h_v_pool[i] = f2bf(rnd());
}
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
int next_slot = 0;
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) {
h_rtt[r * max_ctx + p] = next_slot % pool_size;
next_slot++;
}
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0;
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_lens[b];
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
// Mask: keep first half of each request's kv range, drop the rest —
// exercises the HasMask path with per-request seq_len.
bool* h_mask = (bool*)malloc(sz_mask);
for (int b = 0; b < B; b++)
for (int k = 0; k < max_sl; k++)
h_mask[b * max_sl + k] = (k < seq_lens[b]) && (k % 2 == 0);
cudaMemcpy(d_mask, h_mask, sz_mask, cudaMemcpyHostToDevice);
float* h_q_f = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_pool[i]);
h_v_f[i] = bf2f(h_v_pool[i]);
}
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
cpu_paged_decode_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi,
h_mask, max_sl,
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
PagedAttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
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_q_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = 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;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p); });
cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q);
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f;
bool pass = true;
float max_err = 0.0f;
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_got[i] - h_o_ref[i]);
if (e > max_err) max_err = e;
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
}
if (pass) printf("PASS (max_err=%.4e)\n", max_err);
else printf("FAIL (max_err=%.4e)\n", max_err);
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
free(h_kvi); free(h_mask); free(h_q_f); free(h_k_f); free(h_v_f);
free(h_o_ref); free(h_o_bf); free(h_o_got);
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_mask);
cudaFree(d_op); cudaFree(d_ml);
return pass ? 0 : 1;
}
// ======================================================================
// PREFILL TEST
// ======================================================================
template <int HEAD_DIM>
static int run_prefill_test(int B, int Hq, int Hkv,
std::vector<int>& q_lens,
std::vector<int>& kv_lens,
int causal, int seed) {
int total_q = 0;
int max_sl = 0;
for (int b = 0; b < B; b++) {
total_q += q_lens[b];
max_sl = max(max_sl, kv_lens[b]);
}
int max_ctx = max_sl + 16;
int pool_size = B * max_ctx;
int num_reqs = B + 4;
printf("PREFILL B=%d Hq=%d Hkv=%d D=%d q_lens=[", B, Hq, Hkv, HEAD_DIM);
for (int b = 0; b < B; b++) printf("%d%s", q_lens[b], b < B-1 ? "," : "");
printf("] kv_lens=[");
for (int b = 0; b < B; b++) printf("%d%s", kv_lens[b], b < B-1 ? "," : "");
printf("] causal=%d ... ", causal);
fflush(stdout);
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi;
int *d_kvi, *d_qoi;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi); cudaMalloc(&d_qoi, sz_qoi);
srand(seed);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
h_k_pool[i] = f2bf(rnd());
h_v_pool[i] = f2bf(rnd());
}
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
int next_slot = 0;
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) {
h_rtt[r * max_ctx + p] = next_slot % pool_size;
next_slot++;
}
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0;
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + kv_lens[b];
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
int* h_qoi = (int*)malloc(sz_qoi);
h_qoi[0] = 0;
for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_lens[b];
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
// CPU reference
float* h_q_f = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_pool[i]);
h_v_f[i] = bf2f(h_v_pool[i]);
}
float* h_o_ref = (float*)calloc(total_q * Hq * HEAD_DIM, sizeof(float));
cpu_paged_prefill_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi, h_qoi,
nullptr, 0, 0,
B, Hq, Hkv, HEAD_DIM, max_ctx, causal, h_o_ref);
// Kernel launch
PagedAttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
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.causal_offset = causal ? 0 : -1; p.use_mask = 0;
p.mask = nullptr; p.mask_b_stride = 0;
p.mask_h_stride = 0; p.mask_q_stride = 0;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = 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;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p); });
cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q);
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f;
bool pass = true;
float max_err = 0.0f;
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_got[i] - h_o_ref[i]);
if (e > max_err) max_err = e;
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
}
if (pass) printf("PASS (max_err=%.4e)\n", max_err);
else printf("FAIL (max_err=%.4e)\n", max_err);
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
free(h_kvi); free(h_qoi); free(h_q_f); free(h_k_f); free(h_v_f);
free(h_o_ref); free(h_o_bf); free(h_o_got);
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
return pass ? 0 : 1;
}
// ======================================================================
// PREFILL WITH MASK TEST (regression: 4D causal mask on single request)
// ======================================================================
template <int HEAD_DIM>
static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
srand(seed);
int B = 1;
int total_q = q_len;
int seq_len = q_len; // pure prefill: kv_len == q_len
int max_ctx = seq_len + 16;
int pool_size = B * max_ctx;
int num_reqs = B + 4;
printf("PREFILL-MASK Hq=%d Hkv=%d D=%d q_len=%d ... ", Hq, Hkv, HEAD_DIM, q_len);
fflush(stdout);
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
size_t sz_mask = (size_t)B * q_len * q_len * sizeof(bool);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi;
int *d_kvi, *d_qoi;
bool *d_mask;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi); cudaMalloc(&d_qoi, sz_qoi);
cudaMalloc(&d_mask, sz_mask);
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
bf16* h_q = (bf16*)malloc(sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
bf16* h_k_pool = (bf16*)malloc(sz_kv);
bf16* h_v_pool = (bf16*)malloc(sz_kv);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
h_k_pool[i] = f2bf(rnd());
h_v_pool[i] = f2bf(rnd());
}
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
int next_slot = 0;
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++) {
h_rtt[r * max_ctx + p] = next_slot % pool_size;
next_slot++;
}
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
h_rpi[0] = 0;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0; h_kvi[1] = seq_len;
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
int* h_qoi = (int*)malloc(sz_qoi);
h_qoi[0] = 0; h_qoi[1] = q_len;
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
// 4D causal mask [B, 1, q_len, q_len], True=keep.
bool* h_mask = (bool*)malloc(sz_mask);
for (int qi = 0; qi < q_len; qi++)
for (int kj = 0; kj < q_len; kj++)
h_mask[qi * q_len + kj] = (kj <= qi);
cudaMemcpy(d_mask, h_mask, sz_mask, cudaMemcpyHostToDevice);
float* h_q_f = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
h_k_f[i] = bf2f(h_k_pool[i]);
h_v_f[i] = bf2f(h_v_pool[i]);
}
float* h_o_ref = (float*)calloc(total_q * Hq * HEAD_DIM, sizeof(float));
// CPU ref with causal=0 so it consults the mask (not the causal flag).
cpu_paged_prefill_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi, h_qoi,
h_mask, q_len, q_len,
B, Hq, Hkv, HEAD_DIM, max_ctx, 0, h_o_ref);
PagedAttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = q_len;
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;
p.mask_h_stride = 0; p.mask_q_stride = q_len;
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = 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;
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p); });
cudaDeviceSynchronize();
bf16* h_o_bf = (bf16*)malloc(sz_q);
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
const float atol = 0.02f, rtol = 0.02f;
bool pass = true;
float max_err = 0.0f;
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) {
float e = fabsf(h_o_got[i] - h_o_ref[i]);
if (e > max_err) max_err = e;
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
}
if (pass) printf("PASS (max_err=%.4e)\n", max_err);
else printf("FAIL (max_err=%.4e)\n", max_err);
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
free(h_kvi); free(h_qoi); free(h_mask); free(h_q_f); free(h_k_f); free(h_v_f);
free(h_o_ref); free(h_o_bf); free(h_o_got);
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
cudaFree(d_mask);
return pass ? 0 : 1;
}
// ======================================================================
// BENCH
// ======================================================================
template <int HEAD_DIM>
static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
int max_ctx = seq_len + 16;
int pool_size = B * max_ctx;
int num_reqs = B;
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi;
int *d_kvi;
float *d_op, *d_ml;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi);
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++)
h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size;
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0;
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len;
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
PagedAttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = B;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
p.max_context_len = max_ctx; p.max_seq_len = seq_len;
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 = d_q; p.k_cache = d_k_pool; p.v_cache = 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;
auto launch = [&]() {
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p); });
};
// Decode: q_len=1, query is the last token → attends to all [0, seq_len).
// FLOPs = 2 * (QK^T + PV) = 4 * B * Hq * seq_len * D.
double flops = 4.0 * B * Hq * (double)seq_len * HEAD_DIM;
// HBM: K+V read (Q/O negligible for decode).
size_t nKV = (size_t)B * Hkv * seq_len * HEAD_DIM;
double bytes = 2.0 * nKV * sizeof(bf16);
BenchResult r = bench_kernel(launch, 10, 100, flops, bytes);
char cfg[64];
snprintf(cfg, sizeof(cfg), "DEC B=%2d Hq=%2d Hk=%d kv=%4d D=%3d",
B, Hq, Hkv, seq_len, HEAD_DIM);
print_bench_row(cfg, r);
free(tmp); free(h_rtt); free(h_rpi); free(h_kvi);
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_op); cudaFree(d_ml);
}
template <int HEAD_DIM>
static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int causal) {
int total_q = B * q_len;
int max_ctx = kv_len + 16;
int pool_size = B * max_ctx;
int num_reqs = B;
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
size_t sz_rpi = (size_t)B * sizeof(int64_t);
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
int64_t *d_rtt, *d_rpi;
int *d_kvi, *d_qoi;
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
cudaMalloc(&d_kvi, sz_kvi); cudaMalloc(&d_qoi, sz_qoi);
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
for (int r = 0; r < num_reqs; r++)
for (int p = 0; p < max_ctx; p++)
h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size;
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
for (int b = 0; b < B; b++) h_rpi[b] = b;
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
int* h_kvi = (int*)malloc(sz_kvi);
h_kvi[0] = 0;
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + kv_len;
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
int* h_qoi = (int*)malloc(sz_qoi);
h_qoi[0] = 0;
for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_len;
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
PagedAttentionParams<bf16> p;
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM; p.total_q = total_q;
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 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.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 = d_q; p.k_cache = d_k_pool; p.v_cache = 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;
auto launch = [&]() {
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p); });
};
// FLOPs = 2 * (QK^T + PV) = 4 * effective_qk_pairs * Hq * D.
// Non-causal: effective = q_len * kv_len.
// Causal: Q row qi attends to [0, causal_off + qi + 1) where
// causal_off = kv_len - q_len. Total KV accesses per request:
// sum_{qi=0}^{q_len-1} (kv_len - q_len + qi + 1)
// = q_len * (kv_len - q_len) + q_len * (q_len + 1) / 2.
double eff_kv;
if (causal) {
eff_kv = (double)q_len * (kv_len - q_len)
+ (double)q_len * (q_len + 1) / 2.0;
} else {
eff_kv = (double)q_len * kv_len;
}
double flops = 4.0 * B * Hq * eff_kv * HEAD_DIM;
// HBM: Q read + K read + V read + O write.
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
size_t nQ = (size_t)total_q * Hq * HEAD_DIM;
double bytes = (2.0 * nQ + 2.0 * nKV) * sizeof(bf16);
BenchResult r = bench_kernel(launch, 10, 100, flops, bytes);
char cfg[80];
snprintf(cfg, sizeof(cfg), "PRE B=%d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d c=%d",
B, Hq, Hkv, q_len, kv_len, HEAD_DIM, causal);
print_bench_row(cfg, r);
free(tmp); free(h_rtt); free(h_rpi); free(h_kvi); free(h_qoi);
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
}
int main() {
int fail = 0;
// ===== DECODE TESTS =====
printf("=== Paged Decode Tests ===\n\n");
fail += run_decode_test<128>(1, 32, 4, 512, 0, 1);
fail += run_decode_test<128>(1, 32, 4, 1024, 0, 2);
fail += run_decode_test<128>(4, 32, 4, 512, 0, 3);
fail += run_decode_test<128>(8, 32, 4, 1024, 0, 4);
fail += run_decode_test<128>(4, 32, 8, 2048, 0, 5);
fail += run_decode_test<128>(1, 16, 1, 256, 0, 6);
fail += run_decode_test<128>(2, 8, 2, 512, 1, 7);
fail += run_decode_test<64>(1, 4, 2, 256, 0, 8);
fail += run_decode_test<256>(1, 2, 1, 256, 0, 9);
fail += run_decode_test<128>(16, 32, 4, 2048, 0, 10);
fail += run_decode_test<128>(32, 32, 4, 1024, 0, 11);
// Decode with 2D mask (regression: mixed seq_lens + HasMask)
fail += run_decode_mask_test<128>(2, 8, 2, 256, 30);
fail += run_decode_mask_test<128>(4, 32, 4, 512, 31);
fail += run_decode_mask_test<64>(2, 4, 2, 128, 32);
if (fail) { printf("\nFAILED decode tests\n"); return fail; }
// ===== PREFILL TESTS =====
printf("\n=== Paged Prefill Tests ===\n\n");
// Single request, pure prefill (q_len == kv_len)
{
std::vector<int> ql = {512};
std::vector<int> kl = {512};
fail += run_prefill_test<128>(1, 32, 4, ql, kl, 1, 20);
}
{
std::vector<int> ql = {1024};
std::vector<int> kl = {1024};
fail += run_prefill_test<128>(1, 32, 4, ql, kl, 1, 21);
}
{
std::vector<int> ql = {2048};
std::vector<int> kl = {2048};
fail += run_prefill_test<128>(1, 32, 4, ql, kl, 1, 22);
}
// Ragged batch: different q_lens and kv_lens
{
std::vector<int> ql = {128, 256, 64};
std::vector<int> kl = {128, 256, 64};
fail += run_prefill_test<128>(3, 32, 4, ql, kl, 1, 23);
}
{
std::vector<int> ql = {64, 128, 256, 32};
std::vector<int> kl = {64, 128, 256, 32};
fail += run_prefill_test<128>(4, 32, 4, ql, kl, 1, 24);
}
// Extend: kv_len > q_len (append to existing cache)
{
std::vector<int> ql = {64, 128};
std::vector<int> kl = {256, 512};
fail += run_prefill_test<128>(2, 32, 4, ql, kl, 1, 25);
}
// Non-causal
{
std::vector<int> ql = {256, 128};
std::vector<int> kl = {256, 128};
fail += run_prefill_test<128>(2, 32, 4, ql, kl, 0, 26);
}
// Single token (q_len=1 per request, like decode but via prefill path)
{
std::vector<int> ql = {1, 1, 1, 1};
std::vector<int> kl = {128, 256, 64, 512};
fail += run_prefill_test<128>(4, 32, 4, ql, kl, 1, 27);
}
// D=64
{
std::vector<int> ql = {128, 64};
std::vector<int> kl = {128, 64};
fail += run_prefill_test<64>(2, 4, 2, ql, kl, 1, 28);
}
// D=256
{
std::vector<int> ql = {128, 64};
std::vector<int> kl = {128, 64};
fail += run_prefill_test<256>(2, 2, 1, ql, kl, 1, 29);
}
// Prefill with 4D causal mask (regression: single-request mask path)
fail += run_prefill_mask_test<128>(32, 4, 512, 40);
fail += run_prefill_mask_test<128>(32, 4, 1024, 41);
fail += run_prefill_mask_test<64>(4, 2, 256, 42);
if (fail) { printf("\nFAILED prefill tests\n"); return fail; }
printf("\nAll tests passed!\n");
// ===== BENCH =====
printf("\n===== PAGED DECODE BENCH =====\n");
print_bench_header();
bench_decode<128>(1, 32, 4, 512);
bench_decode<128>(1, 32, 4, 1024);
bench_decode<128>(1, 32, 4, 2048);
bench_decode<128>(1, 32, 4, 4096);
bench_decode<128>(4, 32, 4, 2048);
bench_decode<128>(16, 32, 4, 2048);
bench_decode<128>(32, 32, 4, 1024);
printf("\n===== PAGED PREFILL BENCH =====\n");
print_bench_header();
bench_prefill<128>(1, 32, 4, 512, 512, 0);
bench_prefill<128>(1, 32, 4, 1024, 1024, 0);
bench_prefill<128>(1, 32, 4, 2048, 2048, 0);
bench_prefill<128>(1, 32, 4, 2048, 2048, 1);
bench_prefill<128>(4, 32, 4, 2048, 2048, 1);
bench_prefill<128>(1, 32, 4, 4096, 4096, 1);
return 0;
}