use crate::generator::TokenId;
#[derive(Clone, Debug, Default, PartialEq)]
pub struct Logits {
logits: Vec<f32>,
indices: Vec<TokenId>,
}
impl Logits {
pub fn dense(logits: Vec<f32>) -> Logits {
assert!(logits.len() <= u32::MAX as usize);
let indices = (0..logits.len() as TokenId).collect();
Self { logits, indices }
}
pub fn sparse(logits: Vec<f32>, indices: Vec<TokenId>) -> Logits {
assert_eq!(logits.len(), indices.len());
Self { logits, indices }
}
pub fn into_logits_indices(self) -> (Vec<f32>, Vec<TokenId>) {
(self.logits, self.indices)
}
pub fn len(&self) -> usize {
self.logits.len()
}
pub fn is_empty(&self) -> bool {
self.logits.is_empty()
}
pub fn logits(&self) -> &[f32] {
&self.logits
}
pub fn indices(&self) -> &[TokenId] {
&self.indices
}
pub fn enumerate(&self) -> impl Iterator<Item = (TokenId, f32)> {
self.indices
.iter()
.zip(&self.logits)
.map(|(token_id, logit)| (*token_id, *logit))
}
}