use crate::inference::models::bert::config::PoolingType;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PoolingStrategy {
Mean,
ClsToken,
EosLastToken,
}
impl PoolingStrategy {
pub fn from_pooling_type(pt: PoolingType) -> Option<Self> {
match pt {
PoolingType::Mean => Some(Self::Mean),
PoolingType::Cls => Some(Self::ClsToken),
PoolingType::Last => Some(Self::EosLastToken),
PoolingType::None | PoolingType::Rank => None,
}
}
pub fn as_str(self) -> &'static str {
match self {
Self::Mean => "mean",
Self::ClsToken => "cls",
Self::EosLastToken => "eos_last_token",
}
}
}
pub fn pool_mean(
hidden_states: &[f32],
seq_len: usize,
hidden: usize,
valid_token_count: usize,
) -> Vec<f32> {
assert!(hidden > 0, "pool_mean: hidden must be > 0");
assert!(
valid_token_count > 0,
"pool_mean: valid_token_count must be > 0"
);
assert!(
valid_token_count <= seq_len,
"pool_mean: valid_token_count ({}) > seq_len ({})",
valid_token_count,
seq_len
);
assert_eq!(
hidden_states.len(),
seq_len * hidden,
"pool_mean: hidden_states length {} != seq_len {} × hidden {}",
hidden_states.len(),
seq_len,
hidden
);
let mut out = vec![0.0f32; hidden];
for row in 0..valid_token_count {
let base = row * hidden;
for (d, v) in out.iter_mut().enumerate() {
*v += hidden_states[base + d];
}
}
let inv = 1.0 / valid_token_count as f32;
for v in &mut out {
*v *= inv;
}
out
}
pub fn pool_eos_last_token(
hidden_states: &[f32],
seq_len: usize,
hidden: usize,
valid_token_count: usize,
) -> Vec<f32> {
assert!(hidden > 0, "pool_eos_last_token: hidden must be > 0");
assert!(
valid_token_count > 0,
"pool_eos_last_token: valid_token_count must be > 0"
);
assert!(
valid_token_count <= seq_len,
"pool_eos_last_token: valid_token_count ({}) > seq_len ({})",
valid_token_count,
seq_len
);
assert_eq!(
hidden_states.len(),
seq_len * hidden,
"pool_eos_last_token: hidden_states length {} != seq_len {} × hidden {}",
hidden_states.len(),
seq_len,
hidden
);
let last_row = valid_token_count - 1;
let base = last_row * hidden;
hidden_states[base..base + hidden].to_vec()
}
pub fn l2_normalize(v: &mut [f32], eps: f32) {
let norm_sq: f32 = v.iter().map(|x| x * x).sum();
let inv_norm = 1.0 / (norm_sq + eps).sqrt();
for x in v.iter_mut() {
*x *= inv_norm;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn pool_mean_all_tokens_valid() {
let hs = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
let out = pool_mean(&hs, 3, 2, 3);
assert_eq!(out.len(), 2);
assert!((out[0] - 3.0).abs() < 1e-6, "dim 0: got {}", out[0]);
assert!((out[1] - 4.0).abs() < 1e-6, "dim 1: got {}", out[1]);
}
#[test]
fn pool_mean_masks_padding_rows() {
let hs = vec![1.0f32, 2.0, 3.0, 4.0, 99.0, 99.0];
let out = pool_mean(&hs, 3, 2, 2);
assert_eq!(out.len(), 2);
assert!((out[0] - 2.0).abs() < 1e-6, "dim 0: got {}", out[0]);
assert!((out[1] - 3.0).abs() < 1e-6, "dim 1: got {}", out[1]);
}
#[test]
fn pool_mean_single_token() {
let hs = vec![7.0f32, -3.5, 0.0, 2.0];
let out = pool_mean(&hs, 1, 4, 1);
assert_eq!(out, vec![7.0, -3.5, 0.0, 2.0]);
}
#[test]
fn pool_eos_last_picks_last_valid_row() {
let hs = vec![0.1f32, 0.2, 0.3, 0.4, 0.5, 0.6];
let out = pool_eos_last_token(&hs, 3, 2, 3);
assert_eq!(out.len(), 2);
assert!((out[0] - 0.5).abs() < 1e-7, "dim 0: got {}", out[0]);
assert!((out[1] - 0.6).abs() < 1e-7, "dim 1: got {}", out[1]);
}
#[test]
fn pool_eos_last_ignores_padding() {
let hs = vec![0.1f32, 0.2, 0.5, 0.6, 99.0, 99.0];
let out = pool_eos_last_token(&hs, 3, 2, 2);
assert_eq!(out.len(), 2);
assert!((out[0] - 0.5).abs() < 1e-7, "dim 0: got {}", out[0]);
assert!((out[1] - 0.6).abs() < 1e-7, "dim 1: got {}", out[1]);
}
#[test]
fn pool_eos_last_single_token() {
let hs = vec![3.0f32, -1.0, 2.5];
let out = pool_eos_last_token(&hs, 1, 3, 1);
assert_eq!(out, vec![3.0, -1.0, 2.5]);
}
#[test]
fn l2_normalize_already_unit() {
let mut v = vec![1.0f32, 0.0, 0.0];
l2_normalize(&mut v, 1e-12);
assert!((v[0] - 1.0).abs() < 1e-6);
assert!(v[1].abs() < 1e-6);
assert!(v[2].abs() < 1e-6);
}
#[test]
fn l2_normalize_three_four() {
let mut v = vec![3.0f32, 4.0];
l2_normalize(&mut v, 1e-12);
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-6, "norm = {}", norm);
assert!((v[0] - 0.6).abs() < 1e-6, "v[0] = {}", v[0]);
assert!((v[1] - 0.8).abs() < 1e-6, "v[1] = {}", v[1]);
}
#[test]
fn l2_normalize_all_zeros_no_panic() {
let mut v = vec![0.0f32; 8];
l2_normalize(&mut v, 1e-12);
for x in &v {
assert!(x.is_finite(), "non-finite after normalizing zeros");
}
}
#[test]
fn pooling_strategy_from_pooling_type_round_trips() {
use crate::inference::models::bert::config::PoolingType;
assert_eq!(
PoolingStrategy::from_pooling_type(PoolingType::Mean),
Some(PoolingStrategy::Mean)
);
assert_eq!(
PoolingStrategy::from_pooling_type(PoolingType::Cls),
Some(PoolingStrategy::ClsToken)
);
assert_eq!(
PoolingStrategy::from_pooling_type(PoolingType::Last),
Some(PoolingStrategy::EosLastToken)
);
assert_eq!(PoolingStrategy::from_pooling_type(PoolingType::None), None);
assert_eq!(PoolingStrategy::from_pooling_type(PoolingType::Rank), None);
}
#[test]
fn pooling_strategy_as_str_stable() {
assert_eq!(PoolingStrategy::Mean.as_str(), "mean");
assert_eq!(PoolingStrategy::ClsToken.as_str(), "cls");
assert_eq!(PoolingStrategy::EosLastToken.as_str(), "eos_last_token");
}
#[test]
fn pool_mean_and_eos_last_agree_on_single_token() {
let hs = vec![1.0f32, 2.0, 3.0, 4.0];
let mean_out = pool_mean(&hs, 1, 4, 1);
let last_out = pool_eos_last_token(&hs, 1, 4, 1);
assert_eq!(mean_out, last_out);
}
}