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.

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

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), and from a Modelfile's PARAMETER. bankml generate --sample takes --temp, --top-k, --top-p, --min-p and --seed over the GGUF's defaults.

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 for its cost).
  • Rust practice. No crates: mt19937 and libstdc++'s heap algorithms are written out. No unsafe. Refusals are Errs 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