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:
@@ -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;
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
Reference in New Issue
Block a user