File size: 12,965 Bytes
bbb6388 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 | // 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 <cuda_runtime.h>
#include <algorithm>
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <random>
#include <set>
#include <vector>
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 <typename T> 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 <typename T> 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<uint8_t>(b) : dalloc<uint8_t>(b);
p.v_q4 = host ? halloc<uint8_t>(b) : dalloc<uint8_t>(b);
} else if (fmt == k::kKvInt8) {
p.k_q = host ? halloc<int8_t>(rows * s.head_dim) : dalloc<int8_t>(rows * s.head_dim);
p.v_q = host ? halloc<int8_t>(rows * s.head_dim) : dalloc<int8_t>(rows * s.head_dim);
p.k_scale = host ? halloc<uint16_t>(rows * 4) : dalloc<uint16_t>(rows * 4);
p.v_scale = host ? halloc<uint16_t>(rows * 4) : dalloc<uint16_t>(rows * 4);
} else {
p.k_pool = host ? halloc<uint16_t>(rows * s.head_dim) : dalloc<uint16_t>(rows * s.head_dim);
p.v_pool = host ? halloc<uint16_t>(rows * s.head_dim) : dalloc<uint16_t>(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<int32_t> selection(int64_t n_kv, int64_t width, std::mt19937& rng) {
std::vector<int32_t> 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<int64_t> blocks;
const int64_t want_full = (width - tail) / 4, part = (width - tail) % 4;
std::uniform_int_distribution<int64_t> any(0, n_bid - 1), recent(std::max<int64_t>(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<float> 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<int32_t>(n_blocks);
{
std::vector<int32_t> 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<int32_t>(n_blocks);
m.slot_block = dalloc<int32_t>(n_slots); m.slot_stamp = dalloc<int32_t>(n_slots); m.slot_ref = dalloc<int32_t>(n_slots);
m.miss_block = dalloc<int32_t>(n_slots); m.miss_slot = dalloc<int32_t>(n_slots);
m.ctl = dalloc<int32_t>(k::kKvCtlInts);
m.n_blocks = n_blocks; m.n_slots = n_slots;
k::kv_stream_reset(m, nullptr);
float* kc = dalloc<float>(H * D);
float* vc = dalloc<float>(H * D);
int32_t* step = dalloc<int32_t>(k::kStepCount);
int32_t* steps = dalloc<int32_t>(NQ * k::kStepCount);
int32_t* ids = dalloc<int32_t>(NQ * cap);
float* q = dalloc<float>(NQ * NH * D);
const uint64_t scr = k::qsa_decode_attn_scratch_floats(cap, s);
float* scratch = dalloc<float>(NQ * scr);
float* out_ref = dalloc<float>(NQ * NH * D);
float* out_str = dalloc<float>(NQ * NH * D);
std::vector<float> hk(H * D), hv(H * D), hq(NQ * NH * D), a((size_t) NQ * NH * D), b2((size_t) NQ * NH * D);
std::vector<int32_t> 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<int32_t> hids((size_t) NQ * cap, 0), hst(NQ * 4, 0);
for (int t = 0; t < n_q; ++t) {
const int64_t p = std::max<int64_t>(0, pos - (n_q - 1 - t)), n_kv = p + 1;
const int64_t width = std::min<int64_t>(n_kv, cap);
const std::vector<int32_t> 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<int32_t>(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<int32_t> 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;
}
|