use std::borrow::Cow;
use std::collections::HashMap;
use crate::prefix_cache::{CachedPrefix, PrefixCache};
use crate::runtime::serve::Runtime;
use crate::runtime::transformer::{KvCache, forward_gpt2_cached, forward_llama_cached};
use crate::tokenizer::BpeTokenizer;
fn make_kv_cache(n_layers: usize, hidden: usize, config: &GenerationConfig) -> KvCache {
if config.use_mixed_kv {
KvCache::new_mixed(n_layers, hidden, 64)
} else {
KvCache::new_quantized(n_layers, hidden, config.use_int8_kv)
}
}
fn context_shift(
token_ids: &mut Vec<u32>,
prompt_len: &mut usize,
kv_cache_opt: &mut Option<KvCache>,
max_ctx: usize,
anchor_tokens: usize,
) {
let len = token_ids.len();
let overflow = len - max_ctx + max_ctx / 4;
if anchor_tokens > 0 && overflow > 0 && len > anchor_tokens + overflow {
if let Some(kv) = kv_cache_opt {
kv.shift_anchored(overflow, anchor_tokens);
}
let suffix_start = anchor_tokens + overflow;
let suffix_len = len - suffix_start;
token_ids.copy_within(suffix_start.., anchor_tokens);
token_ids.truncate(anchor_tokens + suffix_len);
if *prompt_len > anchor_tokens {
let non_anchor = *prompt_len - anchor_tokens;
let removed = overflow.min(non_anchor);
*prompt_len = anchor_tokens + non_anchor - removed;
}
} else {
if let Some(kv) = kv_cache_opt {
kv.shift(overflow);
}
let new_len = len - overflow;
token_ids.copy_within(overflow.., 0);
token_ids.truncate(new_len);
*prompt_len = prompt_len.saturating_sub(overflow);
}
}
pub struct GenerationConfig {
pub max_tokens: usize,
pub temperature: f32,
pub top_p: f32,
pub min_p: f32,
pub gamma: usize,
pub use_int8_kv: bool,
pub use_mixed_kv: bool,
pub constraint: Option<std::sync::Arc<dyn crate::constraint::Constraint>>,
pub max_context: Option<usize>,
pub anchor_tokens: usize,
pub stop: Vec<String>,
pub seed: Option<u64>,
pub repetition_penalty: f32,
pub presence_penalty: f32,
pub frequency_penalty: f32,
pub logit_bias: HashMap<u32, f32>,
pub cancel: Option<std::sync::Arc<std::sync::atomic::AtomicBool>>,
}
impl GenerationConfig {
pub fn cancelled(&self) -> bool {
self.cancel
.as_ref()
.map(|c| c.load(std::sync::atomic::Ordering::Relaxed))
.unwrap_or(false)
}
}
impl Clone for GenerationConfig {
fn clone(&self) -> Self {
Self {
max_tokens: self.max_tokens,
temperature: self.temperature,
top_p: self.top_p,
min_p: self.min_p,
gamma: self.gamma,
use_int8_kv: self.use_int8_kv,
use_mixed_kv: self.use_mixed_kv,
constraint: self.constraint.clone(),
max_context: self.max_context,
anchor_tokens: self.anchor_tokens,
stop: self.stop.clone(),
seed: self.seed,
repetition_penalty: self.repetition_penalty,
presence_penalty: self.presence_penalty,
frequency_penalty: self.frequency_penalty,
logit_bias: self.logit_bias.clone(),
cancel: self.cancel.clone(),
}
}
}
impl Default for GenerationConfig {
fn default() -> Self {
Self {
max_tokens: 128,
temperature: 0.0, top_p: 0.0,
min_p: 0.0,
gamma: 0,
use_int8_kv: false,
use_mixed_kv: false,
constraint: None,
max_context: None,
anchor_tokens: 0,
stop: Vec::new(),
seed: None,
repetition_penalty: 1.0,
presence_penalty: 0.0,
frequency_penalty: 0.0,
logit_bias: HashMap::new(),
cancel: None,
}
}
}
pub fn generate(
runtime: &Runtime,
architecture: &str,
hidden: usize,
tokenizer: &BpeTokenizer,
prompt: &str,
config: &GenerationConfig,
draft_model: Option<&dyn crate::draft::DraftModel>,
) -> String {
let generated = generate_token_ids(
runtime,
architecture,
hidden,
tokenizer,
prompt,
config,
draft_model,
);
tokenizer.decode(&generated)
}
pub fn generate_token_ids(
runtime: &Runtime,
architecture: &str,
hidden: usize,
tokenizer: &BpeTokenizer,
prompt: &str,
config: &GenerationConfig,
draft_model: Option<&dyn crate::draft::DraftModel>,
) -> Vec<u32> {
generate_token_ids_with_cache(
runtime,
architecture,
hidden,
tokenizer,
prompt,
config,
None,
draft_model,
)
}
#[allow(clippy::too_many_arguments)]
pub fn generate_token_ids_with_cache(
runtime: &Runtime,
architecture: &str,
hidden: usize,
tokenizer: &BpeTokenizer,
prompt: &str,
config: &GenerationConfig,
mut cache: Option<&mut PrefixCache>,
draft_model: Option<&dyn crate::draft::DraftModel>,
) -> Vec<u32> {
let prompt_ids = tokenizer.encode(prompt);
let lookup = cache.as_ref().and_then(|c| c.lookup(&prompt_ids).ok());
if config.gamma > 0 {
let (text, maybe_kv) = speculative_generate(
runtime,
architecture,
hidden,
tokenizer,
prompt,
config,
lookup,
draft_model,
);
let all_ids = tokenizer.encode(&(prompt.to_string() + &text));
let ids = if all_ids.len() > prompt_ids.len() {
all_ids[prompt_ids.len()..].to_vec()
} else {
Vec::new()
};
if let Some(ref mut c) = cache
&& let Some(kv) = maybe_kv
{
let _ = c.insert(
all_ids,
CachedPrefix {
kv,
last_logits: Vec::new(),
},
);
}
return ids;
}
let (ids, _logprobs, maybe_kv) = generate_core(
runtime,
architecture,
hidden,
tokenizer,
prompt,
config,
lookup,
false,
0,
);
if let Some(ref mut c) = cache
&& let Some(kv) = maybe_kv
{
let mut full_ids = prompt_ids.clone();
full_ids.extend_from_slice(&ids);
let _ = c.insert(
full_ids,
CachedPrefix {
kv,
last_logits: Vec::new(),
},
);
}
ids
}
#[derive(Debug, Clone, PartialEq)]
pub struct TokenLogprob {
pub token: u32,
pub logprob: f32,
pub top_logprobs: Vec<(u32, f32)>,
}
pub fn generate_with_logprobs(
runtime: &Runtime,
architecture: &str,
hidden: usize,
tokenizer: &BpeTokenizer,
prompt: &str,
config: &GenerationConfig,
top_logprobs: usize,
) -> (Vec<u32>, Vec<TokenLogprob>) {
generate_with_logprobs_with_cache(
runtime,
architecture,
hidden,
tokenizer,
prompt,
config,
top_logprobs,
None,
)
}
#[allow(clippy::too_many_arguments)]
pub fn generate_with_logprobs_with_cache(
runtime: &Runtime,
architecture: &str,
hidden: usize,
tokenizer: &BpeTokenizer,
prompt: &str,
config: &GenerationConfig,
top_logprobs: usize,
mut cache: Option<&mut PrefixCache>,
) -> (Vec<u32>, Vec<TokenLogprob>) {
let prompt_ids = tokenizer.encode(prompt);
let lookup = cache.as_ref().and_then(|c| c.lookup(&prompt_ids).ok());
let (ids, logprobs, maybe_kv) = generate_core(
runtime,
architecture,
hidden,
tokenizer,
prompt,
config,
lookup,
true,
top_logprobs,
);
if let Some(ref mut c) = cache
&& let Some(kv) = maybe_kv
{
let mut full_ids = prompt_ids.clone();
full_ids.extend_from_slice(&ids);
let _ = c.insert(
full_ids,
CachedPrefix {
kv,
last_logits: Vec::new(),
},
);
}
(ids, logprobs)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn generate_core(
runtime: &Runtime,
architecture: &str,
hidden: usize,
tokenizer: &BpeTokenizer,
prompt: &str,
config: &GenerationConfig,
cache_lookup: Option<crate::prefix_cache::CacheLookup>,
collect_logprobs: bool,
top_logprobs: usize,
) -> (Vec<u32>, Vec<TokenLogprob>, Option<KvCache>) {
let mut token_ids = tokenizer.encode(prompt);
if token_ids.is_empty() {
return (Vec::new(), Vec::new(), None);
}
let mut prompt_len = token_ids.len();
let n_layers = count_layers(runtime, architecture);
let mut matched_len = 0usize;
let mut kv_cache_opt: Option<KvCache>;
let mut logits: Option<Vec<f32>> = None;
if let Some(look) = cache_lookup {
kv_cache_opt = Some(look.kv);
matched_len = look.matched_len;
logits = look.last_logits; } else {
kv_cache_opt = Some(make_kv_cache(n_layers, hidden, config));
}
if logits.is_none() {
for &id in &token_ids[matched_len..] {
let embedding = token_embedding(runtime, architecture, id, hidden);
logits = match architecture {
"gpt2" => {
forward_gpt2_cached(runtime, embedding.as_ref(), false, &mut kv_cache_opt)
}
"llama" => {
forward_llama_cached(runtime, embedding.as_ref(), false, &mut kv_cache_opt)
}
_ => break,
};
}
}
if logits.is_none() {
return (Vec::new(), Vec::new(), kv_cache_opt.clone());
}
let mut generated = 0usize;
let mut logprobs: Vec<TokenLogprob> = Vec::new();
loop {
if generated >= config.max_tokens || config.cancelled() {
break;
}
let next_token = sample_constrained(
logits.as_deref(),
config,
tokenizer,
&token_ids[prompt_len..],
generated,
);
if collect_logprobs {
logprobs.push(match logits.as_deref() {
Some(l) => logprob_for_step(l, next_token, top_logprobs),
None => TokenLogprob {
token: next_token,
logprob: f32::NEG_INFINITY,
top_logprobs: Vec::new(),
},
});
}
token_ids.push(next_token);
generated += 1;
if !config.stop.is_empty() {
let text = tokenizer.decode(&token_ids[prompt_len..]);
if let Some(pos) = config.stop.iter().filter_map(|s| text.find(s)).min() {
let prefix = &text[..pos];
let prefix_ids = tokenizer.encode(prefix);
let target_len = prompt_len + prefix_ids.len();
token_ids.truncate(target_len);
break;
}
}
if let Some(max_ctx) = config.max_context
&& token_ids.len() > max_ctx
{
context_shift(
&mut token_ids,
&mut prompt_len,
&mut kv_cache_opt,
max_ctx,
config.anchor_tokens,
);
}
if tokenizer.vocab_size() > 0 && next_token as usize >= tokenizer.vocab_size() {
break;
}
let embedding = token_embedding(runtime, architecture, next_token, hidden);
let next_logits = match architecture {
"gpt2" => forward_gpt2_cached(runtime, embedding.as_ref(), false, &mut kv_cache_opt),
"llama" => forward_llama_cached(runtime, embedding.as_ref(), false, &mut kv_cache_opt),
_ => break,
};
match next_logits {
Some(l) => logits = Some(l),
None => break,
}
}
let cut = prompt_len.min(token_ids.len());
let final_kv = kv_cache_opt.clone();
(token_ids[cut..].to_vec(), logprobs, final_kv)
}
fn check_stop_sequences(
tokenizer: &crate::tokenizer::BpeTokenizer,
stop: &[String],
token_ids: &mut Vec<u32>,
prompt_len: usize,
) -> Option<Vec<u32>> {
if stop.is_empty() {
return None;
}
let text = tokenizer.decode(&token_ids[prompt_len..]);
let pos = stop.iter().filter_map(|s| text.find(s)).min()?;
let prefix = &text[..pos];
let prefix_ids = tokenizer.encode(prefix);
token_ids.truncate(prompt_len + prefix_ids.len());
Some(token_ids.clone())
}
fn logprob_for_step(logits: &[f32], chosen: u32, top_n: usize) -> TokenLogprob {
if logits.is_empty() {
return TokenLogprob {
token: chosen,
logprob: f32::NEG_INFINITY,
top_logprobs: Vec::new(),
};
}
let probs = softmax(logits);
let logprob = probs
.get(chosen as usize)
.map(|p| p.ln())
.unwrap_or(f32::NEG_INFINITY);
let top_logprobs = if top_n == 0 {
Vec::new()
} else {
let mut indexed: Vec<(usize, f32)> = probs.iter().copied().enumerate().collect();
indexed.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
indexed
.into_iter()
.take(top_n)
.map(|(i, p)| (i as u32, p.ln()))
.collect()
};
TokenLogprob {
token: chosen,
logprob,
top_logprobs,
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn speculative_generate(
runtime: &Runtime,
architecture: &str,
hidden: usize,
tokenizer: &BpeTokenizer,
prompt: &str,
config: &GenerationConfig,
cache_lookup: Option<crate::prefix_cache::CacheLookup>,
draft_model: Option<&dyn crate::draft::DraftModel>,
) -> (String, Option<KvCache>) {
let mut token_ids = tokenizer.encode(prompt);
if token_ids.is_empty() {
return (String::new(), None);
}
let n_layers = count_layers(runtime, architecture);
let mut prompt_len = token_ids.len();
let n = 3usize; let mut ngram_index = build_ngram_index(&token_ids, n);
let mut matched_len = 0usize;
let mut kv_cache_opt: Option<KvCache>;
if let Some(look) = cache_lookup {
kv_cache_opt = Some(look.kv);
matched_len = look.matched_len;
} else {
kv_cache_opt = Some(make_kv_cache(n_layers, hidden, config));
}
for &id in &token_ids[matched_len..] {
let embedding = token_embedding(runtime, architecture, id, hidden);
let _ = match architecture {
"gpt2" => forward_gpt2_cached(runtime, &embedding, false, &mut kv_cache_opt),
"llama" => forward_llama_cached(runtime, &embedding, false, &mut kv_cache_opt),
_ => break,
};
}
while token_ids.len() - prompt_len < config.max_tokens {
if config.cancelled() {
break;
}
let draft = if let Some(dm) = draft_model {
dm.draft(&token_ids, config.gamma)
} else {
draft_tokens(&ngram_index, &token_ids, n, config.gamma)
};
if draft.is_empty() {
let last_id = token_ids.last().copied().unwrap_or(0);
let embedding = token_embedding(runtime, architecture, last_id, hidden);
let logits = match architecture {
"gpt2" => {
forward_gpt2_cached(runtime, embedding.as_ref(), false, &mut kv_cache_opt)
}
"llama" => {
forward_llama_cached(runtime, embedding.as_ref(), false, &mut kv_cache_opt)
}
_ => break,
};
let step = token_ids.len() - prompt_len;
let next_token = sample_constrained(
logits.as_deref(),
config,
tokenizer,
&token_ids[prompt_len..],
step,
);
token_ids.push(next_token);
ngram_index = update_ngram_index(&ngram_index, &token_ids, n);
if let Some(max_ctx) = config.max_context
&& token_ids.len() > max_ctx
{
context_shift(
&mut token_ids,
&mut prompt_len,
&mut kv_cache_opt,
max_ctx,
config.anchor_tokens,
);
}
if tokenizer.vocab_size() > 0 && next_token as usize >= tokenizer.vocab_size() {
break;
}
if let Some(truncated) =
check_stop_sequences(tokenizer, &config.stop, &mut token_ids, prompt_len)
{
token_ids = truncated;
break;
}
continue;
}
let mut accepted = 0usize;
for &draft_token in &draft {
let last_id = token_ids.last().copied().unwrap_or(0);
let embedding = token_embedding(runtime, architecture, last_id, hidden);
let logits = match architecture {
"gpt2" => {
forward_gpt2_cached(runtime, embedding.as_ref(), false, &mut kv_cache_opt)
}
"llama" => {
forward_llama_cached(runtime, embedding.as_ref(), false, &mut kv_cache_opt)
}
_ => break,
};
let step = token_ids.len() - prompt_len;
let next_token = sample_constrained(
logits.as_deref(),
config,
tokenizer,
&token_ids[prompt_len..],
step,
);
if next_token == draft_token {
token_ids.push(draft_token);
accepted += 1;
if tokenizer.vocab_size() > 0 && draft_token as usize >= tokenizer.vocab_size() {
break;
}
} else {
token_ids.push(next_token);
break;
}
}
ngram_index = update_ngram_index(&ngram_index, &token_ids, n);
if let Some(truncated) =
check_stop_sequences(tokenizer, &config.stop, &mut token_ids, prompt_len)
{
token_ids = truncated;
break;
}
if let Some(max_ctx) = config.max_context
&& token_ids.len() > max_ctx
{
context_shift(
&mut token_ids,
&mut prompt_len,
&mut kv_cache_opt,
max_ctx,
config.anchor_tokens,
);
}
if token_ids.len() - prompt_len >= config.max_tokens {
break;
}
if accepted < draft.len()
&& tokenizer.vocab_size() > 0
&& token_ids.last().copied().unwrap_or(0) as usize >= tokenizer.vocab_size()
{
break;
}
}
let final_kv = kv_cache_opt.clone();
(
tokenizer.decode(&token_ids[prompt_len.min(token_ids.len())..]),
final_kv,
)
}
fn build_ngram_index(tokens: &[u32], n: usize) -> HashMap<Vec<u32>, Vec<u32>> {
let mut index = HashMap::new();
if tokens.len() < n + 1 {
return index;
}
for window in tokens.windows(n + 1) {
let key = window[..n].to_vec();
let next = window[n];
index.entry(key).or_default().push(next);
}
index
}
fn update_ngram_index(
index: &HashMap<Vec<u32>, Vec<u32>>,
tokens: &[u32],
n: usize,
) -> HashMap<Vec<u32>, Vec<u32>> {
let mut new_index = index.clone();
if tokens.len() < n + 1 {
return new_index;
}
let start = tokens.len().saturating_sub(n + 1);
let window = &tokens[start..];
let key = window[..n].to_vec();
let next = window[n];
new_index.entry(key).or_default().push(next);
new_index
}
fn draft_tokens(
index: &HashMap<Vec<u32>, Vec<u32>>,
tokens: &[u32],
n: usize,
gamma: usize,
) -> Vec<u32> {
let mut draft = Vec::new();
let mut context = tokens.to_vec();
for _ in 0..gamma {
let key = if context.len() >= n {
context[context.len() - n..].to_vec()
} else {
context.clone()
};
if let Some(candidates) = index.get(&key) {
let next = most_frequent(candidates);
draft.push(next);
context.push(next);
} else {
break;
}
}
draft
}
fn most_frequent(xs: &[u32]) -> u32 {
let mut counts = HashMap::new();
for &x in xs {
*counts.entry(x).or_insert(0usize) += 1;
}
counts
.into_iter()
.max_by_key(|&(_, c)| c)
.map(|(v, _)| v)
.unwrap_or(0)
}
fn token_embedding<'a>(
runtime: &'a Runtime,
arch: &str,
token_id: u32,
hidden: usize,
) -> Cow<'a, [f32]> {
let weight_name = match arch {
"gpt2" => "transformer.wte.weight",
"llama" => "model.embed_tokens.weight",
_ => return Cow::Owned(vec![0.0; hidden]),
};
if let Some(w) = runtime.get(weight_name) {
let idx = (token_id as usize) * hidden;
if idx + hidden <= w.data.len() {
return Cow::Borrowed(&w.data[idx..idx + hidden]);
}
}
Cow::Owned(vec![0.0; hidden])
}
fn count_layers(runtime: &Runtime, arch: &str) -> usize {
let prefix = match arch {
"gpt2" => "transformer.h.",
"llama" => "model.layers.",
_ => return 0,
};
let mut max = 0usize;
for name in runtime.tensor_names() {
if let Some(rest) = name.strip_prefix(prefix)
&& let Some(n) = rest.split('.').next().and_then(|s| s.parse::<usize>().ok())
{
max = max.max(n);
}
}
max + 1
}
fn apply_repetition_penalty(logits: &mut [f32], generated_tokens: &[u32], penalty: f32) {
if penalty <= 1.0 || generated_tokens.is_empty() {
return;
}
let mut seen = [0u64; 1024]; for &id in generated_tokens {
let idx = id as usize / 64;
let bit = id as u64 % 64;
if idx < seen.len() {
seen[idx] |= 1u64 << bit;
}
}
for (i, logit) in logits.iter_mut().enumerate() {
let idx = i / 64;
let bit = i as u64 % 64;
if idx < seen.len() && (seen[idx] & (1u64 << bit)) != 0 {
if *logit > 0.0 {
*logit /= penalty;
} else {
*logit *= penalty;
}
}
}
}
fn apply_presence_frequency_penalty(
logits: &mut [f32],
generated_tokens: &[u32],
presence_penalty: f32,
frequency_penalty: f32,
) {
if (presence_penalty == 0.0 && frequency_penalty == 0.0) || generated_tokens.is_empty() {
return;
}
let mut counts = [0u16; 65536];
for &id in generated_tokens {
let idx = id as usize;
if idx < counts.len() {
counts[idx] = counts[idx].saturating_add(1);
}
}
for (i, logit) in logits.iter_mut().enumerate() {
let count = counts.get(i).copied().unwrap_or(0) as f32;
if count > 0.0 {
*logit -= presence_penalty;
*logit -= count * frequency_penalty;
}
}
}
fn sample(logits: Option<&[f32]>, config: &GenerationConfig, step: usize) -> u32 {
let logits = logits.unwrap_or(&[]);
if logits.is_empty() {
return 0;
}
if config.temperature <= 0.0 {
return argmax(logits) as u32;
}
let scaled: Vec<f32> = logits.iter().map(|l| l / config.temperature).collect();
let mut probs = softmax(&scaled);
if config.top_p > 0.0 && config.top_p < 1.0 {
apply_top_p(&mut probs, config.top_p);
}
if config.min_p > 0.0 && config.min_p < 1.0 {
apply_min_p(&mut probs, config.min_p);
}
multinomial(&probs, step, config.seed)
}
fn sample_constrained(
logits: Option<&[f32]>,
config: &GenerationConfig,
tokenizer: &BpeTokenizer,
generated_tokens: &[u32],
step: usize,
) -> u32 {
let logits = logits.unwrap_or(&[]);
if logits.is_empty() {
return 0;
}
let has_penalty = config.repetition_penalty > 1.0
|| config.presence_penalty != 0.0
|| config.frequency_penalty != 0.0;
let has_bias = !config.logit_bias.is_empty();
let constraint = config.constraint.as_deref();
if !has_penalty && !has_bias && constraint.is_none() {
return sample(Some(logits), config, step);
}
let mut working = logits.to_vec();
if has_penalty {
apply_repetition_penalty(&mut working, generated_tokens, config.repetition_penalty);
apply_presence_frequency_penalty(
&mut working,
generated_tokens,
config.presence_penalty,
config.frequency_penalty,
);
}
if has_bias {
for (&token_id, &bias) in &config.logit_bias {
let idx = token_id as usize;
if idx < working.len() {
working[idx] += bias;
}
}
}
if let Some(constraint) = constraint {
let prefix = tokenizer.decode(generated_tokens);
let vocab_size = tokenizer.vocab_size();
let mask = constraint.valid_mask(&prefix, vocab_size, tokenizer.cached_token_bytes());
for (i, valid) in mask.iter().enumerate() {
if !valid && i < working.len() {
working[i] = f32::NEG_INFINITY;
}
}
}
sample(Some(&working), config, step)
}
pub(crate) fn apply_top_p(probs: &mut [f32], top_p: f32) {
let mut indexed: Vec<(usize, f32)> = probs.iter().copied().enumerate().collect();
indexed.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let mut cum = 0.0f32;
let mut keep = 0usize;
for (_, p) in &indexed {
cum += *p;
keep += 1;
if cum >= top_p {
break;
}
}
let keep_set: std::collections::HashSet<usize> =
indexed[..keep].iter().map(|(i, _)| *i).collect();
for (i, p) in probs.iter_mut().enumerate() {
if !keep_set.contains(&i) {
*p = 0.0;
}
}
let sum: f32 = probs.iter().sum();
if sum > 0.0 {
for p in probs.iter_mut() {
*p /= sum;
}
}
}
pub(crate) fn apply_min_p(probs: &mut [f32], min_p: f32) {
let max = probs.iter().copied().fold(0.0f32, f32::max);
let threshold = max * min_p;
for p in probs.iter_mut() {
if *p < threshold {
*p = 0.0;
}
}
let sum: f32 = probs.iter().sum();
if sum > 0.0 {
for p in probs.iter_mut() {
*p /= sum;
}
}
}
pub(crate) fn argmax(xs: &[f32]) -> usize {
let mut best = 0usize;
let mut best_val = f32::NEG_INFINITY;
for (i, &x) in xs.iter().enumerate() {
if x > best_val {
best = i;
best_val = x;
}
}
best
}
pub(crate) fn softmax(xs: &[f32]) -> Vec<f32> {
let max = xs.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let exps: Vec<f32> = xs.iter().map(|x| (x - max).exp()).collect();
let sum: f32 = exps.iter().sum();
exps.into_iter().map(|e| e / sum).collect()
}
pub(crate) fn multinomial(probs: &[f32], step: usize, seed: Option<u64>) -> u32 {
let r: f32 = if let Some(s) = seed {
let mut rng = MiniRng::from_seed(s.wrapping_add(step as u64));
rng.gen_f32()
} else {
let mut rng = MiniRng::new();
rng.gen_f32()
};
let mut cum = 0.0f32;
for (i, &p) in probs.iter().enumerate() {
cum += p;
if r < cum {
return i as u32;
}
}
probs.len().saturating_sub(1) as u32
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn argmax_picks_max() {
assert_eq!(argmax(&[1.0, 3.0, 2.0]), 1);
}
#[test]
fn softmax_sums_to_one() {
let p = softmax(&[1.0, 2.0, 3.0]);
let sum: f32 = p.iter().sum();
assert!((sum - 1.0).abs() < 1e-5, "softmax sum = {sum}");
}
#[test]
fn greedy_sample_picks_argmax() {
let config = GenerationConfig {
max_tokens: 1,
temperature: 0.0,
..GenerationConfig::default()
};
assert_eq!(sample(Some(&[1.0, 5.0, 2.0]), &config, 0), 1);
}
#[test]
fn apply_top_p_keeps_nucleus() {
let mut probs = vec![0.5f32, 0.3, 0.15, 0.05];
apply_top_p(&mut probs, 0.8);
assert!((probs[2]).abs() < 1e-6, "token 2 should be zeroed");
assert!((probs[3]).abs() < 1e-6, "token 3 should be zeroed");
let sum: f32 = probs.iter().sum();
assert!((sum - 1.0).abs() < 1e-5, "should renormalize to 1.0");
}
#[test]
fn apply_top_p_with_very_high_p_keeps_all() {
let mut probs = vec![0.4f32, 0.3, 0.2, 0.1];
apply_top_p(&mut probs, 0.99);
let sum: f32 = probs.iter().sum();
assert!((sum - 1.0).abs() < 1e-5, "should still sum to 1.0");
assert!(probs.iter().all(|&p| p > 0.0), "all tokens should be kept");
}
#[test]
fn apply_min_p_drops_low_probability_tokens() {
let mut probs = vec![0.5f32, 0.2, 0.15, 0.15];
apply_min_p(&mut probs, 0.5);
assert!(
(probs[1]).abs() < 1e-6,
"token 1 below threshold should be zeroed"
);
assert!(
(probs[2]).abs() < 1e-6,
"token 2 below threshold should be zeroed"
);
assert!(
(probs[3]).abs() < 1e-6,
"token 3 below threshold should be zeroed"
);
let sum: f32 = probs.iter().sum();
assert!(
(sum - 1.0).abs() < 1e-5,
"kept tokens should renormalize to 1.0"
);
}
#[test]
fn apply_min_p_keeps_tokens_above_threshold() {
let mut probs = vec![0.4f32, 0.3, 0.2, 0.1];
apply_min_p(&mut probs, 0.5);
assert!(probs[0] > 0.0, "token 0 (0.4 >= 0.2) should be kept");
assert!(probs[1] > 0.0, "token 1 (0.3 >= 0.2) should be kept");
assert!(
probs[2] > 0.0,
"token 2 (0.2 >= threshold boundary) should be kept"
);
assert!(
(probs[3]).abs() < 1e-6,
"token 3 (0.1 < 0.2) should be dropped"
);
let sum: f32 = probs.iter().sum();
assert!(
(sum - 1.0).abs() < 1e-5,
"kept tokens should renormalize to 1.0"
);
}
#[test]
fn apply_min_p_with_zero_is_noop() {
let mut probs = vec![0.5f32, 0.3, 0.2];
let original = probs.clone();
apply_min_p(&mut probs, 0.0);
assert_eq!(probs, original);
}
#[test]
fn build_ngram_index_maps_trigrams() {
let tokens = vec![1, 2, 3, 1, 2, 4];
let index = build_ngram_index(&tokens, 3);
assert_eq!(index.get([1, 2, 3].as_slice()).unwrap(), [1].as_slice());
assert_eq!(index.get([2, 3, 1].as_slice()).unwrap(), [2].as_slice());
assert_eq!(index.get([3, 1, 2].as_slice()).unwrap(), [4].as_slice());
}
#[test]
fn draft_tokens_continues_trigram() {
let tokens = vec![1, 2, 3, 1, 2, 4, 1, 2, 3, 5];
let index = build_ngram_index(&tokens, 3);
let draft = draft_tokens(&index, &[1, 2, 3], 3, 2);
assert!(!draft.is_empty(), "should produce at least one draft token");
}
#[test]
fn draft_tokens_stops_when_no_match() {
let index = build_ngram_index(&[1, 2, 3, 4], 3);
let draft = draft_tokens(&index, &[9, 9, 9], 3, 3);
assert!(draft.is_empty(), "unknown n-gram should yield empty draft");
}
#[test]
fn most_frequent_picks_mode() {
assert_eq!(most_frequent(&[1, 2, 2, 3, 2]), 2);
assert_eq!(most_frequent(&[5]), 5);
}
#[test]
fn cancel_flag_defaults_to_not_cancelled() {
let config = GenerationConfig::default();
assert!(!config.cancelled(), "no cancel flag => not cancelled");
}
#[test]
fn cancel_flag_set_is_cancelled() {
let flag = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
let config = GenerationConfig {
cancel: Some(flag.clone()),
..GenerationConfig::default()
};
assert!(!config.cancelled(), "flag false => not cancelled");
flag.store(true, std::sync::atomic::Ordering::Relaxed);
assert!(config.cancelled(), "flag true => cancelled");
}
#[test]
fn update_ngram_index_adds_latest() {
let mut index = build_ngram_index(&[1, 2, 3, 4], 3);
index = update_ngram_index(&index, &[1, 2, 3, 4, 5], 3);
assert_eq!(index.get([2, 3, 4].as_slice()).unwrap(), [5].as_slice());
}
#[test]
fn generate_token_ids_returns_empty_for_unknown_arch() {
let runtime = Runtime::from_raw(&std::collections::HashMap::new());
let tokenizer = BpeTokenizer::byte_fallback();
let config = GenerationConfig {
max_tokens: 10,
temperature: 0.0,
..GenerationConfig::default()
};
let ids = generate_token_ids(&runtime, "unknown", 4, &tokenizer, "hello", &config, None);
assert!(
ids.is_empty(),
"unknown architecture should yield no tokens"
);
}
#[test]
fn logprob_for_step_argmax_has_highest_logprob() {
let lp = logprob_for_step(&[1.0, 5.0, 2.0], 1, 0);
assert_eq!(lp.token, 1);
assert!(lp.logprob.is_finite());
let lp0 = logprob_for_step(&[1.0, 5.0, 2.0], 0, 0);
let lp2 = logprob_for_step(&[1.0, 5.0, 2.0], 2, 0);
assert!(lp.logprob > lp0.logprob);
assert!(lp.logprob > lp2.logprob);
}
#[test]
fn logprob_for_step_exp_sums_to_one() {
let logits = [1.0, 2.0, 3.0, 4.0];
let sum: f32 = logits
.iter()
.enumerate()
.map(|(i, _)| logprob_for_step(&logits, i as u32, 0).logprob.exp())
.sum();
assert!((sum - 1.0).abs() < 1e-5, "exp(logprob) sum = {sum}");
}
#[test]
fn logprob_for_step_top_logprobs_sorted_desc_and_limited() {
let lp = logprob_for_step(&[0.1, 0.9, 5.0, 0.2], 2, 2);
assert_eq!(lp.top_logprobs.len(), 2);
assert_eq!(lp.top_logprobs[0].0, 2);
assert_eq!(lp.top_logprobs[1].0, 1);
assert!(lp.top_logprobs[0].1 > lp.top_logprobs[1].1);
}
#[test]
fn logprob_for_step_empty_logits_yields_neg_infinity() {
let lp = logprob_for_step(&[], 0, 3);
assert_eq!(lp.token, 0);
assert!(lp.logprob.is_infinite() && lp.logprob.is_sign_negative());
assert!(lp.top_logprobs.is_empty());
}
#[test]
fn multinomial_same_seed_same_step_is_deterministic() {
let probs = vec![0.1f32, 0.2, 0.3, 0.4];
let seed = 42u64;
let a = multinomial(&probs, 0, Some(seed));
let b = multinomial(&probs, 0, Some(seed));
assert_eq!(a, b, "same seed + step should produce identical sample");
}
#[test]
fn multinomial_same_seed_different_step_differs() {
let probs = vec![0.1f32, 0.2, 0.3, 0.4];
let seed = 42u64;
let a = multinomial(&probs, 0, Some(seed));
let b = multinomial(&probs, 1, Some(seed));
assert_ne!(
a, b,
"same seed but different step should produce different sample"
);
}
#[test]
fn sample_with_seed_is_deterministic() {
let config = GenerationConfig {
max_tokens: 1,
temperature: 1.0,
seed: Some(123),
..GenerationConfig::default()
};
let logits = vec![1.0f32, 2.0, 3.0];
let a = sample(Some(&logits), &config, 0);
let b = sample(Some(&logits), &config, 0);
assert_eq!(
a, b,
"sample with same seed and step should be deterministic"
);
}
#[test]
fn repetition_penalty_disabled_leaves_logits_unchanged() {
let mut logits = vec![1.0f32, 2.0, -1.0, -2.0];
let original = logits.clone();
apply_repetition_penalty(&mut logits, &[0, 1], 1.0);
assert_eq!(logits, original, "penalty=1.0 should be a no-op");
}
#[test]
fn repetition_penalty_reduces_seen_token_logits() {
let mut logits = vec![2.0f32, 1.0, -1.0, -2.0];
apply_repetition_penalty(&mut logits, &[0], 2.0);
assert!(
(logits[0] - 1.0).abs() < 1e-6,
"positive seen logit should be divided by penalty"
);
assert!(
(logits[1] - 1.0).abs() < 1e-6,
"unseen logit should be unchanged"
);
}
#[test]
fn repetition_penalty_amplifies_negative_logits() {
let mut logits = vec![-2.0f32];
apply_repetition_penalty(&mut logits, &[0], 2.0);
assert!(
(logits[0] - (-4.0)).abs() < 1e-6,
"negative seen logit should be amplified"
);
}
#[test]
fn repetition_penalty_penalizes_each_seen_token_once() {
let mut logits = vec![4.0f32, 3.0, 2.0];
apply_repetition_penalty(&mut logits, &[0, 0, 0], 2.0);
assert!(
(logits[0] - 2.0).abs() < 1e-6,
"duplicate seen tokens should be penalized once"
);
assert!((logits[1] - 3.0).abs() < 1e-6, "unseen token 1 unchanged");
assert!((logits[2] - 2.0).abs() < 1e-6, "unseen token 2 unchanged");
}
#[test]
fn presence_frequency_penalty_disabled_is_noop() {
let mut logits = vec![1.0f32, 2.0, 3.0];
let original = logits.clone();
apply_presence_frequency_penalty(&mut logits, &[0, 1, 2], 0.0, 0.0);
assert_eq!(logits, original, "zero penalties should be a no-op");
}
#[test]
fn presence_penalty_subtracts_once_per_seen_token() {
let mut logits = vec![1.0f32, 2.0, 3.0];
apply_presence_frequency_penalty(&mut logits, &[0, 0], 0.5, 0.0);
assert!(
(logits[0] - 0.5).abs() < 1e-6,
"seen token 0 reduced by presence penalty"
);
assert!((logits[1] - 2.0).abs() < 1e-6, "unseen token 1 unchanged");
assert!((logits[2] - 3.0).abs() < 1e-6, "unseen token 2 unchanged");
}
#[test]
fn frequency_penalty_scales_with_count() {
let mut logits = vec![2.0f32, 1.0];
apply_presence_frequency_penalty(&mut logits, &[0, 0, 0], 0.0, 0.2);
assert!(
(logits[0] - 1.4).abs() < 1e-6,
"seen token 0 reduced by count * frequency penalty"
);
assert!((logits[1] - 1.0).abs() < 1e-6, "unseen token 1 unchanged");
}
#[test]
fn presence_and_frequency_combine() {
let mut logits = vec![5.0f32];
apply_presence_frequency_penalty(&mut logits, &[0, 0], 0.5, 0.3);
assert!(
(logits[0] - 3.9).abs() < 1e-6,
"presence + frequency should both apply: got {}",
logits[0]
);
}
#[test]
fn presence_frequency_penalty_empty_tokens_is_noop() {
let mut logits = vec![1.0f32, 2.0];
let original = logits.clone();
apply_presence_frequency_penalty(&mut logits, &[], 1.0, 1.0);
assert_eq!(
logits, original,
"no generated tokens => no penalty applied"
);
}
#[test]
fn logit_bias_boosts_token_in_greedy() {
let mut bias = HashMap::new();
bias.insert(0u32, 10.0f32);
let config = GenerationConfig {
max_tokens: 1,
temperature: 0.0,
logit_bias: bias,
..GenerationConfig::default()
};
let token = sample_constrained(
Some(&[1.0f32, 5.0, 2.0]),
&config,
&BpeTokenizer::byte_fallback(),
&[],
0,
);
assert_eq!(token, 0, "positive bias should boost token 0 above token 1");
}
#[test]
fn logit_bias_suppresses_token_in_greedy() {
let mut bias = HashMap::new();
bias.insert(1u32, -100.0f32);
let config = GenerationConfig {
max_tokens: 1,
temperature: 0.0,
logit_bias: bias,
..GenerationConfig::default()
};
let token = sample_constrained(
Some(&[1.0f32, 5.0, 2.0]),
&config,
&BpeTokenizer::byte_fallback(),
&[],
0,
);
assert_ne!(token, 1, "negative bias should suppress token 1");
assert_eq!(token, 2, "token 2 should be picked instead");
}
#[test]
fn logit_bias_empty_is_fast_path() {
let config = GenerationConfig {
max_tokens: 1,
temperature: 0.0,
..GenerationConfig::default()
};
assert!(config.logit_bias.is_empty());
let token = sample_constrained(
Some(&[1.0f32, 5.0, 2.0]),
&config,
&BpeTokenizer::byte_fallback(),
&[],
0,
);
assert_eq!(token, 1);
}
}
struct MiniRng(u64);
impl MiniRng {
fn new() -> Self {
let seed = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0xdead_beef_cafe_babe);
Self(Self::mix(seed))
}
fn from_seed(seed: u64) -> Self {
Self(Self::mix(seed))
}
fn next_u64(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
Self::mix(self.0)
}
fn gen_f32(&mut self) -> f32 {
(self.next_u64() >> 40) as f32 / (1u32 << 24) as f32
}
fn mix(x: u64) -> u64 {
let mut z = x.wrapping_add(0x9E37_79B9_7F4A_7C15);
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
}
#[cfg(test)]
mod mini_rng_tests {
use super::MiniRng;
#[test]
fn same_seed_same_sequence() {
let mut a = MiniRng::from_seed(42);
let mut b = MiniRng::from_seed(42);
for _ in 0..100 {
assert_eq!(a.next_u64(), b.next_u64());
}
}
#[test]
fn different_seeds_diverge() {
let mut a = MiniRng::from_seed(1);
let mut b = MiniRng::from_seed(2);
assert_ne!(a.next_u64(), b.next_u64());
}
#[test]
fn f32_in_unit_range() {
let mut rng = MiniRng::from_seed(123);
for _ in 0..10_000 {
let v = rng.gen_f32();
assert!((0.0..1.0).contains(&v));
}
}
}