Download src/kernels/sampler_parity.cpp from WineryLabs/Winery-Strata: direct link, hf CLI and curl.
- Browser
- Download file 71.2 kB
-
https://huggingface.co/WineryLabs/Winery-Strata/resolve/main/src/kernels/sampler_parity.cpp
- Command line
-
hf download hf://WineryLabs/Winery-Strata/src/kernels/sampler_parity.cpp
-
curl -L -o sampler_parity.cpp https://huggingface.co/WineryLabs/Winery-Strata/resolve/main/src/kernels/sampler_parity.cpp
71.2 kB
| // 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. | |
| 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<float>& 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<int> ids; | |
| const int k = p.top_k > 0 ? p.top_k : nv; | |
| std::vector<char> 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<float>& 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<int> ids; | |
| std::vector<char> 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<float>& logits, int n_tokens, const strata::kernels::SamplerParams& p, | |
| const std::vector<int>& want, const std::vector<int>& 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<int> 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<float>& l, const std::vector<int>& 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<int> sel_ids; | |
| std::vector<float> sel_logit; | |
| std::vector<char> 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<float>& l, const std::vector<int>& 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<float> sel; | |
| std::vector<char> 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<int> ids; | |
| std::vector<float> 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<int>& window, const strata::kernels::SamplerParams& p, | |
| int k) { | |
| std::vector<float> 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<char> taken((size_t) nv, 0); | |
| for (int i = 0; i < k; ++i) { | |
| int best = nv; | |
| float bv = -std::numeric_limits<float>::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<float>& logits, int n_tokens_, const std::vector<int>& 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<int> 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<int> 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<int> window_of(const std::vector<int>& 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<int>(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<float> g(0.0f, 3.0f); | |
| for (int T : {1, 4, 8}) { | |
| std::vector<float> 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<float> g(0.0f, 1.0f); | |
| std::vector<float> l((size_t) NV * NT); | |
| for (auto& v : l) v = g(rng); | |
| std::vector<int> 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<float> 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<int> want((size_t) NT), other((size_t) NT); | |
| for (int t = 0; t < NT; ++t) { | |
| const std::vector<float> 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<float> 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<float> g(0.0f, 1.0f); | |
| std::vector<float> l((size_t) NV * NT); | |
| for (auto& v : l) v = g(rng); | |
| std::vector<int> 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<int> 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<float>& l, const std::vector<int>& 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<float> 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<int> hist_a((size_t) NT2 * 4, -1); | |
| for (int t = 0; t < NT2; ++t) hist_a[(size_t) t * 4 + 0] = 0; | |
| std::vector<int> want_a((size_t) NT2), alt_a((size_t) NT2); | |
| for (int t = 0; t < NT2; ++t) { | |
| const std::vector<float> row(la.begin() + (size_t) t * NV2, la.begin() + (size_t) (t + 1) * NV2); | |
| const std::vector<int> 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<float> 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<int> 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<int> want_b((size_t) NT2); | |
| for (int t = 0; t < NT2; ++t) { | |
| const std::vector<float> row(lb.begin() + (size_t) t * NV2, lb.begin() + (size_t) (t + 1) * NV2); | |
| const std::vector<int> 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<float> l((size_t) NV3 * NT3, -5.0f); | |
| std::vector<int> 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<float> l((size_t) NV2 * NT2, -8.0f); | |
| std::vector<int> 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<int> 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<float> row(l.begin() + (size_t) t * NV2, l.begin() + (size_t) (t + 1) * NV2); | |
| const std::vector<int> 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<float> 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<float> 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<int> want((size_t) NT3); | |
| for (int t = 0; t < NT3; ++t) { | |
| const std::vector<float> 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<float> 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<int> 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<float> l((size_t) 2 * NT9, 0.0f); | |
| std::vector<int> hist((size_t) NT9, 0); | |
| std::vector<int> 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<float> row = {lp[0], lp[1], lp[2], lp[3], -30.0f, -30.0f, -30.0f, -30.0f}; | |
| std::vector<float> l; | |
| for (int t = 0; t < NT10; ++t) l.insert(l.end(), row.begin(), row.end()); | |
| std::vector<int> 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<float> g(0.0f, 1.0f); | |
| std::vector<float> l((size_t) NV * NT); | |
| for (auto& v : l) v = g(rng); | |
| std::vector<int> 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<int> 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<int> 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<float>& l, const std::vector<int>& 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<float> g(0.0f, 1.0f); | |
| const int TAIL = 5000; // longer than the widest window below | |
| std::vector<int32_t> 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<int32_t> 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<int32_t> rows((size_t) T * H); | |
| strata::kernels::penalty_rows(tail.data(), TAIL, window.data(), T, H, rows.data()); | |
| std::vector<float> 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<int> gwant((size_t) T), swant((size_t) T), rows_int(rows.begin(), rows.end()); | |
| for (int t = 0; t < T; ++t) { | |
| const std::vector<float> lr(l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV); | |
| const std::vector<int> own(rows.begin() + (size_t) t * H, rows.begin() + (size_t) (t + 1) * H); | |
| const std::vector<int> 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<float> g(0.0f, 1.0f); | |
| const int H = 16; | |
| std::vector<float> l((size_t) NV * NT); | |
| for (auto& v : l) v = g(rng); | |
| std::vector<int> 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<int> gwant((size_t) NT), swant((size_t) NT); | |
| for (int t = 0; t < NT; ++t) { | |
| const std::vector<float> lr(l.begin() + (size_t) t * NV, l.begin() + (size_t) (t + 1) * NV); | |
| const std::vector<int> 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<float> g(0.0f, 1.0f); | |
| std::vector<float> 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<int> 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<int32_t> 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<int32_t> 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<int32_t> seq(tail); | |
| seq.insert(seq.end(), window.begin(), window.begin() + t + 1); | |
| std::vector<int32_t> expect((size_t) h, -1); | |
| const int take = (int) std::min<size_t>((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<int32_t> old((size_t) h, -1); | |
| const int take0 = (int) std::min<int64_t>(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<float> 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<int> 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<float>::infinity(); | |
| const float qnan = std::numeric_limits<float>::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<float> g(0.0f, 1.0f); | |
| std::vector<float> 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<int> 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<SelList> 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<int> 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<float> g(0.0f, 1.5f); | |
| std::vector<float> 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<int> 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<int> 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<SelList> 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<int> 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<int>& got, const std::vector<SelList>& 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<float>& l, std::vector<int>& hist) { | |
| std::mt19937 rng(seed); | |
| std::normal_distribution<float> 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<float> l; | |
| std::vector<int> 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<SelList> 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<int> 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<float> l; | |
| std::vector<int> 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<SelList> 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; | |
| } | |