// src/kernels/sampler_parity.cpp - P2.S2's test for the sampler chain. // // THE CHECK THAT MATTERS IS THAT THE ORDER IS OBSERVABLE. A sampler in the wrong order still returns a valid // token, so "it produced a token" proves nothing; and because temperature is MONOTONIC it does not change which // tokens `top_k` keeps, so a top-k-only fixture cannot see the order either. What it changes is `top_p`'s CUT: // at T < 1 the distribution sharpens, the cumulative mass reaches p sooner, and fewer tokens survive. // // So this builds a fixture where that happens, computes the greedy pick under BOTH orders, and requires them to // DIFFER - then requires the kernel to agree with the specified one. Without the first half, the test would // pass against either order. #include "strata/kernels/sampler.hpp" #include #include #include #include #include #include #include #include #include #include namespace { void check(cudaError_t e, const char* what) { if (e != cudaSuccess) { std::fprintf(stderr, "%s: %s\n", what, cudaGetErrorString(e)); std::exit(1); } } // The host reference for the specified order: top_k -> top_p -> temperature -> argmax. // `temp_first` swaps the first and last stages, which is the intuitive-but-wrong order. int reference_pick(const std::vector& l, const strata::kernels::SamplerParams& p, bool temp_first) { const int nv = (int) l.size(); const float inv_t = p.temperature > 0.0f ? 1.0f / p.temperature : 0.0f; auto val = [&](int v) { return temp_first ? l[(size_t) v] * inv_t : l[(size_t) v]; }; std::vector ids; const int k = p.top_k > 0 ? p.top_k : nv; std::vector taken((size_t) nv, 0); for (int i = 0; i < k; ++i) { int best = -1; float bv = 0; for (int v = 0; v < nv; ++v) { if (taken[(size_t) v]) continue; if (best < 0 || val(v) > bv) { best = v; bv = val(v); } } taken[(size_t) best] = 1; ids.push_back(best); } if (p.top_p < 1.0f) { float mx = val(ids[0]); for (int v : ids) mx = std::fmax(mx, val(v)); double sum = 0; for (int v : ids) sum += std::exp((double) val(v) - (double) mx); double cum = 0; int cut = (int) ids.size(); for (size_t i = 0; i < ids.size(); ++i) { cum += std::exp((double) val(ids[i]) - (double) mx) / sum; if (cum >= (double) p.top_p) { cut = (int) i + 1; break; } } if (cut < p.min_keep) cut = p.min_keep < (int) ids.size() ? p.min_keep : (int) ids.size(); ids.resize((size_t) cut); } return ids[0]; // greedy: the largest SURVIVING logit, and `ids` is in descending order } // The number of survivors AFTER top_p, in the given order - the quantity the order actually changes. int reference_cut(const std::vector& l, const strata::kernels::SamplerParams& p, bool temp_first) { const int nv = (int) l.size(); const float inv_t = p.temperature > 0.0f ? 1.0f / p.temperature : 0.0f; auto val = [&](int v) { return temp_first ? l[(size_t) v] * inv_t : l[(size_t) v]; }; const int k = p.top_k > 0 ? p.top_k : nv; std::vector ids; std::vector taken((size_t) nv, 0); for (int i = 0; i < k; ++i) { int best = -1; float bv = 0; for (int v = 0; v < nv; ++v) { if (taken[(size_t) v]) continue; if (best < 0 || val(v) > bv) { best = v; bv = val(v); } } taken[(size_t) best] = 1; ids.push_back(best); } if (p.top_p < 1.0f) { float mx = val(ids[0]); for (int v : ids) mx = std::fmax(mx, val(v)); double sum = 0; for (int v : ids) sum += std::exp((double) val(v) - (double) mx); double cum = 0; int cut = (int) ids.size(); for (size_t i = 0; i < ids.size(); ++i) { cum += std::exp((double) val(ids[i]) - (double) mx) / sum; if (cum >= (double) p.top_p) { cut = (int) i + 1; break; } } if (cut < p.min_keep) cut = p.min_keep < (int) ids.size() ? p.min_keep : (int) ids.size(); return cut; } return (int) ids.size(); } int run(const char* name, const std::vector& logits, int n_tokens, const strata::kernels::SamplerParams& p, const std::vector& want, const std::vector& hist = {}, int hist_len = 0) { float* d_l = nullptr; int* d_o = nullptr; check(cudaMalloc(&d_l, logits.size() * sizeof(float)), "malloc logits"); check(cudaMalloc(&d_o, (size_t) n_tokens * sizeof(int)), "malloc out"); check(cudaMemcpy(d_l, logits.data(), logits.size() * sizeof(float), cudaMemcpyHostToDevice), "copy"); // -1 in every output slot first: a row the kernel leaves unwritten can never match (a verify window reads // every row, so "no output" is a wrong answer, not a skipped one) check(cudaMemset(d_o, 0xFF, (size_t) n_tokens * sizeof(int)), "fill out"); int* d_h = nullptr; if (hist_len > 0) { check(cudaMalloc(&d_h, hist.size() * sizeof(int)), "malloc hist"); check(cudaMemcpy(d_h, hist.data(), hist.size() * sizeof(int), cudaMemcpyHostToDevice), "copy hist"); } strata::kernels::sample_tokens(d_l, n_tokens, (int) (logits.size() / n_tokens), d_h, hist_len, p, d_o, nullptr); std::vector got((size_t) n_tokens); check(cudaMemcpy(got.data(), d_o, got.size() * sizeof(int), cudaMemcpyDeviceToHost), "back"); int bad = 0; for (int t = 0; t < n_tokens; ++t) if (got[(size_t) t] != want[(size_t) t]) ++bad; std::printf(" %-34s %s (%d of %d differ)", name, bad ? "*** WRONG ***" : "matches", bad, n_tokens); if (bad) std::printf(" first: want %d got %d", want[0], got[0]); std::printf("\n"); cudaFree(d_l); cudaFree(d_o); if (d_h) cudaFree(d_h); return bad; } // The Philox draw, host side - a transcription of the kernel's `philox_uniform` so the SAMPLED pick (not // just the greedy argmax) can be pinned against a reference. `__umulhi(a, b)` is the high half of a 32x32 // multiply, spelled `(uint32_t)(((uint64_t) a * b) >> 32)` here. struct PhiloxRound { uint32_t& c0; uint32_t& c1; uint32_t& c2; uint32_t& c3; void step(uint32_t k0, uint32_t k1) const { const uint32_t hi0 = (uint32_t) (((uint64_t) 0x9E3779B9u * c0) >> 32); const uint32_t hi1 = (uint32_t) (((uint64_t) 0xBB67AE85u * c2) >> 32); const uint32_t lo0 = 0x9E3779B9u * c0; const uint32_t lo1 = 0xBB67AE85u * c2; const uint32_t n0 = hi1 ^ c1 ^ k0; const uint32_t n1 = lo1; const uint32_t n2 = hi0 ^ c3 ^ k1; const uint32_t n3 = lo0; c0 = n0; c1 = n1; c2 = n2; c3 = n3; } }; float host_philox_uniform(uint64_t seed, uint64_t counter) { uint32_t c0 = (uint32_t) counter, c1 = (uint32_t) (counter >> 32); uint32_t c2 = (uint32_t) seed, c3 = (uint32_t) (seed >> 32); PhiloxRound r{c0, c1, c2, c3}; for (int i = 0; i < 10; ++i) r.step((uint32_t) i, 0u); return (float) (c0 >> 8) * (1.0f / 16777216.0f); } // The full SAMPLED chain, host side - the kernel's `sampler_kernel` in serial form, in llama.cpp's order: // penalties on the raw logits during the top_k selection (ties to the lowest index), top_p's cut in double over // the top_k list, the min_p prefix cut on its survivors, the temperature, and one Philox draw at // (seed, counter + row). One penalties stage (issue #53: this reference used to repeat the kernel's second one). int sampled_reference(const std::vector& l, const std::vector& hist, const strata::kernels::SamplerParams& p, int row) { auto penal = [&](float logit, int count) { if (count <= 0) return logit; if (logit <= 0.0f) logit *= p.penalty_repeat; else logit /= p.penalty_repeat; logit -= (float) count * p.penalty_freq + (count > 0 ? 1.0f : 0.0f) * p.penalty_present; return logit; }; auto count = [&](int v) { int c = 0; for (int h : hist) if (h == v) ++c; return c; }; const int nv = (int) l.size(); const int KMAX = 64; // 1..64 as given; 0 (off) and wider keep the widest list, 64 const int k = std::min(nv, (p.top_k > 0 && p.top_k < KMAX) ? p.top_k : KMAX); std::vector sel_ids; std::vector sel_logit; std::vector taken((size_t) nv, 0); for (int i = 0; i < k; ++i) { int best = -1; float bv = 0; for (int v = 0; v < nv; ++v) { if (taken[(size_t) v]) continue; const float s = penal(l[(size_t) v], count(v)); if (best < 0 || s > bv) { best = v; bv = s; } } taken[(size_t) best] = 1; sel_ids.push_back(best); sel_logit.push_back(bv); } // llama.cpp's order (issue #53): top_p over the whole top_k list, then min_p on its survivors const int n_sel = (int) sel_ids.size(); int n_keep = n_sel; if (p.top_p < 1.0f) { double sum = 0.0; for (int i = 0; i < n_sel; ++i) sum += std::exp((double) sel_logit[(size_t) i] - (double) sel_logit[0]); double cum = 0.0; int cut = n_sel; for (int i = 0; i < n_sel; ++i) { cum += std::exp((double) sel_logit[(size_t) i] - (double) sel_logit[0]) / sum; if (cum >= (double) p.top_p) { cut = i + 1; break; } } if (cut < p.min_keep) cut = p.min_keep < n_sel ? p.min_keep : n_sel; n_keep = cut; } if (p.min_p > 0.0f) { const float thresh = sel_logit[0] + std::log(p.min_p); for (int i = 0; i < n_keep; ++i) if (sel_logit[(size_t) i] < thresh) { n_keep = i; break; } } const float inv_t = p.temperature > 0.0f ? 1.0f / p.temperature : 0.0f; auto scaled = [&](int i) { return sel_logit[(size_t) i] * inv_t; }; // one penalties stage, before (#53) float smx = scaled(0); for (int i = 1; i < n_keep; ++i) smx = std::fmax(smx, scaled(i)); double sum = 0.0; for (int i = 0; i < n_keep; ++i) sum += std::exp((double) scaled(i) - (double) smx); const float u = host_philox_uniform(p.seed, p.counter + (uint64_t) row); double cum = 0.0; int pick = sel_ids[(size_t) (n_keep - 1)]; for (int i = 0; i < n_keep; ++i) { cum += std::exp((double) scaled(i) - (double) smx) / sum; if ((double) u < cum) { pick = sel_ids[(size_t) i]; break; } } return pick; } // Survivors after the min_p + top_p cuts, in the sampled chain - the quantity an order or a threshold // actually changes, used to assert a fixture can SEE the feature before asserting the kernel matches. int sampled_cut(const std::vector& l, const std::vector& hist, const strata::kernels::SamplerParams& p) { auto penal = [&](float logit, int count) { if (count <= 0) return logit; if (logit <= 0.0f) logit *= p.penalty_repeat; else logit /= p.penalty_repeat; logit -= (float) count * p.penalty_freq + (count > 0 ? 1.0f : 0.0f) * p.penalty_present; return logit; }; auto count = [&](int v) { int c = 0; for (int h : hist) if (h == v) ++c; return c; }; const int nv = (int) l.size(); const int KMAX = 64; const int k = std::min(nv, (p.top_k > 0 && p.top_k < KMAX) ? p.top_k : KMAX); std::vector sel; std::vector taken((size_t) nv, 0); for (int i = 0; i < k; ++i) { int best = -1; float bv = 0; for (int v = 0; v < nv; ++v) { if (taken[(size_t) v]) continue; const float s = penal(l[(size_t) v], count(v)); if (best < 0 || s > bv) { best = v; bv = s; } } taken[(size_t) best] = 1; sel.push_back(bv); } int n_minp = (int) sel.size(); if (p.min_p > 0.0f) { const float thresh = sel[0] + std::log(p.min_p); for (int i = 0; i < (int) sel.size(); ++i) if (sel[(size_t) i] < thresh) { n_minp = i; break; } } if (p.top_p >= 1.0f) return n_minp; double sum = 0.0; for (int i = 0; i < n_minp; ++i) sum += std::exp((double) sel[(size_t) i] - (double) sel[0]); double cum = 0.0; int cut = n_minp; for (int i = 0; i < n_minp; ++i) { cum += std::exp((double) sel[(size_t) i] - (double) sel[0]) / sum; if (cum >= (double) p.top_p) { cut = i + 1; break; } } if (cut < p.min_keep) cut = p.min_keep < n_minp ? p.min_keep : n_minp; return cut; } // ---- the kernel's own semantics, for the fixtures ---- // // `sampled_reference` picks the first unpicked logit even when it is -inf or NaN; the kernels never pick either, and // a round that finds nothing stores id 0 with a -inf logit (and later rounds treat id 0 as taken, as `sampler_kernel` // does). The mirror below follows the kernels, so rows with -inf, NaN, +inf and more requested than finite logits // can be pinned exactly. On rows without those it is `sampled_reference`. struct SelList { std::vector ids; std::vector logit; }; // The top_k list of `sampler_kernel` for one row: `window` is the counted history (the row's last penalty_last_n). SelList mirror_select(const float* l, int nv, const std::vector& window, const strata::kernels::SamplerParams& p, int k) { std::vector s((size_t) nv); for (int v = 0; v < nv; ++v) { int c = 0; for (int h : window) c += h == v; float x = l[v]; if (c > 0) { if (x <= 0.0f) x *= p.penalty_repeat; else x /= p.penalty_repeat; x -= (float) c * p.penalty_freq + (c > 0 ? 1.0f : 0.0f) * p.penalty_present; } s[(size_t) v] = x; } SelList out; std::vector taken((size_t) nv, 0); for (int i = 0; i < k; ++i) { int best = nv; float bv = -std::numeric_limits::infinity(); for (int v = 0; v < nv; ++v) if (!taken[(size_t) v] && s[(size_t) v] > bv) { bv = s[(size_t) v]; best = v; } const int id = best < nv ? best : 0; taken[(size_t) id] = 1; out.ids.push_back(id); out.logit.push_back(bv); } return out; } // The tail of `sampler_kernel` over a list (its first `k` entries): top_p, min_p, temperature, the Philox draw. int mirror_pick(const SelList& sel, int k, const strata::kernels::SamplerParams& p, int row) { int n_keep = k; float mx = sel.logit[0]; for (int i = 1; i < k; ++i) mx = std::fmax(mx, sel.logit[(size_t) i]); if (p.top_p < 1.0f) { double sum = 0.0; for (int i = 0; i < k; ++i) sum += std::exp((double) sel.logit[(size_t) i] - (double) mx); double cum = 0.0; int cut = k; for (int i = 0; i < k; ++i) { cum += std::exp((double) sel.logit[(size_t) i] - (double) mx) / sum; if (cum >= (double) p.top_p) { cut = i + 1; break; } } if (cut < p.min_keep) cut = p.min_keep < k ? p.min_keep : k; n_keep = cut; } if (p.min_p > 0.0f) { const float thresh = sel.logit[0] + std::log(p.min_p); for (int i = 0; i < n_keep; ++i) if (sel.logit[(size_t) i] < thresh) { n_keep = i; break; } } const float inv_t = p.temperature > 0.0f ? 1.0f / p.temperature : 0.0f; auto scaled = [&](int i) { return sel.logit[(size_t) i] * inv_t; }; float smx = scaled(0); for (int i = 1; i < n_keep; ++i) smx = std::fmax(smx, scaled(i)); double sum = 0.0; for (int i = 0; i < n_keep; ++i) sum += std::exp((double) scaled(i) - (double) smx); const float u = host_philox_uniform(p.seed, p.counter + (uint64_t) row); double cum = 0.0; int pick = sel.ids[(size_t) (n_keep > 0 ? n_keep - 1 : 0)]; for (int i = 0; i < n_keep; ++i) { cum += std::exp((double) scaled(i) - (double) smx) / sum; if ((double) u < cum) { pick = sel.ids[(size_t) i]; break; } } return pick; } int sampled_k(int top_k, int nv) { return std::min(nv, (top_k > 0 && top_k < 64) ? top_k : 64); } // Rows uploaded once and sampled under many parameter sets: `sample` returns the picks of one launch on `stream` // (nullptr: the legacy stream), -1 prefilled as in `run`. struct DeviceRows { float* l = nullptr; int* h = nullptr; int* o = nullptr; int n_tokens = 0, nv = 0, hist_len = 0; DeviceRows(const std::vector& logits, int n_tokens_, const std::vector& hist, int hist_len_) : n_tokens(n_tokens_), nv((int) (logits.size() / (size_t) n_tokens_)), hist_len(hist_len_) { check(cudaMalloc(&l, logits.size() * sizeof(float)), "malloc logits"); check(cudaMalloc(&o, (size_t) n_tokens * sizeof(int)), "malloc out"); check(cudaMemcpy(l, logits.data(), logits.size() * sizeof(float), cudaMemcpyHostToDevice), "copy"); if (hist_len > 0) { check(cudaMalloc(&h, hist.size() * sizeof(int)), "malloc hist"); check(cudaMemcpy(h, hist.data(), hist.size() * sizeof(int), cudaMemcpyHostToDevice), "copy hist"); } } DeviceRows(const DeviceRows&) = delete; DeviceRows& operator=(const DeviceRows&) = delete; ~DeviceRows() { cudaFree(l); cudaFree(o); if (h) cudaFree(h); } std::vector sample(const strata::kernels::SamplerParams& p, cudaStream_t stream) { check(cudaMemset(o, 0xFF, (size_t) n_tokens * sizeof(int)), "fill out"); strata::kernels::sample_tokens(l, n_tokens, nv, h, hist_len, p, o, stream); if (stream != nullptr) check(cudaStreamSynchronize(stream), "stream sync"); std::vector got((size_t) n_tokens); check(cudaMemcpy(got.data(), o, got.size() * sizeof(int), cudaMemcpyDeviceToHost), "back"); return got; } }; // The counted window of row t: the last min(penalty_last_n, hist_len) entries (none without penalties). std::vector window_of(const std::vector& hist, int hist_len, int t, int last_n) { if (hist_len <= 0 || last_n <= 0) return {}; const int h = std::min(last_n, hist_len); const int* row = hist.data() + (size_t) t * hist_len; return std::vector(row + (hist_len - h), row + hist_len); } // `--bench`: the sampled path alone at the engine's vocabulary, per call, on the path the environment selects. void bench_sampled() { const int NV = 248320; cudaStream_t s = nullptr; check(cudaStreamCreate(&s), "stream"); std::mt19937 rng(20); std::normal_distribution g(0.0f, 3.0f); for (int T : {1, 4, 8}) { std::vector l((size_t) NV * T); for (auto& v : l) v = g(rng); float* d_l = nullptr; int* d_o = nullptr; check(cudaMalloc(&d_l, l.size() * sizeof(float)), "bench logits"); check(cudaMalloc(&d_o, (size_t) T * sizeof(int)), "bench out"); check(cudaMemcpy(d_l, l.data(), l.size() * sizeof(float), cudaMemcpyHostToDevice), "bench copy"); for (int k : {20, 64}) { strata::kernels::SamplerParams p; p.top_k = k; p.top_p = 0.95f; p.temperature = 0.7f; p.seed = 1; for (int w = 0; w < 3; ++w) strata::kernels::sample_tokens(d_l, T, NV, nullptr, 0, p, d_o, s); check(cudaStreamSynchronize(s), "bench warmup"); cudaEvent_t e0, e1; check(cudaEventCreate(&e0), "event"); check(cudaEventCreate(&e1), "event"); const int iters = 50; check(cudaEventRecord(e0, s), "record"); for (int it = 0; it < iters; ++it) { p.counter = (uint64_t) it; strata::kernels::sample_tokens(d_l, T, NV, nullptr, 0, p, d_o, s); } check(cudaEventRecord(e1, s), "record"); check(cudaEventSynchronize(e1), "bench sync"); float ms = 0.0f; check(cudaEventElapsedTime(&ms, e0, e1), "elapsed"); std::printf(" bench: n_vocab %d, rows %d, top_k %2d, top_p 0.95: %8.1f us per call\n", NV, T, k, 1000.0 * (double) ms / iters); cudaEventDestroy(e0); cudaEventDestroy(e1); } cudaFree(d_l); cudaFree(d_o); } cudaStreamDestroy(s); } } // namespace int main(int argc, char** argv) { bool selftest = false, bench = false; for (int i = 1; i < argc; ++i) { if (std::string(argv[i]) == "--selftest") selftest = true; else if (std::string(argv[i]) == "--bench") bench = true; else { std::fprintf(stderr, "usage: sampler_parity [--selftest] [--bench]\n"); return 2; } } { // the sampled path under test; ctest runs this binary once per path auto on = [](const char* n) { const char* e = std::getenv(n); return e && *e && std::strcmp(e, "0") != 0; }; std::printf(" sampled path: %s\n", on("STRATA_OLD_SAMPLER") ? "sampler_kernel (STRATA_OLD_SAMPLER)" : on("STRATA_SAMPLER_ONE_BLOCK") ? "one block (STRATA_SAMPLER_ONE_BLOCK)" : "split top_k (default)"); } if (bench) { bench_sampled(); return 0; } int bad = 0; const int NV = 512, NT = 4; // ---- fixture 1: plain greedy. top_k = 0 (disabled), top_p = 1 (disabled), T = 1 -> argmax. { strata::kernels::SamplerParams p; p.top_k = 0; p.top_p = 1.0f; p.temperature = 1.0f; p.greedy = true; std::mt19937 rng(3); std::normal_distribution g(0.0f, 1.0f); std::vector l((size_t) NV * NT); for (auto& v : l) v = g(rng); std::vector want((size_t) NT); for (int t = 0; t < NT; ++t) want[(size_t) t] = reference_pick({l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV}, p, false); bad += run("greedy argmax", l, NT, p, want); } // ---- fixture 2: THE ORDER FIXTURE. T = 0.5 sharpens the distribution enough that top_p = 0.5 cuts // differently before and after the scaling, and the two orders then pick DIFFERENT tokens. { strata::kernels::SamplerParams p; p.top_k = 0; p.top_p = 0.5f; p.temperature = 0.5f; p.min_keep = 1; p.greedy = true; std::vector l((size_t) NV * NT, -1000.0f); for (int t = 0; t < NT; ++t) { // a flat-ish head so the cumulative mass crosses 0.5 inside it, and one clear leader l[(size_t) t * NV + 0] = 3.0f; for (int v = 1; v < 8; ++v) l[(size_t) t * NV + v] = 2.6f - 0.05f * (float) v; } std::vector want((size_t) NT), other((size_t) NT); for (int t = 0; t < NT; ++t) { const std::vector row(l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV); want[(size_t) t] = reference_pick(row, p, false); // the SPECIFIED order other[(size_t) t] = reference_pick(row, p, true); // temperature first } // If the two orders agree on this fixture the test cannot see the order, and saying "the kernel // matches the spec" would be vacuous. // GREEDY CANNOT SEE THE ORDER, and saying otherwise would be a vacuous check: no filter removes the // global argmax, and temperature is monotonic, so the greedy pick is order-independent by // construction. What the order changes is the top_p CUT, so the fixture is asserted to be // order-SENSITIVE at the cut, which is a property of the fixture rather than of the kernel. int cut_spec = 0, cut_alt = 0; for (int t = 0; t < NT; ++t) { const std::vector row(l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV); cut_spec += reference_cut(row, p, false); cut_alt += reference_cut(row, p, true); } const bool distinguishable = (cut_spec != cut_alt); std::printf(" %-34s %s (survivors: spec %d, temp-first %d)\n", "order is observable on this fixture", distinguishable ? "yes" : "*** NO - THE FIXTURE CANNOT SEE THE ORDER ***", cut_spec, cut_alt); if (!distinguishable) ++bad; // greedy is still checked here, but as an ARGMAX check, not an order check bad += run("greedy over this fixture", l, NT, p, want); } // ---- fixture 3: greedy consumes NO random number. Two runs with different seeds must agree, or the // seeded streams diverge between greedy and sampled runs - which docs/sampling.md ยง3 calls out. { strata::kernels::SamplerParams a; a.top_k = 20; a.top_p = 0.95f; a.temperature = 1.0f; a.greedy = true; a.seed = 1; strata::kernels::SamplerParams b = a; b.seed = 999999; std::mt19937 rng(5); std::normal_distribution g(0.0f, 1.0f); std::vector l((size_t) NV * NT); for (auto& v : l) v = g(rng); std::vector wa((size_t) NT); for (int t = 0; t < NT; ++t) wa[(size_t) t] = reference_pick({l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV}, a, false); float* d_l = nullptr; int *d_a = nullptr, *d_b = nullptr; check(cudaMalloc(&d_l, l.size() * sizeof(float)), "m1"); check(cudaMalloc(&d_a, (size_t) NT * sizeof(int)), "m2"); check(cudaMalloc(&d_b, (size_t) NT * sizeof(int)), "m3"); check(cudaMemcpy(d_l, l.data(), l.size() * sizeof(float), cudaMemcpyHostToDevice), "c1"); strata::kernels::sample_tokens(d_l, NT, NV, nullptr, 0, a, d_a, nullptr); strata::kernels::sample_tokens(d_l, NT, NV, nullptr, 0, b, d_b, nullptr); std::vector ga((size_t) NT), gb((size_t) NT); check(cudaMemcpy(ga.data(), d_a, ga.size() * sizeof(int), cudaMemcpyDeviceToHost), "g1"); check(cudaMemcpy(gb.data(), d_b, gb.size() * sizeof(int), cudaMemcpyDeviceToHost), "g2"); int mismatch = 0, wrong = 0; for (int t = 0; t < NT; ++t) { if (ga[(size_t) t] != gb[(size_t) t]) ++mismatch; if (ga[(size_t) t] != wa[(size_t) t]) ++wrong; } std::printf(" %-34s %s (seed-independent: %d differ; vs reference: %d wrong)\n", "greedy ignores the seed", (!mismatch && !wrong) ? "matches" : "*** WRONG ***", mismatch, wrong); bad += mismatch + wrong; cudaFree(d_l); cudaFree(d_a); cudaFree(d_b); } // ---- fixture 4: PENALTIES. Two sub-cases, each built so the rule it tests decides the answer. { // host reference for the penalty stage, transcribed from llama_sampler_penalties_apply auto penal = [](float logit, int count, const strata::kernels::SamplerParams& p) { if (count <= 0) return logit; if (logit <= 0.0f) logit *= p.penalty_repeat; else logit /= p.penalty_repeat; logit -= (float) count * p.penalty_freq + (count > 0 ? 1.0f : 0.0f) * p.penalty_present; return logit; }; auto pick = [&](const std::vector& l, const std::vector& hist, const strata::kernels::SamplerParams& p, bool divide_unconditionally) { int best = 0; float bv = 0; bool first = true; for (int v = 0; v < (int) l.size(); ++v) { int c = 0; for (int h : hist) if (h == v) ++c; float s; if (divide_unconditionally && c > 0) { s = l[(size_t) v] / p.penalty_repeat - (float) c * p.penalty_freq - (c > 0 ? 1.0f : 0.0f) * p.penalty_present; } else { s = penal(l[(size_t) v], c, p); } if (first || s > bv) { bv = s; best = v; first = false; } } return best; }; const int NV2 = 8, NT2 = 2; strata::kernels::SamplerParams p; p.top_k = 0; p.top_p = 1.0f; p.temperature = 1.0f; p.greedy = true; p.penalty_last_n = 4; p.penalty_repeat = 2.0f; // A: ALL logits negative, so the multiply-or-divide rule decides the argmax std::vector la((size_t) NV2 * NT2, -8.0f); for (int t = 0; t < NT2; ++t) { la[(size_t) t * NV2 + 0] = -1.0f; // in the history -> penalised la[(size_t) t * NV2 + 1] = -1.2f; // not penalised -> should win } std::vector hist_a((size_t) NT2 * 4, -1); for (int t = 0; t < NT2; ++t) hist_a[(size_t) t * 4 + 0] = 0; std::vector want_a((size_t) NT2), alt_a((size_t) NT2); for (int t = 0; t < NT2; ++t) { const std::vector row(la.begin() + (size_t) t * NV2, la.begin() + (size_t) (t + 1) * NV2); const std::vector h(hist_a.begin() + (size_t) t * 4, hist_a.begin() + (size_t) (t + 1) * 4); want_a[(size_t) t] = pick(row, h, p, false); alt_a[(size_t) t] = pick(row, h, p, true); // divide unconditionally } const bool A_visible = want_a[0] != alt_a[0]; std::printf(" %-34s %s (multiply-rule %d, divide-always %d)\n", "multiply-or-divide is observable", A_visible ? "yes" : "*** NO ***", want_a[0], alt_a[0]); if (!A_visible) ++bad; else bad += run("penalties: repeat on negatives", la, NT2, p, want_a, hist_a, 4); // B: the PRESENCE penalty is a boolean, so two occurrences cost the same as one. The runner's margin // is inside the difference between one and two applications. std::vector lb((size_t) NV2 * NT2, -8.0f); for (int t = 0; t < NT2; ++t) { lb[(size_t) t * NV2 + 0] = 5.0f; // seen twice -> penalised ONCE (presence) + freq*2 lb[(size_t) t * NV2 + 1] = 3.4f; // unseen } std::vector hist_b((size_t) NT2 * 4, -1); for (int t = 0; t < NT2; ++t) { hist_b[(size_t) t * 4 + 0] = 0; hist_b[(size_t) t * 4 + 1] = 0; // twice } strata::kernels::SamplerParams q = p; q.penalty_present = 1.5f; q.penalty_freq = 0.0f; std::vector want_b((size_t) NT2); for (int t = 0; t < NT2; ++t) { const std::vector row(lb.begin() + (size_t) t * NV2, lb.begin() + (size_t) (t + 1) * NV2); const std::vector h(hist_b.begin() + (size_t) t * 4, hist_b.begin() + (size_t) (t + 1) * 4); want_b[(size_t) t] = pick(row, h, q, false); } std::printf(" %-34s want token %d (with present=1.5, token 0 goes 5.0/2 - 1.5 = 1.0 vs token 1 at " "3.4)\n", "presence penalty is a boolean", want_b[0]); bad += run("penalties: presence is boolean", lb, NT2, q, want_b, hist_b, 4); } // ---- fixture 5: TEMPERATURE 0 MUST STILL RETURN THE ARGMAX. Regression test for a real bug. // // The greedy branch used to read `apply_penalties(l[v] * inv_t, ...)`. `inv_t` is 0.0f whenever // temperature <= 0, so at temperature 0 EVERY logit became 0.0f and the argmax returned index 0 - the // sampler emitted token 0 forever, whatever the model predicted. OpenAI clients send `temperature: 0` // for greedy decoding, so this was reachable from any ordinary client. // // It survived because EVERY other greedy fixture in this file sets temperature = 1.0f, where inv_t = 1.0 // and the extra multiply is harmless. The bug needs temperature <= 0 to appear, and no fixture used it. // The fixture below makes token 0 the WORST token in every row, so returning 0 is unambiguously wrong. { const int NV3 = 512, NT3 = 4; std::vector l((size_t) NV3 * NT3, -5.0f); std::vector want((size_t) NT3); for (int t = 0; t < NT3; ++t) { const int best = 100 + t; // the argmax is never token 0 l[(size_t) t * NV3 + best] = 3.0f; l[(size_t) t * NV3 + 0] = -9.0f; // token 0 is the worst in the row want[(size_t) t] = best; } strata::kernels::SamplerParams p0; p0.top_k = 0; p0.top_p = 1.0f; p0.temperature = 0.0f; p0.greedy = false; bad += run("T=0 greedy=false is the argmax", l, NT3, p0, want); strata::kernels::SamplerParams p1 = p0; p1.greedy = true; bad += run("T=0 greedy=true is the argmax", l, NT3, p1, want); strata::kernels::SamplerParams p2 = p0; p2.greedy = true; p2.temperature = 1.0f; bad += run("T=1 greedy=true is the argmax", l, NT3, p2, want); } // ---- fixture 6: PENALTIES IN THE SAMPLED CHAIN. Fixture 4 pins the greedy (argmax) path; the sampled // chain gets its own reference (the full chain with the host Philox) and its own observability check: with // the penalties on, the history row's favourite must LOSE a pick it would win penalty-free. { const int NV2 = 8, NT2 = 2; strata::kernels::SamplerParams p; p.top_k = 5; p.top_p = 0.9f; p.temperature = 0.8f; p.seed = 9; p.counter = 0; p.penalty_last_n = 4; p.penalty_repeat = 3.0f; p.penalty_freq = 0.2f; p.penalty_present = 0.6f; std::vector l((size_t) NV2 * NT2, -8.0f); std::vector hist((size_t) NT2 * 4, -1); for (int t = 0; t < NT2; ++t) { float* row = l.data() + (size_t) t * NV2; row[0] = 9.0f; row[1] = 4.5f; row[2] = 4.4f; row[3] = 4.3f; // token 0 leads clean (9 vs 4.5) hist[(size_t) t * 4 + 0] = 0; // and falls to 2.2/2.0 penalised hist[(size_t) t * 4 + 1] = t == 1 ? 0 : -1; // (repeat 3, freq, presence) } std::vector want((size_t) NT2), clean((size_t) NT2); strata::kernels::SamplerParams clean_p = p; clean_p.penalty_last_n = 0; clean_p.penalty_repeat = 1.0f; clean_p.penalty_freq = 0.0f; clean_p.penalty_present = 0.0f; for (int t = 0; t < NT2; ++t) { const std::vector row(l.begin() + (size_t) t * NV2, l.begin() + (size_t) (t + 1) * NV2); const std::vector h(hist.begin() + (size_t) t * 4, hist.begin() + (size_t) (t + 1) * 4); want[(size_t) t] = sampled_reference(row, h, p, t); clean[(size_t) t] = sampled_reference(row, h, clean_p, t); } const bool visible = want[0] != clean[0] || want[1] != clean[1]; std::printf(" %-34s %s (penalised picks %d/%d, clean %d/%d)\n", "sampled penalties are observable", visible ? "yes" : "*** NO ***", want[0], want[1], clean[0], clean[1]); if (!visible) ++bad; else bad += run("sampled chain: penalties + top_k/p", l, NT2, p, want, hist, 4); } // ---- fixture 7: MIN_P. The cut is a PREFIX of the descending top_k list (logit >= max + log(min_p)), // so the fixture asserts the survivor count moves with the threshold (the observability half) and that // the kernel's pick equals the reference's through the full sampled chain (the correctness half). { const int NV3 = 8, NT3 = 2; std::vector l((size_t) NV3 * NT3, -8.0f); for (int t = 0; t < NT3; ++t) { float* row = l.data() + (size_t) t * NV3; row[0] = 4.0f; row[1] = 3.5f; row[2] = 3.2f; row[3] = 3.1f; // gaps keep the cut off the row[4] = 2.0f; // logf/rounding knife edge } strata::kernels::SamplerParams base; base.top_k = 6; base.top_p = 1.0f; base.temperature = 0.9f; base.seed = 77; int c0 = 0, c05 = 0, c09 = 0; for (int t = 0; t < NT3; ++t) { const std::vector row(l.begin() + (size_t) t * NV3, l.begin() + (size_t) (t + 1) * NV3); strata::kernels::SamplerParams q = base; q.min_p = 0.0f; c0 += sampled_cut(row, {}, q); q.min_p = 0.5f; c05 += sampled_cut(row, {}, q); q.min_p = 0.9f; c09 += sampled_cut(row, {}, q); } const bool visible = c0 > c05 && c05 > c09 && c09 >= NT3; std::printf(" %-34s %s (survivors: min_p 0 -> %d, 0.5 -> %d, 0.9 -> %d)\n", "min_p cut is observable", visible ? "yes" : "*** NO ***", c0, c05, c09); if (!visible) ++bad; for (float mp : {0.0f, 0.5f, 0.9f}) { strata::kernels::SamplerParams q = base; q.min_p = mp; std::vector want((size_t) NT3); for (int t = 0; t < NT3; ++t) { const std::vector row(l.begin() + (size_t) t * NV3, l.begin() + (size_t) (t + 1) * NV3); want[(size_t) t] = sampled_reference(row, {}, q, t); } char name[64]; std::snprintf(name, sizeof name, "sampled chain: min_p=%.1f", (double) mp); bad += run(name, l, NT3, q, want); } } // ---- fixture 8: THE PENALTY WINDOW IS THE TAIL. With an 8-entry history and penalty_last_n = 4, only // the LAST four entries count: a token punished in the old half must come back to full strength, and one // punished in the tail half stays down. The reference counts the same tail; the observability check runs // the reference once more WITHOUT the clamp (counting all 8) and requires the picks to differ. { const int NV4 = 8; strata::kernels::SamplerParams p; p.top_k = 0; p.top_p = 1.0f; p.temperature = 1.0f; p.greedy = true; p.penalty_last_n = 4; p.penalty_repeat = 3.0f; p.penalty_freq = 0.3f; p.penalty_present = 0.5f; std::vector l((size_t) NV4, -8.0f); l[0] = 6.0f; l[3] = 6.5f; // token 0 leads clean; token 3 is the tail offender std::vector hist = {0, 0, 0, 0, 3, 3, 3, 3}; // token 0 old (out), token 3 in the tail auto pick_clamped = [&](bool clamp) { int best = 0; float bv = 0; bool first = true; for (int v = 0; v < NV4; ++v) { int c = 0; for (int i = 0; i < (clamp ? 4 : 8); ++i) if (hist[(size_t) (8 - (clamp ? 4 : 8) + i)] == v) ++c; float logit = l[(size_t) v]; if (c > 0) { logit = logit <= 0.0f ? logit * p.penalty_repeat : logit / p.penalty_repeat; logit -= (float) c * p.penalty_freq + p.penalty_present; } if (first || logit > bv) { bv = logit; best = v; first = false; } } return best; }; const int want = pick_clamped(true), unclamped = pick_clamped(false); const bool visible = want != unclamped; std::printf(" %-34s %s (clamped pick %d, full-history pick %d)\n", "penalty window clamp is observable", visible ? "yes" : "*** NO ***", want, unclamped); if (!visible) ++bad; else bad += run("penalty window: tail only", {l.begin(), l.end()}, 1, p, {want}, hist, 8); } // ---- fixture 9: ONE PENALTIES STAGE (issue #53), against an independently computed distribution, not the // reference above (which had copied the kernel's mistake). Two tokens with equal logits, token 0 in the // history, presence penalty 1.5, temperature 0.7: llama.cpp's chain gives token 0 the logit (0 - 1.5) / 0.7, // P = 1 / (1 + exp(1.5 / 0.7)) = 0.1050; a second penalty after the temperature made it 0.0255. Over 8,000 // draws the kernel's share of token 0 must be near 0.105 (4 sigma = 0.014), and its picks must equal the // reference's draw for draw. { const int NT9 = 8000; strata::kernels::SamplerParams p; p.top_k = 2; p.top_p = 1.0f; p.min_p = 0.0f; p.temperature = 0.7f; p.seed = 53; p.counter = 0; p.penalty_last_n = 1; p.penalty_repeat = 1.0f; p.penalty_freq = 0.0f; p.penalty_present = 1.5f; std::vector l((size_t) 2 * NT9, 0.0f); std::vector hist((size_t) NT9, 0); std::vector want((size_t) NT9); int zeros = 0; for (int t = 0; t < NT9; ++t) { want[(size_t) t] = sampled_reference({0.0f, 0.0f}, {0}, p, t); zeros += want[(size_t) t] == 0; } const double expect = 1.0 / (1.0 + std::exp(1.5 / 0.7)), share = (double) zeros / NT9; const bool near = std::fabs(share - expect) < 0.014; std::printf(" %-34s %s (token 0 drawn %.4f of %d, expected %.4f; twice-penalised would be 0.0255)\n", "one penalties stage (#53)", near ? "yes" : "*** NO ***", share, NT9, expect); if (!near) ++bad; bad += run("sampled chain: #53's example", l, NT9, p, want, hist, 1); } // ---- fixture 10: TOP_P BEFORE MIN_P (llama.cpp's order). Probabilities 0.4 / 0.3 / 0.2 / 0.1, top_p 0.75, // min_p 0.3 (keeps p >= 0.12): top_p over all four keeps three (0.4 + 0.3 < 0.75 <= 0.9) and min_p keeps them; // min_p first would drop 0.1, renormalise, and top_p would then stop at two (0.444 + 0.333 >= 0.75). So token // 2 must be drawn sometimes - never in the old order - and every pick must equal the reference's. { const int NT10 = 256; strata::kernels::SamplerParams p; p.top_k = 4; p.top_p = 0.75f; p.min_p = 0.3f; p.temperature = 1.0f; p.seed = 10; p.counter = 0; const float lp[4] = {std::log(0.4f), std::log(0.3f), std::log(0.2f), std::log(0.1f)}; std::vector row = {lp[0], lp[1], lp[2], lp[3], -30.0f, -30.0f, -30.0f, -30.0f}; std::vector l; for (int t = 0; t < NT10; ++t) l.insert(l.end(), row.begin(), row.end()); std::vector want((size_t) NT10); int twos = 0; for (int t = 0; t < NT10; ++t) { want[(size_t) t] = sampled_reference(row, {}, p, t); twos += want[(size_t) t] == 2; } std::printf(" %-34s %s (token 2 drawn %d of %d times)\n", "top_p before min_p is observable", twos > 0 ? "yes" : "*** NO ***", twos, NT10); if (twos == 0) ++bad; bad += run("sampled chain: top_p then min_p", l, NT10, p, want); } // ---- fixture 11: A STALE HISTORY WITH last_n = 0 IS INERT (PR #59). A caller can hand over a history buffer // from a previous penalised request while this request disables the penalties - the run must equal the // no-history run in both kernels, and the bitmap the launch did not size must stay untouched. { std::mt19937 rng(11); std::normal_distribution g(0.0f, 1.0f); std::vector l((size_t) NV * NT); for (auto& v : l) v = g(rng); std::vector hist((size_t) NT * 8, -1); for (int t = 0; t < NT; ++t) hist[(size_t) t * 8] = 3; strata::kernels::SamplerParams p; p.top_k = 20; p.top_p = 0.95f; p.temperature = 0.8f; p.seed = 5; std::vector want((size_t) NT); for (int t = 0; t < NT; ++t) want[(size_t) t] = sampled_reference({l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV}, {}, p, t); bad += run("stale history, last_n=0 (sampled)", l, NT, p, want, hist, 8); strata::kernels::SamplerParams gp = p; gp.greedy = true; gp.top_k = 0; gp.top_p = 1.0f; std::vector gwant((size_t) NT); for (int t = 0; t < NT; ++t) gwant[(size_t) t] = reference_pick({l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV}, gp, false); bad += run("stale history, last_n=0 (greedy)", l, NT, gp, gwant, hist, 8); } // ---- fixture 12: ONE HISTORY PER ROW (engine 0.1.19). A verify window samples T rows, and row t's pick // follows the drafts 1..t: its penalties must count them. The engine used to stage row 0 alone, so rows // 1..T-1 read slots nobody wrote. Here every row gets `penalty_rows`' history and is pinned against a // per-row scalar reference; the fixture first asserts it can SEE the difference - with row 0's history // copied to every row (the nearest well-defined stand-in for the old staging) some row must pick differently. // Row t's own newest token (window[t]) is its favourite by a margin the presence penalty overturns. { auto greedy_pen = [](const std::vector& l, const std::vector& hist, const strata::kernels::SamplerParams& p) { int best = 0; float bv = 0; bool first = true; for (int v = 0; v < (int) l.size(); ++v) { int c = 0; for (int h : hist) if (h == v) ++c; float s = l[(size_t) v]; if (c > 0) { if (s <= 0.0f) s *= p.penalty_repeat; else s /= p.penalty_repeat; s -= (float) c * p.penalty_freq + p.penalty_present; } if (first || s > bv) { bv = s; best = v; first = false; } } return best; }; std::mt19937 rng(12); std::normal_distribution g(0.0f, 1.0f); const int TAIL = 5000; // longer than the widest window below std::vector tail((size_t) TAIL); for (auto& v : tail) v = (int32_t) (rng() % NV); int observable = 0, rows_checked = 0; for (int T : {1, 2, 4, 8}) { for (int H : {1, 64, 1024, 4096}) { std::vector window((size_t) T); for (int t = 0; t < T; ++t) window[(size_t) t] = (int32_t) (100 + 37 * t); // distinct, in range std::vector rows((size_t) T * H); strata::kernels::penalty_rows(tail.data(), TAIL, window.data(), T, H, rows.data()); std::vector l((size_t) T * NV); for (auto& v : l) v = g(rng); for (int t = 0; t < T; ++t) l[(size_t) t * NV + window[(size_t) t]] = 5.0f; strata::kernels::SamplerParams gp; gp.greedy = true; gp.temperature = 0.0f; gp.top_k = 0; gp.top_p = 1.0f; gp.penalty_last_n = H; gp.penalty_present = 4.0f; strata::kernels::SamplerParams sp2; sp2.top_k = 20; sp2.top_p = 0.9f; sp2.temperature = 0.7f; sp2.seed = 1000 + (uint64_t) H; sp2.counter = 77; sp2.penalty_last_n = H; sp2.penalty_present = 4.0f; sp2.penalty_repeat = 1.1f; std::vector gwant((size_t) T), swant((size_t) T), rows_int(rows.begin(), rows.end()); for (int t = 0; t < T; ++t) { const std::vector lr(l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV); const std::vector own(rows.begin() + (size_t) t * H, rows.begin() + (size_t) (t + 1) * H); const std::vector row0(rows.begin(), rows.begin() + H); gwant[(size_t) t] = greedy_pen(lr, own, gp); swant[(size_t) t] = sampled_reference(lr, own, sp2, t); if (t > 0 && greedy_pen(lr, row0, gp) != gwant[(size_t) t]) ++observable; ++rows_checked; } char name[64]; std::snprintf(name, sizeof name, "per-row history T=%d H=%d greedy", T, H); bad += run(name, l, T, gp, gwant, rows_int, H); std::snprintf(name, sizeof name, "per-row history T=%d H=%d sampled", T, H); bad += run(name, l, T, sp2, swant, rows_int, H); } } std::printf(" %-34s %s (%d drafted rows pick differently with row 0's history, of %d rows)\n", "per-row histories are observable", observable > 0 ? "yes" : "*** NO ***", observable, rows_checked); if (observable == 0) ++bad; } // ---- fixture 13: HISTORY IDS OUTSIDE THE VOCABULARY ARE IGNORED. The bitmap is sized for n_vocab bits; // an id >= n_vocab used to set a bit past its end (a shared-memory write out of bounds - compute-sanitizer // memcheck reports it). They can never be a candidate, so the result equals the reference without them. { std::mt19937 rng(13); std::normal_distribution g(0.0f, 1.0f); const int H = 16; std::vector l((size_t) NV * NT); for (auto& v : l) v = g(rng); std::vector hist((size_t) NT * H, -1), valid_only((size_t) NT * H, -1); const int junk[] = {NV, NV + 1000, 0x7fffffff, -5, 1 << 20}; for (int t = 0; t < NT; ++t) { for (int j = 0; j < H; ++j) { const bool bogus = j % 3 == 0; const int v = bogus ? junk[(size_t) (j / 3) % 5] : (int) (rng() % NV); hist[(size_t) t * H + j] = v; if (!bogus) valid_only[(size_t) t * H + j] = v; } l[(size_t) t * NV + hist[(size_t) t * H + 1]] = 4.0f; // a penalised favourite, so penalties matter } strata::kernels::SamplerParams gp; gp.greedy = true; gp.temperature = 0.0f; gp.top_k = 0; gp.top_p = 1.0f; gp.penalty_last_n = H; gp.penalty_present = 3.0f; strata::kernels::SamplerParams sp2 = gp; sp2.greedy = false; sp2.temperature = 0.8f; sp2.top_k = 20; sp2.top_p = 0.95f; sp2.seed = 13; std::vector gwant((size_t) NT), swant((size_t) NT); for (int t = 0; t < NT; ++t) { const std::vector lr(l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV); const std::vector ok(valid_only.begin() + (size_t) t * H, valid_only.begin() + (size_t) (t + 1) * H); { // the greedy reference with the penalty (reference_pick has none) int best = 0; float bv = 0; bool first = true; for (int v = 0; v < NV; ++v) { int c = 0; for (int h : ok) if (h == v) ++c; float s = lr[(size_t) v]; if (c > 0) { if (s <= 0.0f) s *= gp.penalty_repeat; else s /= gp.penalty_repeat; s -= gp.penalty_present; } if (first || s > bv) { bv = s; best = v; first = false; } } gwant[(size_t) t] = best; } swant[(size_t) t] = sampled_reference(lr, ok, sp2, t); } bad += run("out-of-vocab history ids (greedy)", l, NT, gp, gwant, hist, H); bad += run("out-of-vocab history ids (sampled)", l, NT, sp2, swant, hist, H); } // ---- fixture 14: THE top_k CONTRACT. 1..64 as given; 0 ("off") and anything wider use the widest list the // kernel keeps, 64 - and every row is written (the sampled kernel used to print an error for 0 and leave the // row unwritten, which a verify window then read as a token; `run` pre-fills -1 so that fails here). { std::mt19937 rng(14); std::normal_distribution g(0.0f, 1.0f); std::vector l((size_t) NV * NT); for (auto& v : l) v = g(rng) * 0.3f; // flat: the 64-wide list matters to the draw strata::kernels::SamplerParams p64; p64.top_k = 64; p64.top_p = 1.0f; p64.temperature = 1.5f; p64.seed = 14; std::vector want((size_t) NT); for (int t = 0; t < NT; ++t) want[(size_t) t] = sampled_reference({l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV}, {}, p64, t); bad += run("sampled top_k=64", l, NT, p64, want); strata::kernels::SamplerParams p0 = p64; p0.top_k = 0; bad += run("sampled top_k=0 means 64", l, NT, p0, want); strata::kernels::SamplerParams p100 = p64; p100.top_k = 100; bad += run("sampled top_k=100 means 64", l, NT, p100, want); strata::kernels::SamplerParams pneg = p64; pneg.top_k = -3; bad += run("sampled top_k=-3 means 64", l, NT, pneg, want); } // ---- fixture 15: `penalty_rows`, host only, against the plain definition: row t = the last h tokens of // tail + window[0..t], -1 padded in front. Covers a tail shorter than, equal to and longer than h, and row 0 // equal to the single row the engine staged before 0.1.19. { int wrong = 0, cases = 0; for (int n_tail : {0, 1, 5, 63, 64, 65, 300}) { for (int T : {1, 3, 8}) { for (int h : {1, 4, 64, 100}) { std::vector tail((size_t) n_tail), window((size_t) T); for (int i = 0; i < n_tail; ++i) tail[(size_t) i] = 1000 + i; for (int i = 0; i < T; ++i) window[(size_t) i] = 5000 + i; std::vector rows((size_t) T * h, 12345); strata::kernels::penalty_rows(tail.data(), n_tail, window.data(), T, h, rows.data()); for (int t = 0; t < T; ++t) { std::vector seq(tail); seq.insert(seq.end(), window.begin(), window.begin() + t + 1); std::vector expect((size_t) h, -1); const int take = (int) std::min((size_t) h, seq.size()); for (int j = 0; j < take; ++j) expect[(size_t) (h - take + j)] = seq[seq.size() - take + j]; if (!std::equal(expect.begin(), expect.end(), rows.begin() + (size_t) t * h)) ++wrong; ++cases; } // row 0 = the old single-row staging: consumed tail, then the fed-back head last std::vector old((size_t) h, -1); const int take0 = (int) std::min(h, (int64_t) n_tail + 1); for (int j = 0; j < take0 - 1; ++j) old[(size_t) (h - take0 + j)] = tail[(size_t) (n_tail - (take0 - 1) + j)]; old[(size_t) (h - 1)] = window[0]; if (!std::equal(old.begin(), old.end(), rows.begin())) ++wrong; ++cases; } } } std::printf(" %-34s %s (%d of %d rows differ)\n", "penalty_rows layout", wrong ? "*** WRONG ***" : "matches", wrong, cases); bad += wrong; } // A continuous stream and individual decode calls consume the same draw counters. { constexpr int count = 32, vocab = 16; std::vector uniform(count * vocab, 0.0f); float* input = nullptr; int* output = nullptr; check(cudaMalloc(&input, uniform.size() * sizeof(float)), "counter logits"); check(cudaMalloc(&output, count * sizeof(int)), "counter output"); check(cudaMemcpy(input, uniform.data(), uniform.size() * sizeof(float), cudaMemcpyHostToDevice), "counter upload"); strata::kernels::SamplerParams p; p.top_k = vocab; p.top_p = 1.0f; p.seed = 123; p.counter = (uint64_t(1) << 32) + 7; strata::kernels::sample_tokens(input, count, vocab, nullptr, 0, p, output, nullptr); std::vector batch(count), singles(count), repeated(count); check(cudaMemcpy(batch.data(), output, count * sizeof(int), cudaMemcpyDeviceToHost), "counter batch"); for (int i = 0; i < count; ++i) { auto one = p; one.counter += i; strata::kernels::sample_tokens(input, 1, vocab, nullptr, 0, one, output + i, nullptr); } check(cudaMemcpy(singles.data(), output, count * sizeof(int), cudaMemcpyDeviceToHost), "counter singles"); strata::kernels::sample_tokens(input, count, vocab, nullptr, 0, p, output, nullptr); check(cudaMemcpy(repeated.data(), output, count * sizeof(int), cudaMemcpyDeviceToHost), "counter repeated"); bool varies = false; for (int i = 1; i < count; ++i) varies |= batch[i] != batch[0]; const bool valid = batch == singles && batch == repeated && varies; std::printf(" sampler draw counter segmentation/repeat: %s\n", valid ? "PASS" : "FAIL"); bad += !valid; cudaFree(input); cudaFree(output); } // ---- fixture 16: THE WHOLE top_k LIST, POSITION BY POSITION, UNDER TIES. A pick shows the list // through one draw; this reads the list itself. Every row holds a +inf logit, so the tail's arithmetic is NaN // (inf - inf), no cut fires and no draw lands: the chain returns its LAST kept entry, sel_ids[k - 1]. Launching // top_k = 1..64 then reads the list one position at a time - its set and its order. The rows make the order // rest on the tie rule: hundreds of logits share the top finite values, spread over every split block, warp and // lane; -0 and +0 tie; -inf and NaN are mixed in; a row has fewer finite logits than 64 (the sentinel id 0 must // come out); a penalised row lands its penalised tokens exactly on other tokens' values. Vocabularies: the // engine's 248,320 (a partial last split block), 100,003, 262,144 (the widest split), 262,145 (one more: the // one-block fallback) and 1,000. { const float inf = std::numeric_limits::infinity(); const float qnan = std::numeric_limits::quiet_NaN(); int wrong = 0, probes = 0, sentinels = 0; for (int nv : {248320, 100003, 262144, 262145, 1000}) { const int T = 4, H = 64; std::mt19937 rng((unsigned) (1600 + nv)); std::normal_distribution g(0.0f, 1.0f); std::vector l((size_t) nv * T); for (auto& v : l) v = std::floor(g(rng) * 4.0f) / 4.0f; // quarter steps: ties everywhere auto spread = [&](int j) { return (int) (((int64_t) j * 7919 + 13) % nv); }; // row 0: 20 logits at 6.25 and 300 at 6.0 over the whole row, two +inf float* r0 = l.data(); for (int j = 0; j < 320; ++j) r0[spread(j)] = j < 20 ? 6.25f : 6.0f; r0[nv / 2] = inf; r0[nv - 1] = inf; // row 1: nothing above 0, every zero signed by its id's parity (-0 at even ids), one +inf float* r1 = l.data() + (size_t) nv; for (int v = 0; v < nv; ++v) { r1[v] = -std::fabs(r1[v]); if (r1[v] == 0.0f) r1[v] = (v & 1) ? 0.0f : -0.0f; } r1[3] = inf; // row 2: -inf everywhere but ten finite logits (two values), two NaN and a +inf: 11 candidates in all float* r2 = l.data() + (size_t) 2 * nv; for (int v = 0; v < nv; ++v) r2[v] = -inf; for (int j = 0; j < 10; ++j) r2[spread(j + 400)] = j < 5 ? 1.0f : 0.5f; r2[spread(500)] = qnan; r2[spread(501)] = qnan; r2[spread(502)] = inf; // row 3: penalties. Window ids on block and warp edges, repeats, and ids outside the vocabulary; the // penalised tokens sit at 9.0, which repeat 2 / freq 0.25 / present 0.5 turns into 4.0 - 0.25 x count // (3.75, 3.5, ...), values the quarter-step logits share. float* r3 = l.data() + (size_t) 3 * nv; std::vector hist((size_t) T * H, -1); int* h3 = hist.data() + (size_t) 3 * H; const int edges[] = {0, 1023, 1024, 4095, 4096, 8191, 8192, 12345, nv / 2 + 1, nv - 2}; int hn = 0; for (int e : edges) if (e < nv && e != nv / 2) { h3[hn++] = e; if (hn % 3 == 0) h3[hn++] = e; // counted twice r3[e] = 9.0f; } h3[hn++] = nv; // ignored: outside h3[hn++] = nv + 77; h3[hn++] = -5; for (int j = 0; j < 30; ++j) r3[spread(j + 600)] = j < 15 ? 3.75f : 3.5f; r3[nv / 2] = inf; // not in the window strata::kernels::SamplerParams base; base.top_p = 1.0f; base.min_p = 0.0f; base.temperature = 0.8f; base.seed = 16; base.penalty_last_n = H; base.penalty_repeat = 2.0f; base.penalty_freq = 0.25f; base.penalty_present = 0.5f; const int kmax = sampled_k(64, nv); std::vector lists; for (int t = 0; t < T; ++t) lists.push_back(mirror_select(l.data() + (size_t) t * nv, nv, window_of(hist, H, t, H), base, kmax)); DeviceRows rows(l, T, hist, H); for (int top_k = 0; top_k <= 66; ++top_k) { strata::kernels::SamplerParams p = base; p.top_k = top_k == 65 ? 100 : top_k == 66 ? -3 : top_k; // 0, 100 and -3 mean 64 p.top_p = (top_k & 1) ? 0.5f : 1.0f; // both tail branches (NaN: no cut) p.counter = (uint64_t) top_k; const int k = sampled_k(p.top_k, nv); const std::vector got = rows.sample(p, nullptr); for (int t = 0; t < T; ++t) { const int want = lists[(size_t) t].ids[(size_t) (k - 1)]; sentinels += (t == 2 && k > 11); ++probes; if (got[(size_t) t] != want) { if (wrong < 8) std::printf(" n_vocab %d row %d top_k %d: position %d want id %d got %d\n", nv, t, p.top_k, k - 1, want, got[(size_t) t]); ++wrong; } } } } // the fixture must reach the sentinel (row 2 has 11 candidates) or it cannot see the "nothing left" rule std::printf(" %-34s %s (%d of %d positions differ; %d sentinel positions)\n", "top_k list under ties", wrong || !sentinels ? "*** WRONG ***" : "matches", wrong, probes, sentinels); bad += wrong + (sentinels == 0); } // ---- fixture 17: SAMPLED DRAWS UNDER TIES. The realistic chain - finite logits on half steps, // so dozens of tokens share each value near the top - through every stage: top_k 1 / 20 / 64, top_p 0.9 / 1, // min_p 0 / 0.05, a hot temperature that spreads the draws over the whole list, penalties off and on (half of // each window on the row's head, so they reorder it). 17 rows (the split's scratch is first sized for 16: this // regrows it), on the legacy stream and on a created one. Observability: some picks must be tokens that tie // with another kept token, or the tie rule would go untested. { const int T = 17, H = 64; cudaStream_t cs = nullptr; check(cudaStreamCreate(&cs), "fixture 17 stream"); int wrong = 0, draws = 0, tied = 0; for (int nv : {248320, 512}) { std::mt19937 rng((unsigned) (1700 + nv)); std::normal_distribution g(0.0f, 1.5f); std::vector l((size_t) nv * T); for (auto& v : l) v = std::floor(g(rng) * 2.0f) / 2.0f; if (nv == 512) for (auto& v : l) v = std::floor(v / 2.0f); // whole steps: a few big tie groups std::vector hist((size_t) T * H); for (int t = 0; t < T; ++t) { const float* row = l.data() + (size_t) t * nv; const float mx = *std::max_element(row, row + nv); std::vector head; for (int v = 0; v < nv; ++v) if (row[v] >= mx - 1.0f) head.push_back(v); for (int j = 0; j < H; ++j) hist[(size_t) t * H + j] = (j & 1) ? (int) (rng() % (unsigned) nv) : head[(size_t) (rng() % (unsigned) head.size())]; } DeviceRows rows(l, T, hist, H); int config = 0; for (int pen = 0; pen < 2; ++pen) { strata::kernels::SamplerParams base; base.temperature = 2.5f; base.seed = 17; base.penalty_last_n = pen ? H : 0; base.penalty_repeat = 1.25f; base.penalty_freq = 0.25f; base.penalty_present = 0.5f; const int kmax = sampled_k(64, nv); std::vector lists; for (int t = 0; t < T; ++t) lists.push_back(mirror_select(l.data() + (size_t) t * nv, nv, window_of(hist, H, t, base.penalty_last_n), base, kmax)); for (int top_k : {1, 20, 64}) for (float top_p : {0.9f, 1.0f}) for (float min_p : {0.0f, 0.05f}) for (cudaStream_t st : {(cudaStream_t) nullptr, cs}) { strata::kernels::SamplerParams p = base; p.top_k = top_k; p.top_p = top_p; p.min_p = min_p; p.counter = (uint64_t) (1000 * ++config); const int k = sampled_k(top_k, nv); const std::vector got = rows.sample(p, st); for (int t = 0; t < T; ++t) { const SelList& sl = lists[(size_t) t]; const int want = mirror_pick(sl, k, p, t); ++draws; int same = 0; for (int i = 0; i < k; ++i) { if (sl.ids[(size_t) i] != want) continue; for (int j = 0; j < k; ++j) same += sl.logit[(size_t) j] == sl.logit[(size_t) i]; break; } tied += same > 1; if (got[(size_t) t] != want) { if (wrong < 8) std::printf(" n_vocab %d row %d top_k %d top_p %.2f min_p %.2f pen %d: " "want %d got %d\n", nv, t, top_k, (double) top_p, (double) min_p, pen, want, got[(size_t) t]); ++wrong; } } } } } cudaStreamDestroy(cs); std::printf(" %-34s %s (%d of %d draws differ; %d picks tie with another kept token)\n", "sampled draws under ties", wrong || !tied ? "*** WRONG ***" : "matches", wrong, draws, tied); bad += wrong + (tied == 0); } // ---- fixture 18: THE DEFAULT PATH'S AUTOMATIC FALLBACKS. (a) A stream under graph capture // (ThreadLocal mode) gets the one-block kernel, penalty bitmap in dynamic shared memory, inside the graph: the // capture must succeed and the replayed graph (twice) must pick the mirror's tokens; the same stream uncaptured // then takes the split path (its scratch is first allocated after the capture) and picks the same. (b) More // rows than a split launch takes (64) fall back to one block; exactly 64 stay split. { int wrong = 0, draws = 0; const int H = 64; auto compare = [&](const char* what, const std::vector& got, const std::vector& lists, int k, const strata::kernels::SamplerParams& p) { for (size_t t = 0; t < got.size(); ++t) { const int want = mirror_pick(lists[t], k, p, (int) t); ++draws; if (got[t] != want) { if (wrong < 8) std::printf(" %s row %zu: want %d got %d\n", what, t, want, got[t]); ++wrong; } } }; auto make = [&](int nv, int T, unsigned seed, std::vector& l, std::vector& hist) { std::mt19937 rng(seed); std::normal_distribution g(0.0f, 1.5f); l.assign((size_t) nv * T, 0.0f); for (auto& v : l) v = std::floor(g(rng) * 2.0f) / 2.0f; hist.assign((size_t) T * H, 0); for (auto& h : hist) h = (int) (rng() % (unsigned) nv); }; { const int nv = 248320, T = 3; std::vector l; std::vector hist; make(nv, T, 1800u, l, hist); strata::kernels::SamplerParams p; p.top_k = 20; p.top_p = 0.9f; p.temperature = 2.5f; p.seed = 18; p.counter = 77; p.penalty_last_n = H; p.penalty_repeat = 1.25f; p.penalty_freq = 0.25f; p.penalty_present = 0.5f; const int k = sampled_k(p.top_k, nv); std::vector lists; for (int t = 0; t < T; ++t) lists.push_back(mirror_select(l.data() + (size_t) t * nv, nv, window_of(hist, H, t, H), p, k)); DeviceRows rows(l, T, hist, H); cudaStream_t cs = nullptr; check(cudaStreamCreate(&cs), "fixture 18 stream"); check(cudaMemset(rows.o, 0xFF, (size_t) T * sizeof(int)), "fixture 18 fill"); cudaGraph_t graph = nullptr; cudaGraphExec_t exec = nullptr; check(cudaStreamBeginCapture(cs, cudaStreamCaptureModeThreadLocal), "begin capture"); strata::kernels::sample_tokens(rows.l, T, nv, rows.h, H, p, rows.o, cs); check(cudaStreamEndCapture(cs, &graph), "end capture"); check(cudaGraphInstantiate(&exec, graph, 0), "instantiate"); for (int replay = 0; replay < 2; ++replay) { check(cudaMemset(rows.o, 0xFF, (size_t) T * sizeof(int)), "fixture 18 refill"); check(cudaGraphLaunch(exec, cs), "graph launch"); check(cudaStreamSynchronize(cs), "graph sync"); std::vector got((size_t) T); check(cudaMemcpy(got.data(), rows.o, got.size() * sizeof(int), cudaMemcpyDeviceToHost), "back"); compare("captured graph", got, lists, k, p); } cudaGraphExecDestroy(exec); cudaGraphDestroy(graph); compare("same stream, uncaptured", rows.sample(p, cs), lists, k, p); cudaStreamDestroy(cs); } for (int T : {64, 70}) { const int nv = 512; std::vector l; std::vector hist; make(nv, T, 1810u + (unsigned) T, l, hist); strata::kernels::SamplerParams p; p.top_k = 64; p.top_p = 1.0f; p.temperature = 2.5f; p.seed = 18; p.counter = (uint64_t) T; const int k = sampled_k(p.top_k, nv); std::vector lists; for (int t = 0; t < T; ++t) lists.push_back(mirror_select(l.data() + (size_t) t * nv, nv, {}, p, k)); DeviceRows rows(l, T, hist, 0); compare(T > 64 ? "70 rows (one block)" : "64 rows (split)", rows.sample(p, nullptr), lists, k, p); } std::printf(" %-34s %s (%d of %d draws differ)\n", "fallbacks: graph capture, row cap", wrong ? "*** WRONG ***" : "matches", wrong, draws); bad += wrong; } std::printf("\nsampler: %d failures\n", bad); if (bad) return 1; if (selftest) std::printf("sampler_parity OK\n"); return 0; }