use super::model::Qwen35Model;
use crate::error::InferenceError;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum HiddenPooling {
#[default]
LastToken,
Mean,
}
impl Qwen35Model {
pub fn embed_tokens(
&self,
tokens: &[u32],
pooling: HiddenPooling,
) -> Result<Vec<f32>, InferenceError> {
let hiddens = self.final_hidden_states(tokens)?;
if hiddens.is_empty() {
return Err(InferenceError::Inference(
"embed_tokens: no hidden states were produced".to_string(),
));
}
Ok(pool_hidden(&hiddens, pooling))
}
pub fn hidden_size(&self) -> usize {
self.config.hidden_size
}
}
fn pool_hidden(hiddens: &[Vec<f32>], pooling: HiddenPooling) -> Vec<f32> {
debug_assert!(
!hiddens.is_empty(),
"pool_hidden requires a non-empty slice"
);
let dim = hiddens[0].len();
match pooling {
HiddenPooling::LastToken => hiddens[hiddens.len() - 1].clone(),
HiddenPooling::Mean => {
let n = hiddens.len() as f32;
let mut acc = vec![0.0_f32; dim];
for h in hiddens {
for (a, &x) in acc.iter_mut().zip(h.iter()) {
*a += x;
}
}
for a in &mut acc {
*a /= n;
}
acc
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::qwen35::test_support::tiny_zero_model;
fn discriminating_input() -> Vec<Vec<f32>> {
vec![vec![1.0, 3.0], vec![2.0, 3.0], vec![9.0, 9.0]]
}
#[test]
fn last_token_pooling_takes_the_final_position() {
let h = discriminating_input();
assert_eq!(pool_hidden(&h, HiddenPooling::LastToken), vec![9.0, 9.0]);
}
#[test]
fn mean_pooling_averages_across_positions() {
let h = discriminating_input();
assert_eq!(pool_hidden(&h, HiddenPooling::Mean), vec![4.0, 5.0]);
}
#[test]
fn the_two_modes_disagree_on_this_input() {
let h = discriminating_input();
assert_ne!(
pool_hidden(&h, HiddenPooling::LastToken),
pool_hidden(&h, HiddenPooling::Mean)
);
}
#[test]
fn single_position_makes_both_modes_agree() {
let h = vec![vec![7.0, -2.0]];
assert_eq!(
pool_hidden(&h, HiddenPooling::LastToken),
pool_hidden(&h, HiddenPooling::Mean)
);
}
#[test]
fn default_pooling_is_last_token() {
assert_eq!(HiddenPooling::default(), HiddenPooling::LastToken);
}
#[test]
fn empty_input_is_rejected_rather_than_pooled_to_zeros() {
let model = tiny_zero_model();
let err = model
.embed_tokens(&[], HiddenPooling::default())
.expect_err("empty token slice must not produce a vector");
assert!(
err.to_string().contains("at least 1 token"),
"unexpected error: {err}"
);
}
#[test]
fn out_of_vocab_token_is_rejected() {
let model = tiny_zero_model();
let bad = model.config.vocab_size as u32;
let err = model
.embed_tokens(&[bad], HiddenPooling::default())
.expect_err("out-of-vocab id must be rejected");
assert!(
err.to_string().contains("vocab_size"),
"unexpected error: {err}"
);
}
#[test]
fn embedding_length_is_hidden_size() {
let model = tiny_zero_model();
let v = model
.embed_tokens(&[0, 1], HiddenPooling::default())
.expect("tiny model should embed two in-vocab tokens");
assert_eq!(v.len(), model.hidden_size());
}
}