#[cfg(feature = "realizar-inference")]
use realizar::generate::{LogitProcessor, LogitProcessorContext};
use crate::tokenizer::special_tokens;
#[derive(Debug, Clone)]
pub struct WhisperTokenSuppressor {
suppress_ids: Vec<u32>,
suppress_timestamps: bool,
n_vocab: usize,
}
impl WhisperTokenSuppressor {
#[must_use]
pub fn new() -> Self {
let mut suppress_ids = vec![
special_tokens::SOT,
special_tokens::NO_SPEECH,
special_tokens::TRANSLATE,
special_tokens::TRANSCRIBE,
special_tokens::PREV,
special_tokens::SPEAKER_TURN,
special_tokens::NO_TIMESTAMPS,
];
for lang_id in special_tokens::LANG_BASE..special_tokens::TRANSLATE {
suppress_ids.push(lang_id);
}
Self {
suppress_ids,
suppress_timestamps: true,
n_vocab: 51865,
}
}
#[must_use]
pub fn with_tokens(tokens: Vec<u32>) -> Self {
Self {
suppress_ids: tokens,
suppress_timestamps: true,
n_vocab: 51865,
}
}
#[must_use]
pub fn with_timestamp_suppression(mut self, suppress: bool) -> Self {
self.suppress_timestamps = suppress;
self
}
#[must_use]
pub fn with_vocab_size(mut self, n_vocab: usize) -> Self {
self.n_vocab = n_vocab;
self
}
pub fn add_suppression(&mut self, token: u32) {
if !self.suppress_ids.contains(&token) {
self.suppress_ids.push(token);
}
}
#[must_use]
pub fn suppressed_tokens(&self) -> &[u32] {
&self.suppress_ids
}
#[must_use]
pub const fn suppresses_timestamps(&self) -> bool {
self.suppress_timestamps
}
pub fn apply(&self, logits: &mut [f32]) {
for &token_id in &self.suppress_ids {
let idx = token_id as usize;
if idx < logits.len() {
logits[idx] = f32::NEG_INFINITY;
}
}
if self.suppress_timestamps {
let timestamp_start = special_tokens::TIMESTAMP_BASE as usize;
for logit in logits
.iter_mut()
.skip(timestamp_start)
.take(self.n_vocab.saturating_sub(timestamp_start))
{
*logit = f32::NEG_INFINITY;
}
}
}
}
impl Default for WhisperTokenSuppressor {
fn default() -> Self {
Self::new()
}
}
#[cfg(feature = "realizar-inference")]
impl LogitProcessor for WhisperTokenSuppressor {
fn process(&self, logits: &mut [f32], _ctx: &LogitProcessorContext<'_>) {
self.apply(logits);
}
fn name(&self) -> &'static str {
"whisper_token_suppressor"
}
}
#[cfg(feature = "realizar-inference")]
pub use realizar::generate::{
LogitProcessorChain, RepetitionPenalty, TemperatureScaler, TokenSuppressor,
};
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_whisper_suppressor_default() {
let suppressor = WhisperTokenSuppressor::default();
assert!(suppressor
.suppressed_tokens()
.contains(&special_tokens::SOT));
assert!(suppressor
.suppressed_tokens()
.contains(&special_tokens::LANG_BASE));
assert!(suppressor.suppresses_timestamps());
}
#[test]
fn test_whisper_suppressor_apply() {
let suppressor = WhisperTokenSuppressor::default();
let mut logits = vec![1.0f32; 51865];
suppressor.apply(&mut logits);
assert!(logits[special_tokens::SOT as usize] == f32::NEG_INFINITY);
assert!(logits[special_tokens::LANG_BASE as usize] == f32::NEG_INFINITY);
assert!(logits[special_tokens::EOT as usize] != f32::NEG_INFINITY);
}
#[test]
fn test_whisper_suppressor_timestamps() {
let suppressor = WhisperTokenSuppressor::new();
let mut logits = vec![1.0f32; 51865];
suppressor.apply(&mut logits);
assert!(logits[special_tokens::TIMESTAMP_BASE as usize] == f32::NEG_INFINITY);
}
#[test]
fn test_whisper_suppressor_no_timestamps() {
let suppressor = WhisperTokenSuppressor::new().with_timestamp_suppression(false);
let mut logits = vec![1.0f32; 51865];
suppressor.apply(&mut logits);
assert!(logits[special_tokens::TIMESTAMP_BASE as usize] != f32::NEG_INFINITY);
}
#[test]
fn test_whisper_suppressor_custom_tokens() {
let suppressor = WhisperTokenSuppressor::with_tokens(vec![100, 200, 300]);
let mut logits = vec![1.0f32; 51865];
suppressor.apply(&mut logits);
assert!(logits[100] == f32::NEG_INFINITY);
assert!(logits[200] == f32::NEG_INFINITY);
assert!(logits[300] == f32::NEG_INFINITY);
assert!(logits[101] != f32::NEG_INFINITY);
}
#[test]
fn test_whisper_suppressor_add_token() {
let mut suppressor = WhisperTokenSuppressor::with_tokens(vec![]);
suppressor.add_suppression(500);
assert!(suppressor.suppressed_tokens().contains(&500));
}
#[test]
fn test_whisper_suppressor_bounds_check() {
let suppressor = WhisperTokenSuppressor::with_tokens(vec![100000]); let mut logits = vec![1.0f32; 100];
suppressor.apply(&mut logits);
}
#[test]
fn property_suppressor_idempotent() {
let suppressor = WhisperTokenSuppressor::default();
let mut logits1 = vec![1.0f32; 51865];
suppressor.apply(&mut logits1);
let mut logits2 = logits1.clone();
suppressor.apply(&mut logits2);
for (a, b) in logits1.iter().zip(logits2.iter()) {
assert!(
(a - b).abs() < 1e-6 || (a.is_infinite() && b.is_infinite()),
"Suppression should be idempotent"
);
}
}
#[test]
fn property_suppressor_preserves_eot() {
let suppressor = WhisperTokenSuppressor::default();
let mut logits = vec![1.0f32; 51865];
suppressor.apply(&mut logits);
assert!(
logits[special_tokens::EOT as usize].is_finite(),
"EOT must never be suppressed"
);
}
#[test]
fn test_with_vocab_size_affects_timestamp_suppression() {
let suppressor = WhisperTokenSuppressor::new().with_vocab_size(50400);
let mut logits = vec![1.0f32; 51865];
suppressor.apply(&mut logits);
assert!(
logits[50400].is_finite(),
"tokens beyond n_vocab should not be suppressed"
);
}
#[test]
fn test_with_vocab_size_large() {
let suppressor = WhisperTokenSuppressor::new().with_vocab_size(100_000);
let mut logits = vec![1.0f32; 100_000];
suppressor.apply(&mut logits);
assert!(logits[special_tokens::TIMESTAMP_BASE as usize] == f32::NEG_INFINITY);
}
#[test]
fn test_with_vocab_size_chain() {
let suppressor = WhisperTokenSuppressor::new()
.with_vocab_size(32000)
.with_timestamp_suppression(false);
let mut logits = vec![1.0f32; 51865];
suppressor.apply(&mut logits);
assert!(logits[special_tokens::TIMESTAMP_BASE as usize].is_finite());
}
#[test]
fn test_add_suppression_deduplication() {
let mut suppressor = WhisperTokenSuppressor::with_tokens(vec![100]);
suppressor.add_suppression(100); suppressor.add_suppression(200); assert_eq!(
suppressor
.suppressed_tokens()
.iter()
.filter(|&&t| t == 100)
.count(),
1
);
assert!(suppressor.suppressed_tokens().contains(&200));
}
#[test]
fn test_with_vocab_size_sets_field_directly() {
let suppressor = WhisperTokenSuppressor::new().with_vocab_size(32000);
let mut logits = vec![1.0f32; 51865];
suppressor.apply(&mut logits);
assert!(
logits[32000].is_finite(),
"token at 32000 should not be suppressed when n_vocab=32000"
);
}
#[test]
fn test_with_vocab_size_zero() {
let suppressor = WhisperTokenSuppressor::new().with_vocab_size(0);
let mut logits = vec![1.0f32; 51865];
suppressor.apply(&mut logits);
assert!(
logits[special_tokens::TIMESTAMP_BASE as usize].is_finite(),
"timestamp base should not be suppressed when n_vocab=0"
);
}
}