#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FinishReason {
Stop,
Length,
}
impl FinishReason {
#[must_use]
pub fn from_generation(stopped: bool, completion_tokens: usize, max_tokens: usize) -> Self {
if !stopped && completion_tokens >= max_tokens {
Self::Length
} else {
Self::Stop
}
}
#[must_use]
pub fn from_decode(generated: &[u32], stop_tokens: &[u32], budget: usize) -> Self {
let ended_on_stop = generated.last().is_some_and(|t| stop_tokens.contains(t));
Self::from_generation(ended_on_stop, generated.len(), budget)
}
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
Self::Stop => "stop",
Self::Length => "length",
}
}
}
impl std::fmt::Display for FinishReason {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[must_use]
pub fn decode_budget(
max_tokens: usize,
prompt_len: usize,
context_length: usize,
clamps_to_context: bool,
) -> usize {
if clamps_to_context {
max_tokens.min(context_length.saturating_sub(prompt_len))
} else {
max_tokens
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct RunReport {
pub finish_reason: Option<FinishReason>,
pub context_length: Option<usize>,
}
#[cfg(test)]
mod tests {
use super::*;
const EOS: u32 = 7;
#[test]
fn from_decode_case_table() {
let rows: &[(&str, &[u32], usize, FinishReason)] = &[
(
"pushed stop, under budget",
&[5, 6, EOS],
8,
FinishReason::Stop,
),
(
"pushed stop as the last budgeted token",
&[5, 6, EOS],
3,
FinishReason::Stop,
),
(
"consumed stop, under budget",
&[5, 6],
8,
FinishReason::Stop,
),
(
"consumed stop, nothing generated",
&[],
8,
FinishReason::Stop,
),
("budget spent, no stop", &[5, 6, 9], 3, FinishReason::Length),
("zero budget", &[], 0, FinishReason::Length),
("clamped budget spent", &[5, 6], 2, FinishReason::Length),
];
for (name, generated, budget, want) in rows {
assert_eq!(
FinishReason::from_decode(generated, &[EOS], *budget),
*want,
"{name}"
);
}
}
#[test]
fn clamped_cut_is_length_only_against_the_budget_it_ran_with() {
let generated = [5, 6, 9, 10];
let requested = 256;
let ran_with = decode_budget(requested, 60, 64, true);
assert_eq!(ran_with, 4);
assert_eq!(
FinishReason::from_decode(&generated, &[EOS], requested),
FinishReason::Stop
);
assert_eq!(
FinishReason::from_decode(&generated, &[EOS], ran_with),
FinishReason::Length
);
}
#[test]
fn decode_budget_clamps_only_when_the_loop_does() {
assert_eq!(decode_budget(100, 10, 64, true), 54);
assert_eq!(decode_budget(100, 10, 64, false), 100);
assert_eq!(decode_budget(20, 10, 64, true), 20);
assert_eq!(decode_budget(100, 70, 64, true), 0);
}
#[test]
fn wire_strings_are_the_openai_ones() {
assert_eq!(FinishReason::Stop.as_str(), "stop");
assert_eq!(FinishReason::Length.to_string(), "length");
}
#[test]
fn an_unreported_run_is_none_not_a_default_reason() {
let report = RunReport::default();
assert_eq!(report.finish_reason, None);
assert_eq!(report.context_length, None);
}
const PYGMY_CONTEXT: usize = 32;
fn pygmy_file() -> tempfile::NamedTempFile {
use std::io::Write;
let mut file = tempfile::NamedTempFile::with_suffix(".gguf").expect("temp gguf");
file.write_all(&crate::gguf::test_factory::build_executable_pygmy_gguf())
.expect("write gguf");
file.flush().expect("flush gguf");
file
}
fn run(
file: &tempfile::NamedTempFile,
prompt: Vec<u32>,
max_tokens: usize,
stop: Vec<u32>,
) -> crate::error::Result<(super::super::InferenceResult, RunReport)> {
let config = super::super::InferenceConfig::new(file.path())
.with_input_tokens(prompt)
.with_max_tokens(max_tokens)
.with_stop_tokens(stop)
.without_gpu();
super::super::run_inference_report(&config)
}
#[test]
fn context_clamped_cut_reports_length_with_the_counts() {
let file = pygmy_file();
let prompt = vec![1, 2, 3, 4];
let (result, report) = run(&file, prompt.clone(), 1000, vec![]).expect("run");
assert_eq!(result.input_token_count, prompt.len());
assert_eq!(report.context_length, Some(PYGMY_CONTEXT));
assert_eq!(
result.generated_token_count,
PYGMY_CONTEXT - prompt.len(),
"control: the loop must run to the context edge, or this is not measuring a clamp"
);
assert_eq!(report.finish_reason, Some(FinishReason::Length));
}
#[test]
fn max_tokens_cut_reports_length() {
let file = pygmy_file();
let (result, report) = run(&file, vec![1, 2, 3, 4], 3, vec![]).expect("run");
assert_eq!(result.generated_token_count, 3);
assert_eq!(report.finish_reason, Some(FinishReason::Length));
}
#[test]
fn a_stop_token_under_budget_reports_stop() {
let file = pygmy_file();
let prompt = vec![1, 2, 3, 4];
let (first, _) = run(&file, prompt.clone(), 1, vec![]).expect("probe run");
let first_token = first.tokens[prompt.len()];
let (result, report) = run(&file, prompt, 10, vec![first_token]).expect("run");
assert_eq!(
result.generated_token_count, 0,
"control: greedy is deterministic, so the probed token must stop the run at once"
);
assert_eq!(report.finish_reason, Some(FinishReason::Stop));
}
#[test]
fn a_fixed_text_prompt_counts_the_tokens_the_tokenizer_feeds() {
use std::io::Write;
let mut vocab: Vec<String> = ["<unk>", "<s>", "</s>", "▁Hello", "▁world", "!"]
.map(String::from)
.to_vec();
vocab.extend((vocab.len()..32).map(|i| format!("<filler_{i}>")));
let pieces: Vec<&str> = vocab.iter().map(String::as_str).collect();
let bytes = crate::gguf::test_factory::build_executable_pygmy_gguf_with(|b| {
b.add_string("tokenizer.ggml.model", "llama")
.add_string_array("tokenizer.ggml.tokens", &pieces)
.add_u32("tokenizer.ggml.bos_token_id", 1)
.add_u32("tokenizer.ggml.eos_token_id", 2)
});
let mut file = tempfile::NamedTempFile::with_suffix(".gguf").expect("temp gguf");
file.write_all(&bytes).expect("write gguf");
file.flush().expect("flush gguf");
let config = super::super::InferenceConfig::new(file.path())
.with_prompt("Hello world!")
.with_max_tokens(2)
.without_gpu();
let (result, report) = super::super::run_inference_report(&config).expect("run");
assert_eq!(
result.tokens.get(..4),
Some(&[1, 3, 4, 5][..]),
"the prompt must be fed as <s> ▁Hello ▁world !"
);
assert_eq!(result.input_token_count, 4, "prompt_tokens is what was fed");
assert_eq!(report.context_length, Some(PYGMY_CONTEXT));
}
#[test]
fn a_prompt_longer_than_the_context_is_refused() {
let file = pygmy_file();
let prompt: Vec<u32> = (0..(PYGMY_CONTEXT as u32 + 4)).map(|t| t % 32).collect();
let err = run(&file, prompt, 4, vec![]).expect_err("an over-long prompt must be refused");
let msg = err.to_string();
assert!(
msg.contains(&(PYGMY_CONTEXT + 4).to_string())
&& msg.contains(&PYGMY_CONTEXT.to_string()),
"the refusal must name the prompt length and the context: {msg}"
);
}
}