use kopitiam_core::Result;
use kopitiam_tokenizer::Tokenizer;
use crate::constraint::{ConstrainedSampler, ConstraintError, TokenConstraint};
use crate::sampling::{GreedySampler, Sampler};
use crate::traits::Model;
#[derive(Debug, Clone)]
pub struct GenerationConfig {
pub max_new_tokens: usize,
pub eos_token_id: Option<u32>,
}
impl Default for GenerationConfig {
fn default() -> Self {
Self { max_new_tokens: 256, eos_token_id: None }
}
}
pub fn generate<M: Model>(
model: &M,
tokenizer: &dyn Tokenizer,
prompt: &str,
config: &GenerationConfig,
on_token: impl FnMut(u32, &str),
) -> Result<String> {
generate_with_sampler(model, tokenizer, prompt, config, &mut GreedySampler, on_token)
}
pub fn generate_with_sampler<M: Model>(
model: &M,
tokenizer: &dyn Tokenizer,
prompt: &str,
config: &GenerationConfig,
sampler: &mut dyn Sampler,
mut on_token: impl FnMut(u32, &str),
) -> Result<String> {
let prompt_ids = tokenizer.encode(prompt)?;
let mut cache = model.new_cache();
let mut generated_ids: Vec<u32> = Vec::new();
if prompt_ids.is_empty() {
return Ok(String::new());
}
let logits = model.forward(&prompt_ids, &mut cache)?;
let mut next = sampler.sample(&last_row(&logits, model.vocab_size())?);
for _ in 0..config.max_new_tokens {
if config.eos_token_id == Some(next) {
break;
}
generated_ids.push(next);
let token_text = tokenizer.decode(&[next])?;
on_token(next, &token_text);
let logits = model.forward(&[next], &mut cache)?;
next = sampler.sample(&last_row(&logits, model.vocab_size())?);
}
tokenizer.decode(&generated_ids)
}
#[derive(Debug)]
pub enum ConstrainedGenerateError {
Runtime(kopitiam_core::Error),
Constraint(ConstraintError),
}
impl std::fmt::Display for ConstrainedGenerateError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ConstrainedGenerateError::Runtime(e) => write!(f, "{e}"),
ConstrainedGenerateError::Constraint(e) => write!(f, "constrained decoding failed: {e}"),
}
}
}
impl std::error::Error for ConstrainedGenerateError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
ConstrainedGenerateError::Runtime(e) => Some(e),
ConstrainedGenerateError::Constraint(e) => Some(e),
}
}
}
impl From<kopitiam_core::Error> for ConstrainedGenerateError {
fn from(e: kopitiam_core::Error) -> Self {
ConstrainedGenerateError::Runtime(e)
}
}
impl From<ConstraintError> for ConstrainedGenerateError {
fn from(e: ConstraintError) -> Self {
ConstrainedGenerateError::Constraint(e)
}
}
pub fn generate_constrained<M, C, S>(
model: &M,
tokenizer: &dyn Tokenizer,
prompt: &str,
config: &GenerationConfig,
constrained: &mut ConstrainedSampler<C, S>,
mut on_token: impl FnMut(u32, &str),
) -> std::result::Result<String, ConstrainedGenerateError>
where
M: Model,
C: TokenConstraint,
S: Sampler,
{
let prompt_ids = tokenizer.encode(prompt)?;
let mut cache = model.new_cache();
let mut generated_ids: Vec<u32> = Vec::new();
if prompt_ids.is_empty() {
return Ok(String::new());
}
let logits = model.forward(&prompt_ids, &mut cache)?;
let mut next = constrained.try_sample(&last_row(&logits, model.vocab_size())?)?;
for _ in 0..config.max_new_tokens {
if config.eos_token_id == Some(next) {
break;
}
generated_ids.push(next);
let token_text = tokenizer.decode(&[next])?;
on_token(next, &token_text);
let logits = model.forward(&[next], &mut cache)?;
next = constrained.try_sample(&last_row(&logits, model.vocab_size())?)?;
}
Ok(tokenizer.decode(&generated_ids)?)
}
fn last_row(logits: &kopitiam_tensor::Tensor, vocab_size: usize) -> Result<Vec<f32>> {
let data = logits.to_vec_f32()?;
let start = data.len() - vocab_size;
Ok(data[start..].to_vec())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::QwenModel;
use crate::test_support::synthetic_gguf::{build, write_temp_gguf, SyntheticModelSpec};
use kopitiam_tensor::Tensor;
fn tiny_tokenizer() -> kopitiam_tokenizer::BpeTokenizer {
let mut vocab: Vec<Vec<u8>> = (0u16..=255).map(|b| vec![b as u8]).collect();
vocab.push(b"ab".to_vec()); vocab.push(b"abc".to_vec()); let merges = vec![(b"a".to_vec(), b"b".to_vec()), (b"ab".to_vec(), b"c".to_vec())];
kopitiam_tokenizer::BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap()
}
fn model_matching_tokenizer() -> QwenModel {
let spec = SyntheticModelSpec { vocab_size: 258, ..SyntheticModelSpec::default() };
let bytes = build(&spec);
let path = write_temp_gguf(&bytes, "generate-e2e");
let loaded = kopitiam_loader::load_model(&path).unwrap();
QwenModel::from_loaded_model(&loaded).unwrap()
}
#[test]
fn generate_produces_at_most_max_new_tokens_and_streams_every_one() {
let model = model_matching_tokenizer();
let tokenizer = tiny_tokenizer();
let config = GenerationConfig { max_new_tokens: 5, eos_token_id: None };
let mut streamed: Vec<u32> = Vec::new();
let text = generate(&model, &tokenizer, "abc", &config, |id, _text| streamed.push(id)).unwrap();
assert!(streamed.len() <= 5);
assert!(!streamed.is_empty(), "greedy decoding with no EOS configured must run the full budget");
assert_eq!(streamed.len(), 5);
assert_eq!(text, tokenizer_decode_all(&tokenizer, &streamed));
}
fn tokenizer_decode_all(tokenizer: &kopitiam_tokenizer::BpeTokenizer, ids: &[u32]) -> String {
tokenizer.decode(ids).unwrap()
}
#[test]
fn an_eos_token_stops_generation_before_it_is_emitted() {
let model = model_matching_tokenizer();
let tokenizer = tiny_tokenizer();
let unbounded = GenerationConfig { max_new_tokens: 1, eos_token_id: None };
let mut first_token = None;
generate(&model, &tokenizer, "abc", &unbounded, |id, _| first_token = Some(id)).unwrap();
let first_token = first_token.expect("greedy decoding must produce a first token");
let with_eos = GenerationConfig { max_new_tokens: 10, eos_token_id: Some(first_token) };
let mut streamed = Vec::new();
let text = generate(&model, &tokenizer, "abc", &with_eos, |id, _| streamed.push(id)).unwrap();
assert!(streamed.is_empty(), "the EOS token itself must never be streamed");
assert_eq!(text, "");
}
#[test]
fn an_empty_prompt_generates_nothing() {
let model = model_matching_tokenizer();
let tokenizer = tiny_tokenizer();
let config = GenerationConfig { max_new_tokens: 5, eos_token_id: None };
let mut calls = 0;
let text = generate(&model, &tokenizer, "", &config, |_, _| calls += 1).unwrap();
assert_eq!(calls, 0);
assert_eq!(text, "");
}
#[test]
fn last_row_extracts_the_final_positions_logits() {
let logits = Tensor::from_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], [3, 2]).unwrap();
assert_eq!(last_row(&logits, 2).unwrap(), vec![5.0, 6.0]);
}
#[test]
fn zero_max_new_tokens_generates_nothing_but_still_runs_prefill() {
let model = model_matching_tokenizer();
let tokenizer = tiny_tokenizer();
let config = GenerationConfig { max_new_tokens: 0, eos_token_id: None };
let mut calls = 0;
let text = generate(&model, &tokenizer, "abc", &config, |_, _| calls += 1).unwrap();
assert_eq!(calls, 0);
assert_eq!(text, "");
}
#[test]
fn generate_with_a_greedy_sampler_matches_generates_own_output() {
use crate::sampling::GreedySampler;
let model = model_matching_tokenizer();
let tokenizer = tiny_tokenizer();
let config = GenerationConfig { max_new_tokens: 6, eos_token_id: None };
let via_generate = generate(&model, &tokenizer, "abc", &config, |_, _| {}).unwrap();
let via_sampler =
generate_with_sampler(&model, &tokenizer, "abc", &config, &mut GreedySampler, |_, _| {}).unwrap();
assert_eq!(via_generate, via_sampler);
}
#[test]
fn generate_with_sampler_and_a_fixed_seed_is_reproducible_end_to_end() {
use crate::sampling::{SamplingConfig, StochasticSampler};
let model = model_matching_tokenizer();
let tokenizer = tiny_tokenizer();
let config = GenerationConfig { max_new_tokens: 8, eos_token_id: None };
let run = || {
let mut sampler = StochasticSampler::new(SamplingConfig {
temperature: 1.0,
top_k: Some(5),
seed: 7,
..SamplingConfig::default()
});
generate_with_sampler(&model, &tokenizer, "abc", &config, &mut sampler, |_, _| {}).unwrap()
};
assert_eq!(run(), run(), "the same seed must reproduce the exact same completion end to end");
}
#[test]
fn constrained_generation_only_ever_emits_allowed_tokens_end_to_end() {
use crate::constraint::{AllowedTokens, ConstrainedSampler};
use crate::sampling::GreedySampler;
let model = model_matching_tokenizer();
let tokenizer = tiny_tokenizer();
let config = GenerationConfig { max_new_tokens: 12, eos_token_id: None };
let allowed: Vec<u32> = vec![10, 20, 30, 40];
let constraint = AllowedTokens::new(allowed.clone()).unwrap();
let mut sampler = ConstrainedSampler::new(constraint, GreedySampler);
let mut streamed: Vec<u32> = Vec::new();
let text =
generate_constrained(&model, &tokenizer, "abc", &config, &mut sampler, |id, _| streamed.push(id)).unwrap();
assert_eq!(streamed.len(), 12, "no EOS configured, so the full budget must run");
for id in &streamed {
assert!(allowed.contains(id), "constrained decode streamed disallowed token {id}");
}
assert_eq!(text, tokenizer.decode(&streamed).unwrap());
}
#[test]
fn constrained_generation_surfaces_a_dead_constraint_as_an_error() {
use crate::constraint::{AllowedTokens, ConstrainedSampler, ConstraintError};
use crate::sampling::GreedySampler;
let model = model_matching_tokenizer();
let tokenizer = tiny_tokenizer();
let config = GenerationConfig { max_new_tokens: 4, eos_token_id: None };
let constraint = AllowedTokens::new([9999]).unwrap();
let mut sampler = ConstrainedSampler::new(constraint, GreedySampler);
let err = generate_constrained(&model, &tokenizer, "abc", &config, &mut sampler, |_, _| {}).unwrap_err();
assert!(matches!(err, ConstrainedGenerateError::Constraint(ConstraintError::NoTokenAllowed)));
}
#[test]
fn generate_with_sampler_temperature_zero_matches_greedy_generate_end_to_end() {
use crate::sampling::{SamplingConfig, StochasticSampler};
let model = model_matching_tokenizer();
let tokenizer = tiny_tokenizer();
let config = GenerationConfig { max_new_tokens: 6, eos_token_id: None };
let greedy_text = generate(&model, &tokenizer, "abc", &config, |_, _| {}).unwrap();
let mut sampler = StochasticSampler::new(SamplingConfig { temperature: 0.0, ..SamplingConfig::default() });
let stochastic_text =
generate_with_sampler(&model, &tokenizer, "abc", &config, &mut sampler, |_, _| {}).unwrap();
assert_eq!(
greedy_text, stochastic_text,
"temperature=0.0 must reproduce plain greedy decoding through the full generate loop, not just per-call"
);
}
}