use rten_tensor::prelude::*;
use rten_tensor::{NdTensor, NdTensorView};
use crate::generator::TokenId;
pub trait LogitsFilter {
fn filter(
&self,
logits: NdTensorView<f32, 1>,
prev_tokens: &[TokenId],
) -> Option<NdTensor<f32, 1>>;
}
struct TokenIdFilter<F: Fn(TokenId) -> bool> {
predicate: F,
}
impl<F: Fn(TokenId) -> bool> LogitsFilter for TokenIdFilter<F> {
fn filter(
&self,
logits: NdTensorView<f32, 1>,
_prev_tokens: &[TokenId],
) -> Option<NdTensor<f32, 1>> {
Some(NdTensor::from_fn(logits.shape(), |[i]| {
let token_id = i as TokenId;
if (self.predicate)(token_id) {
logits[[i]]
} else {
f32::NEG_INFINITY
}
}))
}
}
pub fn token_id_filter<F: Fn(TokenId) -> bool>(predicate: F) -> impl LogitsFilter {
TokenIdFilter { predicate }
}
#[cfg(test)]
mod tests {
use rten_tensor::NdTensor;
use rten_tensor::prelude::*;
use super::{LogitsFilter, token_id_filter};
#[test]
fn test_token_id_filter() {
let logits = NdTensor::from([0., 1., 2., 3., 4.]);
let filter = token_id_filter(|id| id % 2 == 0);
let output = filter.filter(logits.view(), &[]);
assert_eq!(
output,
Some(NdTensor::from([
0.,
f32::NEG_INFINITY,
2.,
f32::NEG_INFINITY,
4.
]))
);
}
}