bankml / docs /modules /sampler.md
Gregory-L's picture
bankML: the whole source (github.com/cryptoAGI/bankml @ 12ae409) and its page, with the bankML persona; the live engine (Dockerfile, hf/start.sh) ready for Docker hardware
28c70af verified
|
Raw History Blame Contribute Delete
14.7 kB
# `bankML/sampler.rs` — llama-server's sampler chain, token for token with the same seed
## Summary
`sampler.rs` picks the next token from the logits as llama.cpp b11192's sampler chain does when llama-server builds
it for a chat request (`common/sampling.cpp`). The full chain is penalties → dry → top-n-σ → top-k → typical-p →
top-p → min-p → xtc → temperature → dist. bankml reproduces **all of it** (0.3.7, unreleased), each step from
`src/llama-sampler.cpp` in its float order, with llama.cpp's `sorted` state tracked through the chain. With the same
seed it draws the same tokens as llama-server. It is step eleven of P3 in [../oracles.md](../oracles.md).
With the Bonsai GGUF's own defaults (top-k 20, top-p 0.85, min-p 0, temperature 0.5) and the penalties, DRY,
top-n-σ, typical-p and XTC neutral, the active part of the chain is top-k, top-p, min-p, temperature and the draw;
these were reproduced first (P3).
The penalties came first, in 0.3.6 (O2): the repeat, frequency and presence penalties over the last `repeat_last_n`
tokens. They let the coach's `ollama_predict` (`repeat_penalty 1.3`) and mindXtrain's imprint gate run on bankml;
without a penalty, greedy `mindx-gen39` degenerates into `,,,,` (CHANGELOG 0.3.6).
It also computes each token's probability as llama-server reports it in `logprobs` (`token_probs`, 0.3.8, unreleased).
Callers: the native engine (`native.rs`, `Native::complete_with`), and through it `bankml serve --native` (`/v1`,
`/api/*`) and the C API; `bankml generate --sample` (`main.rs`); the Ollama layer, which checks `Params` from Ollama's
options before a model is loaded (`ollama.rs`); `serve.rs`, which does the same for `/v1`. Under a grammar,
`grammar.rs` supplies the check and the mask to `sample_constrained`.
## Technical usage
```rust
pub struct Params { // impl Default: llama.cpp's common.h
pub temp: f32, pub top_k: i32, pub top_p: f32, pub min_p: f32, pub min_keep: usize, pub seed: u32,
pub penalty_last_n: i32, pub penalty_repeat: f32, pub penalty_freq: f32, pub penalty_present: f32, // 0.3.6
pub typical_p: f32, pub top_n_sigma: f32, pub xtc_probability: f32, pub xtc_threshold: f32, // 0.3.7
pub dynatemp_range: f32, pub dynatemp_exponent: f32,
pub dry_multiplier: f32, pub dry_base: f32, pub dry_allowed_length: i32, pub dry_penalty_last_n: i32,
pub dry_sequence_breakers: Vec<String>,
pub n_probs: usize, // 0.3.8
}
impl Params { pub fn from_gguf(path: &std::path::Path) -> Result<Self, String> }
pub const DEFAULT_SEED: u32 = 0xFFFF_FFFF;
impl Sampler {
pub fn new(p: Params) -> Result<Self, String>
pub fn wants_dry_breakers(&self) -> bool
pub fn dry_sequence_breakers(&self) -> &[String]
pub fn set_dry_breakers(&mut self, b: std::sync::Arc<DryBreakers>)
pub fn accept(&mut self, t: u32)
pub fn sample(&mut self, logits: &[f32]) -> u32
pub fn sample_constrained(&mut self, logits: &[f32], allows: impl Fn(u32) -> bool, mask: impl Fn(&mut [Cand])) -> (u32, bool)
}
pub type DryBreakers = HashMap<u32, Vec<Vec<u32>>>;
pub fn dry_breakers(breakers: &[String], pieces: &[Vec<u8>], encode: impl Fn(&str) -> Vec<u32>) -> Result<DryBreakers, String>
pub struct TokenProb { pub id: u32, pub p: f32 }
pub fn token_probs(logits: &[f32], sampled: u32, n_top: usize) -> (f32, Vec<TokenProb>)
pub struct Mt19937 { … } // new(seed), next_u32(), uniform(), uniform_f32()
pub struct Cand { pub id: u32, pub logit: f32, pub p: f32 } // llama_token_data
pub fn partial_sort(v: &mut [Cand], middle: usize) // libstdc++'s std::partial_sort
pub fn sort_by<T: Copy>(v: &mut [T], less: &impl Fn(&T, &T) -> bool) // libstdc++'s std::sort (introsort)
```
- `Params::from_gguf` resolves as llama-server does without a request: the GGUF's `general.sampling.*` (`temp`,
`top_k`, `top_p`, `min_p`, `penalty_last_n`, `penalty_repeat`, `xtc_probability`, `xtc_threshold`) over llama.cpp's
defaults (temperature 0.8, top-k 40, top-p 0.95, min-p 0.05, `penalty_last_n` 64, penalties neutral, typical-p 1,
top-n-σ −1, XTC probability 0 and threshold 0.1, dynamic temperature off, DRY off with base 1.75, allowed length 2,
window 64 and breakers `\n`, `:`, `"`, `*`), and `DEFAULT_SEED`, which means a seed from the clock and process id,
as llama.cpp seeds from the random device.
- `Sampler::new` refuses what is not reproduced (top-k outside 1–128) and what llama-server refuses, with its
message word for word: a negative `repeat_last_n`, `dry_allowed_length` or `dry_penalty_last_n`; a repeat penalty
that is not finite and above 0; a non-finite frequency or presence penalty; an empty `dry_sequence_breakers`.
- DRY needs its breakers as token sequences of the model's vocabulary. When `wants_dry_breakers()`, the caller builds
them with `dry_breakers` (each breaker cut to 40 bytes, each tail to 20 tokens, as `llama_sampler_init_dry`) and
hands them over with `set_dry_breakers`. `native.rs` keeps the last list's result, so a conversation builds it once.
- `accept(t)` puts a token into the penalties' and DRY's windows. llama-server accepts **every prompt token** before
the first draw, then each token drawn; callers do the same.
- `token_probs(logits, sampled, n_top)` is llama-server's `get_token_probabilities` on the raw logits, before any
sampler: a partial sort of the whole vocabulary for the top `n_top`, then `expf(logit − max)` summed in float in the
order the partial sort left the array, each divided by the sum. It returns the sampled token's probability and the
top `n_top`. The native engine calls it per token when `n_probs` > 0; `serve.rs` writes the `logprob` as `logf` of
it.
- `sample(logits)` runs the chain once. The RNG advances once per token.
- `sample_constrained` is `common_sampler_sample` with `grammar_first = false`: the chain on the raw logits; if
`allows` the token, it stands. Otherwise the logits are taken afresh, `mask` sets rejected tokens to −∞, and the
chain runs again: a second draw from the same generator. The `bool` says whether it redrew.
The chain, step by step:
| step | what it does |
|---|---|
| penalties | for a token seen `count` times in the window: logit divided by the repeat penalty if positive, multiplied if not; then `count · freq + present` taken off. Skipped when `last_n` is 0 or all three are neutral. Under a grammar's redraw the chain runs again, so the penalties apply again |
| DRY (0.3.7) | over the last `dry_penalty_last_n` tokens: the nearest restart sequence (a breaker) caps the repeat length; a reverse Z-algorithm finds each suffix's repeat; a token that would extend a repeat of at least `dry_allowed_length` loses `multiplier · base^(len − allowed)` (libm `pow`, the exponent clamped); single-token breakers are never penalised. Breakers are built from the vocabulary's pieces once per list and cached by the engine |
| top-n-σ (0.3.7) | over the finite logits: those below `max − n·σ` become −∞ (the squares in double, as C++'s `pow(float, 2)`) |
| top-k | libstdc++'s `std::partial_sort` (heap select, then heap sort), ported exactly, because tied logits are common on a 1-bit model and the order it leaves them in decides the draw |
| typical-p (0.3.7) | the softmax's entropy, each token's distance from it, the tokens sorted by that distance with libstdc++'s `std::sort` (ported: introsort leaves ties in its own order), kept while the running probability ≤ p; leaves the set unsorted |
| top-p | a float softmax over the kept tokens (in their current order), sorted afterwards if typical-p left them unsorted, a float running sum cut where it reaches p (respecting `min_keep`) |
| min-p | a cut at `max + logf(p)`: on an unsorted set the filter keeps order (and falls back to the sorted cut if fewer than `min_keep` pass) |
| XTC (0.3.7) | its own `mt19937` draws a float (`generate_canonical<float, 24>`); when it falls within the probability, the most probable tokens above the threshold are dropped, all but the last of them |
| temperature | `logit / temp`; at temp ≤ 0 every logit but the first maximum becomes −∞. Dynamic (0.3.7) when `dynatemp_range` > 0: the temperature follows the softmax's normalised entropy, `min + (max − min)·entropy^exponent` |
| dist | `expf(logit − max)` summed in double; one `uniform_real_distribution<double>` draw from `std::mt19937` (two 32-bit outputs, libstdc++'s `generate_canonical`) |
Request fields that reach `Params` (native engine, `native::sampling`): `temperature`, `top_k`, `top_p`, `min_p`,
`min_keep`, `seed`, `repeat_last_n`, `repeat_penalty`, `frequency_penalty`, `presence_penalty`, and since 0.3.7
`typical_p`, `top_n_sigma`, `xtc_probability`, `xtc_threshold`, `dynatemp_range`, `dynatemp_exponent`,
`dry_multiplier`, `dry_base`, `dry_allowed_length`, `dry_penalty_last_n`, `dry_sequence_breakers`; `n_probs` comes from
`/v1`'s `logprobs` and `top_logprobs` (0.3.8). llama-server's soft limits clamp (`top_p`, `min_p` and the XTC fields to
[0, 1], temperature to ≥ 0; a `dry_base` below 1 becomes 1.75). On `/api/*` the fields Ollama has come as `options`
(through `typical_p`; see [ollama.md](ollama.md)), and from a Modelfile's `PARAMETER`. `bankml generate --sample`
takes `--temp`, `--top-k`, `--top-p`, `--min-p` and `--seed` over the GGUF's defaults.
```rust
let mut s = Sampler::new(Params::from_gguf(model)?)?;
for &t in &prompt { s.accept(t); } // the prompt fills the penalties' window
let next = s.sample(&w.logits(&rn)?);
s.accept(next);
```
## How it is verified
- Unit tests: `mt19937_reference_value` (the C++ standard's check: the 10,000th output of a 5489-seeded mt19937 is
4,123,659,995) and `partial_sort_orders_the_top`.
- `oracle_sample_llama_server` (`forward.rs`, `#[ignore]`, in the gate): `testing/sample_oracle.py` has llama-server
sample 40 continuations with fixed seeds over temperature 0–1.5, top-k 5–128, top-p and min-p, and keeps the
parameters the server reports. Replayed through bankml's forward pass and sampler: **40 of 40, 1,175 tokens**.
`oracle_llama_server_bonsai_1_7b` and `oracle_llama_server_llama_f16` do the same on the O4 models (40 of 40 each).
- `oracle_samplers` (`native.rs`, 0.3.7): `testing/penalty_oracle.py --kind sampler`, 23 variants × 4 prompts:
mindx-gen39 **76 / 76** (3,576 tokens) and Bonsai-1.7B **76 / 76** (2,587 tokens), **16 / 16** refusals each with
llama-server's message. `oracle_samplers_8b` replays a Bonsai-8B record in the gate; CHANGELOG 0.3.7 does
not yet give its count.
`oracle_std_sort`: libstdc++'s own `std::sort` (`testing/sort_oracle.cpp`) on 876 key arrays, sizes 0 to 1,000,
heavy ties, sorted, reversed and equal keys: **876 / 876** orders identical.
- `sampler_oracle_live` (gate): the sampler cases through a running `serve --native` on mindx-gen39.
- `logprobs_oracle_live` (gate, 0.3.8): `testing/logprobs_oracle.py`, **14 / 14** with five streamed; every logprob
the same 32-bit float as llama-server's, which checks `token_probs` and its partial sort (CHANGELOG 0.3.8).
- `oracle_penalties`, `oracle_penalties_8b` (`native.rs`): `testing/penalty_oracle.py`, llama-server b11192 from an
empty cache, greedy and seeded, 17 variants × 4 prompts made to repeat, `repeat_last_n` smaller than, equal to and
larger than the prompt. mindx-gen39 **56 / 56** (2,478 tokens), Bonsai-1.7B **56 / 56** (1,895 tokens), Bonsai-8B
**56 / 56** (1,568 tokens); **12 / 12** refusals each with llama-server's message (CHANGELOG 0.3.6).
- `penalty_oracle_live` (gate): the recorded cases through a running `serve --native`, via `/v1` and `/api/chat`:
**85 / 85**, refusals as 400s.
- Under a grammar: `oracle_json_mode`, `oracle_json_schema*` check the redraw's use of the generator (152 of 860
tokens redrawn on Bonsai-8B in JSON mode).
## Advantages and efficiency
- **Reproducible sampling.** A seed gives the same answer as llama-server, so seeded answers can be checked by an
oracle like greedy ones, and a grammar's redraw consumes the generator exactly as llama-server's does.
- **Small after top-k.** Top-k runs a heap select over the vocabulary in place and truncates to at most 128
candidates; typical-p, top-p, min-p, XTC, temperature and the draw then work on that short list.
- **Penalties cost nothing when off.** They are skipped entirely when disabled, as llama.cpp disables them. The window
is a ring (`VecDeque`) with a per-token count (`HashMap`), updated in constant time per accepted token.
- **The redraw only when needed.** Under a grammar the whole-vocabulary mask runs only when the first draw breaks
the grammar (see [grammar.md](grammar.md) for its cost).
- **Rust practice.** No crates: mt19937 and libstdc++'s heap algorithms are written out. No `unsafe`. Refusals are
`Err`s carrying llama-server's own text.
- **DRY's breakers cost once.** Building them scans every token's piece (151,669 on the Qwen3 models); the engine
keeps the result per breaker list, so a conversation pays it once.
- **Logprobs only when asked.** `token_probs` runs only when `n_probs` > 0; it is one partial sort of the
vocabulary per generated token.
- **Next:** nothing in llama-server's default chain is left. Open (docs/OLLAMA.md, O2): mirostat and a custom
sampler order, refused today.
## Limitations
- Not reproduced, refused: mirostat, and a custom `samplers` order (`native::sampling`). DRY, XTC, top-n-σ and
dynamic temperature are not Ollama options, so they come through `/v1` and the C API only; on `/api/*` they are
refused as unknown options.
- `n_probs` is reached only through `/v1`'s `logprobs`; llama-server's `/completion` endpoint (and its `n_probs`
field) is not served natively.
- A DRY breaker whose split point falls inside a multi-byte character is refused (bankML tokenizes whole characters).
- Top-k 0 or above 128 is refused (llama.cpp sorts larger sets another way).
- mindXtrain's `no_repeat_ngram_size` is a transformers rule, not a llama.cpp sampler; it is not here (docs/OLLAMA.md).
- `DEFAULT_SEED` seeds from the clock and process id, so such a run is not reproducible, as in llama.cpp.
- `bankml generate --sample` has no penalty flags; penalties there come from the GGUF's defaults.
## See also
- [../oracles.md](../oracles.md) §1d (step eleven), §5c
- [../OLLAMA.md](../OLLAMA.md) — O2, the sampler chain
- [../TODO.md](../TODO.md) — 0.4.0
- [../usage.md](../usage.md) §13 — `bankml generate --sample`
- Sibling pages: [grammar.md](grammar.md), [forward.md](forward.md), [native.md](native.md),
[serve.md](serve.md), [ollama.md](ollama.md)