use crate::decode::ctc::SPACE_IDX;
use crate::tensor::Tensor;
pub struct IntraWordPooler {
pub layers: Vec<(Tensor, Tensor, Tensor, Tensor)>,
}
pub fn build_intra_word_pooler(
dim: usize,
n_layers: usize,
store: Option<&crate::weights::WeightStore>,
) -> IntraWordPooler {
let mut layers = Vec::new();
if let Some(store) = store {
for i in 0..n_layers {
let linear_idx = i * 3;
let ln_idx = linear_idx + 1;
let p = "word_segmenter.intra_word_pooler.mlp.";
if let (Ok(w), Ok(b), Ok(lnw), Ok(lnb)) = (
crate::weights::get(store, &format!("{p}{linear_idx}.weight")),
crate::weights::get(store, &format!("{p}{linear_idx}.bias")),
crate::weights::get(store, &format!("{p}{ln_idx}.weight")),
crate::weights::get(store, &format!("{p}{ln_idx}.bias")),
) {
layers.push((
crate::weights::param_to_tensor(w),
crate::weights::param_to_tensor(b),
crate::weights::param_to_tensor(lnw),
crate::weights::param_to_tensor(lnb),
));
}
}
}
if layers.is_empty() {
for li in 0..n_layers {
let in_d = if li == 0 { dim } else { dim };
let out_d = dim;
layers.push((
identity_linear(in_d, out_d),
Tensor::zeros(&[out_d]),
Tensor::from_vec(vec![1.0; out_d], vec![out_d]),
Tensor::zeros(&[out_d]),
));
}
}
IntraWordPooler { layers }
}
fn identity_linear(in_d: usize, out_d: usize) -> Tensor {
let mut data = vec![0.0f32; in_d * out_d];
for i in 0..in_d.min(out_d) {
data[i * out_d + i] = 1.0;
}
Tensor::from_vec(data, vec![out_d, in_d])
}
impl IntraWordPooler {
pub fn forward_frames(&self, frames: &Tensor) -> Tensor {
let (n, d) = (frames.shape[0], frames.shape[1]);
let mut x = frames.reshape(&[1, n, d]);
for (w, b, lnw, lnb) in &self.layers {
x = x.linear(w, Some(b)).layer_norm(lnw, lnb, 1e-5).gelu();
}
let x = x.reshape(&[n, d]);
let mut mean = vec![0.0f32; d];
for fi in 0..n {
for j in 0..d {
mean[j] += x.data[fi * d + j];
}
}
for v in &mut mean {
*v /= n as f32;
}
Tensor::from_vec(mean, vec![d])
}
}
pub struct CTCSpaceSegmenter {
pub include_blanks: bool,
pub min_word_frames: usize,
pub pooler: IntraWordPooler,
}
impl CTCSpaceSegmenter {
pub fn forward(&self, z_final: &Tensor, ctc_logits: &Tensor) -> Vec<Tensor> {
assert_eq!(z_final.ndim(), 3);
let (b, t, d) = (z_final.shape[0], z_final.shape[1], z_final.shape[2]);
let c = ctc_logits.shape[2];
let mut results = Vec::with_capacity(b);
for bi in 0..b {
let mut preds = vec![0usize; t];
for ti in 0..t {
let base = (bi * t + ti) * c;
let mut best = 0usize;
let mut best_v = ctc_logits.data[base];
for j in 1..c {
if ctc_logits.data[base + j] > best_v {
best_v = ctc_logits.data[base + j];
best = j;
}
}
preds[ti] = best;
}
let mut segments: Vec<Vec<usize>> = Vec::new();
let mut current: Vec<usize> = Vec::new();
for ti in 0..t {
if preds[ti] == SPACE_IDX {
if !current.is_empty() {
segments.push(current);
current = Vec::new();
}
} else if self.include_blanks || preds[ti] != 0 {
current.push(ti);
}
}
if !current.is_empty() {
segments.push(current);
}
segments.retain(|s| s.len() >= self.min_word_frames);
if segments.is_empty() {
let mut all_frames = Vec::new();
for ti in 0..t {
all_frames
.push(z_final.data[(bi * t + ti) * d..(bi * t + ti + 1) * d].to_vec());
}
let flat: Vec<f32> = all_frames.into_iter().flatten().collect();
let frames = Tensor::from_vec(flat, vec![t, d]);
results.push(self.pooler.forward_frames(&frames).reshape(&[1, d]));
} else {
let mut embeds = Vec::new();
for seg in segments {
let n = seg.len();
let mut flat = vec![0.0f32; n * d];
for (si, &ti) in seg.iter().enumerate() {
let src = (bi * t + ti) * d;
flat[si * d..(si + 1) * d].copy_from_slice(&z_final.data[src..src + d]);
}
let frames = Tensor::from_vec(flat, vec![n, d]);
embeds.push(self.pooler.forward_frames(&frames));
}
let n_words = embeds.len();
let flat: Vec<f32> = embeds.into_iter().flat_map(|t| t.data).collect();
results.push(Tensor::from_vec(flat, vec![n_words, d]));
}
}
results
}
}