// src/kernels/kv_stream_parity.cpp - KV streaming (kv_stream.hpp) against a fully resident pool (GPU, no model). // // Decodes a synthetic sequence twice: into a fully resident pool with the identity page table, and into a streamed // layer (host copy + a small VRAM slot pool that must evict constantly). Every few cells a batch of 1-8 queries // with selections shaped like the indexer's (whole 4-cell blocks, a partial threshold block, the tail block, half of // them recent) is attended in both; `kv_stream_resolve` runs before the streamed attention. Checks: // 1. the attention outputs are BITWISE equal (the streamed readers see exactly the resident values); // 2. the residency map is consistent after every call (slot_block and page_table invert each other); // 3. no call overflowed, and the hit/miss counters add up; // 4. a ring (the MTP drafter's layout) restored from the host copy reads the same values as the resident pool. // INT8, FP16 and Q4_0 (PR #21) pools. #include "strata/kernels/kv_q4.hpp" #include "strata/kernels/kv_q8.hpp" #include "strata/kernels/kv_stream.hpp" #include "strata/kernels/qsa.hpp" #include "strata/kernels/qsa_decode_attn.hpp" #include #include #include #include #include #include #include #include namespace k = strata::kernels; namespace { int g_fail = 0; void ck(cudaError_t e, const char* w) { if (e != cudaSuccess) { std::fprintf(stderr, "%s: %s\n", w, cudaGetErrorString(e)); std::exit(2); } } template T* dalloc(size_t n) { T* p = nullptr; ck(cudaMalloc(&p, n * sizeof(T) + 64), "malloc"); ck(cudaMemset(p, 0, n * sizeof(T) + 64), "memset"); return p; } template T* halloc(size_t n) { // pinned, mapped; returns the device pointer void* h = nullptr; void* d = nullptr; ck(cudaHostAlloc(&h, n * sizeof(T) + 64, cudaHostAllocMapped), "hostalloc"); std::memset(h, 0, n * sizeof(T) + 64); ck(cudaHostGetDevicePointer(&d, h, 0), "devptr"); return (T*) d; } struct Pools { // one K/V pool set of `pages` pages k::KvHostPools p; // reused as a plain pointer bundle void alloc(int64_t pages, const k::QsaShapes& s, int fmt, bool host) { const size_t rows = (size_t) pages * s.n_head_kv * s.page_size; if (fmt == k::kKvQ4) { const size_t b = rows * k::kv_q4_bytes_per_head((int) s.head_dim); p.k_q4 = host ? halloc(b) : dalloc(b); p.v_q4 = host ? halloc(b) : dalloc(b); } else if (fmt == k::kKvInt8) { p.k_q = host ? halloc(rows * s.head_dim) : dalloc(rows * s.head_dim); p.v_q = host ? halloc(rows * s.head_dim) : dalloc(rows * s.head_dim); p.k_scale = host ? halloc(rows * 4) : dalloc(rows * 4); p.v_scale = host ? halloc(rows * 4) : dalloc(rows * 4); } else { p.k_pool = host ? halloc(rows * s.head_dim) : dalloc(rows * s.head_dim); p.v_pool = host ? halloc(rows * s.head_dim) : dalloc(rows * s.head_dim); } } k::QsaAttnPools attn(const int32_t* table) const { k::QsaAttnPools a; a.k_pool = p.k_pool; a.v_pool = p.v_pool; a.k_q = p.k_q; a.v_q = p.v_q; a.k_scale = p.k_scale; a.v_scale = p.v_scale; a.k_q4 = p.k_q4; a.v_q4 = p.v_q4; a.page_table = table; return a; } }; void append(const Pools& pl, const int32_t* table, const int32_t* step, const float* kc, const float* vc, const k::QsaShapes& s, int fmt, const k::KvHostPools* host) { if (fmt == k::kKvQ4) k::kv_append_q4_step(pl.p.k_q4, pl.p.v_q4, table, step, kc, vc, s, nullptr, host); else if (fmt == k::kKvInt8) k::kv_append_q8_step(pl.p.k_q, pl.p.v_q, pl.p.k_scale, pl.p.v_scale, table, step, kc, vc, s, nullptr, host); else k::kv_append_step(pl.p.k_pool, pl.p.v_pool, table, step, kc, vc, s, nullptr, host); } // A selection like qsa_block_topk's: ascending cells, whole blocks, the tail block's cells, and one partial block. std::vector selection(int64_t n_kv, int64_t width, std::mt19937& rng) { std::vector ids; if (n_kv <= width) { for (int64_t c = 0; c < n_kv; ++c) ids.push_back((int32_t) c); return ids; } const int64_t n_bid = n_kv / 4, tail = n_kv - n_bid * 4; std::set blocks; const int64_t want_full = (width - tail) / 4, part = (width - tail) % 4; std::uniform_int_distribution any(0, n_bid - 1), recent(std::max(0, n_bid - 2048), n_bid - 1); while ((int64_t) blocks.size() < want_full + (part ? 1 : 0)) blocks.insert(rng() % 2 ? recent(rng) : any(rng)); int64_t partial_block = part ? *std::next(blocks.begin(), (long) (rng() % blocks.size())) : -1; for (int64_t b : blocks) { const int64_t take = b == partial_block ? part : 4; for (int64_t i = 0; i < take; ++i) ids.push_back((int32_t) (b * 4 + i)); } for (int64_t c = n_bid * 4; c < n_kv; ++c) ids.push_back((int32_t) c); std::sort(ids.begin(), ids.end()); return ids; } bool run(int fmt) { const char* name = fmt == k::kKvQ4 ? "q4_0" : fmt == k::kKvInt8 ? "int8" : "fp16"; k::QsaShapes s = k::qsa_real_shapes(); const int64_t N = 40000, n_blocks = (N + 3) / 4, n_slots = 8 * 516 + 700; // must evict: slots < blocks const int64_t cap = k::qsa_selection_width(k::kTopkMaxCells, s), NQ = 8; const int64_t H = s.n_head_kv, D = s.head_dim, NH = s.n_head; std::mt19937 rng(7 + 4 * fmt); std::normal_distribution nd(0.f, 1.f); Pools ref, slots, host; ref.alloc(n_blocks, s, fmt, false); slots.alloc(n_slots, s, fmt, false); host.alloc(n_blocks, s, fmt, true); int32_t* ident = dalloc(n_blocks); { std::vector t(n_blocks); for (int64_t i = 0; i < n_blocks; ++i) t[i] = (int32_t) i; ck(cudaMemcpy(ident, t.data(), n_blocks * 4, cudaMemcpyHostToDevice), "ident"); } k::KvStreamMap m; m.page_table = dalloc(n_blocks); m.slot_block = dalloc(n_slots); m.slot_stamp = dalloc(n_slots); m.slot_ref = dalloc(n_slots); m.miss_block = dalloc(n_slots); m.miss_slot = dalloc(n_slots); m.ctl = dalloc(k::kKvCtlInts); m.n_blocks = n_blocks; m.n_slots = n_slots; k::kv_stream_reset(m, nullptr); float* kc = dalloc(H * D); float* vc = dalloc(H * D); int32_t* step = dalloc(k::kStepCount); int32_t* steps = dalloc(NQ * k::kStepCount); int32_t* ids = dalloc(NQ * cap); float* q = dalloc(NQ * NH * D); const uint64_t scr = k::qsa_decode_attn_scratch_floats(cap, s); float* scratch = dalloc(NQ * scr); float* out_ref = dalloc(NQ * NH * D); float* out_str = dalloc(NQ * NH * D); std::vector hk(H * D), hv(H * D), hq(NQ * NH * D), a((size_t) NQ * NH * D), b2((size_t) NQ * NH * D); std::vector pt(n_blocks), sb(n_slots); int batches = 0, bad = 0; for (int64_t pos = 0; pos < N; ++pos) { for (auto& x : hk) x = nd(rng) * ((pos % 17 == 0) ? 30.f : 1.f); for (auto& x : hv) x = nd(rng); const int32_t st[4] = {(int32_t) pos, (int32_t) (pos + 1), (int32_t) ((pos + 1) / 4), 0}; ck(cudaMemcpy(kc, hk.data(), hk.size() * 4, cudaMemcpyHostToDevice), "k"); ck(cudaMemcpy(vc, hv.data(), hv.size() * 4, cudaMemcpyHostToDevice), "v"); ck(cudaMemcpy(step, st, sizeof(st), cudaMemcpyHostToDevice), "step"); append(ref, ident, step, kc, vc, s, fmt, nullptr); append(slots, m.page_table, step, kc, vc, s, fmt, &host.p); if (pos % 131 != 130 && pos != N - 1) continue; // a batch: n_q queries at the last n_q positions (as a verify window), each with its own selection const int n_q = 1 + (int) (rng() % NQ); std::vector hids((size_t) NQ * cap, 0), hst(NQ * 4, 0); for (int t = 0; t < n_q; ++t) { const int64_t p = std::max(0, pos - (n_q - 1 - t)), n_kv = p + 1; const int64_t width = std::min(n_kv, cap); const std::vector sel = selection(n_kv, width, rng); if ((int64_t) sel.size() != width) { std::fprintf(stderr, "selection size %zu != %lld\n", sel.size(), (long long) width); return false; } std::copy(sel.begin(), sel.end(), hids.begin() + (size_t) t * cap); hst[t * 4 + 0] = (int32_t) p; hst[t * 4 + 1] = (int32_t) n_kv; hst[t * 4 + 2] = (int32_t) (n_kv / 4); hst[t * 4 + 3] = (int32_t) width; } for (auto& x : hq) x = nd(rng); ck(cudaMemcpy(ids, hids.data(), hids.size() * 4, cudaMemcpyHostToDevice), "ids"); ck(cudaMemcpy(steps, hst.data(), hst.size() * 4, cudaMemcpyHostToDevice), "steps"); ck(cudaMemcpy(q, hq.data(), hq.size() * 4, cudaMemcpyHostToDevice), "q"); k::qsa_decode_attn_batch(q, ref.attn(ident), ids, steps, cap, s, scratch, out_ref, n_q, nullptr); k::kv_stream_resolve(m, slots.attn(m.page_table), host.p, fmt, ids, steps, n_q, cap, s, nullptr); k::qsa_decode_attn_batch(q, slots.attn(m.page_table), ids, steps, cap, s, scratch, out_str, n_q, nullptr); ck(cudaDeviceSynchronize(), "batch"); ck(cudaMemcpy(a.data(), out_ref, (size_t) n_q * NH * D * 4, cudaMemcpyDeviceToHost), "a"); ck(cudaMemcpy(b2.data(), out_str, (size_t) n_q * NH * D * 4, cudaMemcpyDeviceToHost), "b"); if (std::memcmp(a.data(), b2.data(), (size_t) n_q * NH * D * 4) != 0) { if (bad++ < 5) std::fprintf(stderr, " %s pos %lld n_q %d: streamed attention differs\n", name, (long long) pos, n_q); } // the map inverts itself ck(cudaMemcpy(pt.data(), m.page_table, n_blocks * 4, cudaMemcpyDeviceToHost), "pt"); ck(cudaMemcpy(sb.data(), m.slot_block, n_slots * 4, cudaMemcpyDeviceToHost), "sb"); int64_t resident = 0; for (int64_t bk = 0; bk < n_blocks; ++bk) { if (pt[bk] < -1 || pt[bk] >= n_slots || (pt[bk] >= 0 && sb[pt[bk]] != bk)) { if (bad++ < 5) std::fprintf(stderr, " map broken at block %lld (table %d)\n", (long long) bk, pt[bk]); break; } resident += pt[bk] >= 0; } for (int64_t sl = 0; sl < n_slots; ++sl) if (sb[sl] >= 0 && pt[sb[sl]] != sl) { if (bad++ < 5) std::fprintf(stderr, " slot %lld not in the table\n", (long long) sl); break; } ++batches; } const k::KvStreamCounters c = k::kv_stream_counters(m); std::printf(" %s: %d batches, %llu block lookups, %llu misses (%.1f%% hit), overflow %d, %d failures\n", name, batches, (unsigned long long) c.lookups, (unsigned long long) c.misses, c.lookups ? 100.0 * (double) (c.lookups - c.misses) / (double) c.lookups : 0.0, (int) c.overflow, bad); if (c.overflow || c.calls != (uint64_t) batches || c.misses == 0 || c.misses >= c.lookups) ++bad; // 4. a ring over the same host copy: blocks [b1 - R, b1) restored, read through `b % R` { const int64_t R = 1500, b1 = n_blocks, b0 = b1 - R; Pools ring; ring.alloc(R, s, fmt, false); int32_t* rt = dalloc(n_blocks); k::kv_ring_table(rt, n_blocks, R, nullptr); k::kv_ring_restore(ring.attn(rt), host.p, fmt, b0, b1, R, s, nullptr); const int64_t width = cap; std::vector hids((size_t) cap), hst = {(int32_t) (N - 1), (int32_t) N, (int32_t) (N / 4), (int32_t) width}; for (int64_t i = 0; i < width; ++i) hids[i] = (int32_t) (N - width + i); // the window's last cells ck(cudaMemcpy(ids, hids.data(), hids.size() * 4, cudaMemcpyHostToDevice), "ids"); ck(cudaMemcpy(steps, hst.data(), 16, cudaMemcpyHostToDevice), "steps"); k::qsa_decode_attn_batch(q, ref.attn(ident), ids, steps, cap, s, scratch, out_ref, 1, nullptr); k::qsa_decode_attn_batch(q, ring.attn(rt), ids, steps, cap, s, scratch, out_str, 1, nullptr); ck(cudaDeviceSynchronize(), "ring"); ck(cudaMemcpy(a.data(), out_ref, (size_t) NH * D * 4, cudaMemcpyDeviceToHost), "a"); ck(cudaMemcpy(b2.data(), out_str, (size_t) NH * D * 4, cudaMemcpyDeviceToHost), "b"); const bool ok = std::memcmp(a.data(), b2.data(), (size_t) NH * D * 4) == 0; std::printf(" %s ring restore: %s\n", name, ok ? "identical" : "DIFFERS"); if (!ok) ++bad; } return bad == 0; } } // namespace int main() { std::printf("kv_stream_parity: streamed vs resident KV, bitwise\n"); const bool a = run(k::kKvInt8), b = run(k::kKvF16), c = run(k::kKvQ4); if (!a || !b || !c) ++g_fail; std::printf(g_fail ? "FAIL\n" : "PASS\n"); return g_fail ? 1 : 0; }