use std::sync::Arc;
use crate::errors::{Result, TrustformersError};
use crate::tensor::Tensor;
#[derive(Debug, Clone, Default)]
pub struct KVCache {
pub keys: Vec<Tensor>,
pub values: Vec<Tensor>,
pub seq_len: usize,
}
impl KVCache {
pub fn new() -> Self {
Self {
keys: Vec::new(),
values: Vec::new(),
seq_len: 0,
}
}
pub fn push_layer(&mut self, key: Tensor, value: Tensor) -> Result<()> {
if self.keys.len() != self.values.len() {
return Err(TrustformersError::invalid_input(
"Key-value cache size mismatch".to_string(),
));
}
self.keys.push(key);
self.values.push(value);
Ok(())
}
pub fn set_layer(&mut self, layer_idx: usize, key: Tensor, value: Tensor) -> Result<()> {
if layer_idx >= self.keys.len() || layer_idx >= self.values.len() {
return Err(TrustformersError::invalid_input(format!(
"layer {layer_idx} is not present in a cache with {} layers",
self.keys.len()
)));
}
self.keys[layer_idx] = key;
self.values[layer_idx] = value;
Ok(())
}
pub fn advance(&mut self, tokens: usize) {
self.seq_len += tokens;
}
pub fn num_layers(&self) -> usize {
self.keys.len()
}
pub fn clear(&mut self) {
self.keys.clear();
self.values.clear();
self.seq_len = 0;
}
pub fn get_layer(&self, layer_idx: usize) -> Option<(&Tensor, &Tensor)> {
if layer_idx < self.keys.len() {
Some((&self.keys[layer_idx], &self.values[layer_idx]))
} else {
None
}
}
}
#[derive(Debug, Clone)]
pub struct Beam {
pub tokens: Vec<usize>,
pub score: f32,
pub finished: bool,
pub cache: Option<Arc<KVCache>>,
}
impl Beam {
pub fn new(tokens: Vec<usize>, score: f32) -> Self {
Self {
tokens,
score,
finished: false,
cache: None,
}
}
pub fn with_cache(mut self, cache: KVCache) -> Self {
self.cache = Some(Arc::new(cache));
self
}
pub fn extend(&self, token: usize, score: f32) -> Self {
let mut new_tokens = Vec::with_capacity(self.tokens.len() + 1);
new_tokens.extend_from_slice(&self.tokens);
new_tokens.push(token);
Self {
tokens: new_tokens,
score: self.score + score,
finished: false,
cache: self.cache.clone(),
}
}
pub fn cache_mut(&mut self) -> Option<&mut KVCache> {
self.cache.as_mut().map(Arc::make_mut)
}
pub fn finalize(&mut self) {
self.finished = true;
}
pub fn get_normalized_score(&self) -> f32 {
if self.tokens.is_empty() {
0.0
} else {
self.score / self.tokens.len() as f32
}
}
pub fn length_normalized_score(&self, length_penalty: f32) -> f32 {
let length = self.tokens.len().max(1) as f32;
self.score / length.powf(length_penalty)
}
}