Skip to main content

mean_pool_batch

Function mean_pool_batch 

Source
pub fn mean_pool_batch(
    hidden_states: &[f32],
    masks: &[&[i32]],
    max_seq_len: usize,
    dim: usize,
) -> Vec<Vec<f32>>
Expand description

Batched mean pooling over multiple sequences in a single output tensor.

The model output is [batch, max_seq_len, dim] flattened row-major. Each sequence is mean-pooled using its own attention mask to exclude padding.