use super::{AprKVCache, AprTransformer, GenerateConfig};
use crate::error::{RealizarError, Result};
use rand::rngs::StdRng;
use rand::SeedableRng;
#[inline]
fn is_eos_token(token: u32, stop_tokens: &[u32]) -> bool {
token == 0 || stop_tokens.contains(&token)
}
#[inline]
fn argmax_logits(logits: &[f32]) -> u32 {
logits
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map_or(0, |(idx, _)| idx as u32)
}
fn sample_from_logits(logits: &[f32], config: &GenerateConfig, rng: &mut StdRng) -> u32 {
if crate::sampling::is_greedy(config.temperature, config.top_k) {
return argmax_logits(logits);
}
crate::sampling::draw_seeded(logits, config.temperature, config.top_k, config.top_p, rng)
}
fn process_prompt_tokens(
model: &AprTransformer,
prompt: &[u32],
cache: &mut AprKVCache,
trace: bool,
) -> Result<Vec<f32>> {
if trace {
eprintln!("[TRACE] Processing {} prompt tokens...", prompt.len());
}
let mut logits = Vec::new();
for (pos, &token) in prompt.iter().enumerate() {
let start = std::time::Instant::now();
logits = model.forward_with_cache(token, cache, pos)?;
if trace {
eprintln!("[TRACE] Prompt token {}: {:?}", pos, start.elapsed());
}
}
Ok(logits)
}
fn generate_next_tokens(
model: &AprTransformer,
cache: &mut AprKVCache,
output: &mut Vec<u32>,
initial_logits: Vec<f32>,
config: &GenerateConfig,
trace: bool,
) -> Result<()> {
let mut logits = initial_logits;
let mut rng = StdRng::seed_from_u64(config.seed);
for i in 0..config.max_tokens {
if config.cancel.is_cancelled() {
break;
}
let next_token = sample_from_logits(&logits, config, &mut rng);
output.push(next_token);
if is_eos_token(next_token, &config.stop_tokens) {
break;
}
if i < config.max_tokens - 1 {
let start = std::time::Instant::now();
logits = model.forward_with_cache(next_token, cache, output.len() - 1)?;
if trace {
eprintln!(
"[TRACE] Gen token {} (pos {}): {:?}",
i,
output.len() - 1,
start.elapsed()
);
}
}
}
Ok(())
}
pub(crate) fn generate_with_cache(
model: &AprTransformer,
prompt: &[u32],
config: &GenerateConfig,
) -> Result<Vec<u32>> {
if prompt.is_empty() {
return Err(RealizarError::InvalidShape {
reason: "Prompt cannot be empty".to_string(),
});
}
let trace = std::env::var("REALIZE_TRACE").is_ok();
let mut cache = AprKVCache::new(&model.config);
let mut output = prompt.to_vec();
let logits = process_prompt_tokens(model, prompt, &mut cache, trace)?;
generate_next_tokens(model, &mut cache, &mut output, logits, config, trace)?;
if trace {
eprintln!(
"[TRACE] Generation complete. Total output tokens: {}",
output.len()
);
}
Ok(output)
}
fn forward_with_trace(
model: &AprTransformer,
token: u32,
cache: &mut AprKVCache,
pos: usize,
step: usize,
trace: bool,
) -> Result<Vec<f32>> {
let start = std::time::Instant::now();
let logits = model.forward_with_cache(token, cache, pos)?;
if trace {
eprintln!(
"[TRACE] Gen token {} (pos {}): {:?}",
step,
pos,
start.elapsed()
);
}
Ok(logits)
}
fn trace_generation_complete(trace: bool, total_tokens: usize) {
if trace {
eprintln!(
"[TRACE] Streaming generation complete. Total output tokens: {}",
total_tokens
);
}
}
pub(crate) fn generate_with_cache_streaming<F>(
model: &AprTransformer,
prompt: &[u32],
config: &GenerateConfig,
mut on_token: F,
) -> Result<Vec<u32>>
where
F: FnMut(u32) -> bool,
{
if prompt.is_empty() {
return Err(RealizarError::InvalidShape {
reason: "Prompt cannot be empty".to_string(),
});
}
let trace = std::env::var("REALIZE_TRACE").is_ok();
let mut cache = AprKVCache::new(&model.config);
let mut output = prompt.to_vec();
let logits = process_prompt_tokens(model, prompt, &mut cache, trace)?;
let mut logits = logits;
let mut rng = StdRng::seed_from_u64(config.seed);
for i in 0..config.max_tokens {
if config.cancel.is_cancelled() {
break;
}
let next_token = sample_from_logits(&logits, config, &mut rng);
output.push(next_token);
if is_eos_token(next_token, &config.stop_tokens) {
break;
}
if !on_token(next_token) {
break;
}
if i < config.max_tokens - 1 {
logits = forward_with_trace(model, next_token, &mut cache, output.len() - 1, i, trace)?;
}
}
trace_generation_complete(trace, output.len());
Ok(output)
}
#[cfg(test)]
mod sampler_tests {
use super::{argmax_logits, sample_from_logits};
use crate::apr_transformer::GenerateConfig;
use rand::rngs::StdRng;
use rand::SeedableRng;
fn cfg(temperature: f32, top_k: usize, top_p: f32) -> GenerateConfig {
GenerateConfig {
max_tokens: 1,
temperature,
top_p,
top_k,
..GenerateConfig::default()
}
}
const FLAT: [f32; 6] = [1.0, 0.9, 1.1, 0.95, 1.05, 0.85];
fn draws(config: &GenerateConfig, seed: u64, n: usize) -> Vec<u32> {
let mut rng = StdRng::seed_from_u64(seed);
(0..n)
.map(|_| sample_from_logits(&FLAT, config, &mut rng))
.collect()
}
#[test]
fn a_sampled_step_draws_off_the_argmax() {
let greedy = argmax_logits(&FLAT);
let tokens = draws(&cfg(1.0, 0, 1.0), 7, 64);
assert!(
tokens.iter().any(|&t| t != greedy),
"64 draws at temperature 1.0 all returned the argmax: the sampler does not draw"
);
}
#[test]
fn the_same_seed_reproduces_and_another_seed_changes_the_draws() {
let config = cfg(1.0, 0, 1.0);
assert_eq!(draws(&config, 7, 32), draws(&config, 7, 32));
assert_ne!(
draws(&config, 7, 32),
draws(&config, 8, 32),
"the seed is not used"
);
}
#[test]
fn temperature_zero_or_top_k_one_is_the_argmax() {
let expected = argmax_logits(&FLAT);
for seed in 0..8 {
assert!(draws(&cfg(0.0, 0, 1.0), seed, 8)
.iter()
.all(|&t| t == expected));
assert!(draws(&cfg(1.0, 1, 1.0), seed, 8)
.iter()
.all(|&t| t == expected));
assert!(draws(&cfg(0.0, 40, 0.1), seed, 8)
.iter()
.all(|&t| t == expected));
}
}
#[test]
fn top_k_bounds_the_draw() {
let tokens = draws(&cfg(5.0, 2, 1.0), 3, 200);
assert!(tokens.iter().all(|&t| t == 2 || t == 4), "{tokens:?}");
assert!(
tokens.contains(&2) && tokens.contains(&4),
"both survivors are drawn"
);
}
}