use unicode_categories::UnicodeCategories;
use crate::{
MultiArray,
audio::whisper::{
backend::{AlignmentMatrix, AlignmentView},
constants::{SAMPLE_RATE, SECONDS_PER_TIME_TOKEN},
decode::AlignmentRows,
error::{
AlignmentPitchUnavailable, AlignmentPitchUnexpectedLayout, InvalidAlignmentShape,
SegmentError,
},
options::{AlignmentGather, DecodingOptions, WordGrouping},
result::{DecodingResult, TranscriptionSegment, WordTiming},
tokenizer::WhisperTokenizer,
},
};
pub(crate) fn window_span(seek: usize, samples: usize) -> (f32, f32) {
let start = seek as f32 / SAMPLE_RATE as f32;
let end = start + samples as f32 / SAMPLE_RATE as f32;
if samples > 0 && end <= start {
return (start, start.next_up());
}
(start, end)
}
pub(crate) fn shift_span(start: f32, end: f32, offset_seconds: f32) -> (f32, f32) {
let shifted_start = start + offset_seconds;
let shifted_end = end + offset_seconds;
if end > start && shifted_end <= shifted_start {
return (shifted_start, shifted_start.next_up());
}
(shifted_start, shifted_end)
}
pub fn find_seek_point_and_segments(
decoding: &DecodingResult,
options: &DecodingOptions,
all_segments_count: usize,
current_seek: usize,
segment_size: usize,
tokenizer: &WhisperTokenizer,
) -> Result<(usize, Option<Vec<TranscriptionSegment>>), SegmentError> {
let special = tokenizer.special_tokens();
let time_token = special.time_token_begin();
let special_token_begin = special.special_token_begin();
let mut seek = current_seek;
let (time_offset, window_end) = window_span(current_seek, segment_size);
if let Some(threshold) = options.no_speech_threshold() {
let mut should_skip = decoding.no_speech_prob() > threshold;
if let Some(logprob_threshold) = options.logprob_threshold()
&& decoding.avg_logprob() > logprob_threshold
{
should_skip = false;
}
if should_skip {
return Ok((seek + segment_size, None));
}
}
let current_tokens = decoding.tokens_slice();
let current_log_probs = decoding.token_log_probs_slice();
let is_timestamp_token: Vec<bool> = current_tokens.iter().map(|&t| t >= time_token).collect();
let single_timestamp_ending = matches!(is_timestamp_token.as_slice(), [.., false, true, false]);
let no_timestamp_ending = matches!(is_timestamp_token.as_slice(), [.., false, false, false]);
let mut slice_indexes: Vec<usize> = Vec::new();
let mut previous_is_timestamp = false;
for (index, &is_timestamp) in is_timestamp_token.iter().enumerate() {
if previous_is_timestamp && is_timestamp {
slice_indexes.push(index);
}
previous_is_timestamp = is_timestamp;
}
let mut segments: Vec<TranscriptionSegment> = Vec::new();
if slice_indexes.is_empty() {
let mut segment_end = window_end;
let timestamp_tokens: Vec<u32> = current_tokens
.iter()
.copied()
.filter(|&t| t > time_token)
.collect();
if let Some(&last_timestamp) = timestamp_tokens.last() {
segment_end = time_offset + (last_timestamp - time_token) as f32 * SECONDS_PER_TIME_TOKEN;
}
let word_tokens: Vec<u32> = current_tokens
.iter()
.copied()
.filter(|&t| t < special_token_begin)
.collect();
let segment_text_tokens: &[u32] = if options.skip_special_tokens() {
&word_tokens
} else {
current_tokens
};
let segment_text = tokenizer.decode(segment_text_tokens, false)?;
segments.push(
TranscriptionSegment::new()
.with_id(all_segments_count + segments.len())
.with_seek(seek)
.with_start(time_offset)
.with_end(segment_end)
.with_text(segment_text)
.with_tokens(current_tokens)
.with_token_log_probs(current_log_probs)
.with_temperature(decoding.temperature())
.with_avg_logprob(decoding.avg_logprob())
.with_compression_ratio(decoding.compression_ratio())
.with_no_speech_prob(decoding.no_speech_prob()),
);
seek += segment_size;
} else {
if single_timestamp_ending {
let single_ending_index = is_timestamp_token
.iter()
.rposition(|&t| t)
.expect("single_timestamp_ending's pattern requires a `true` entry");
slice_indexes.push(single_ending_index + 1);
} else if no_timestamp_ending {
slice_indexes.push(current_tokens.len());
}
let mut last_slice_start = 0usize;
for ¤t_slice_end in &slice_indexes {
let sliced_tokens = ¤t_tokens[last_slice_start..current_slice_end];
let sliced_log_probs = ¤t_log_probs[last_slice_start..current_slice_end];
let timestamp_tokens: Vec<u32> = sliced_tokens
.iter()
.copied()
.filter(|&t| t >= time_token)
.collect();
let start_ts = *timestamp_tokens
.first()
.expect("slice bounded by a timestamp pair contains a timestamp token");
let end_ts = *timestamp_tokens
.last()
.expect("slice bounded by a timestamp pair contains a timestamp token");
let start_seconds = (start_ts - time_token) as f32 * SECONDS_PER_TIME_TOKEN;
let end_seconds = (end_ts - time_token) as f32 * SECONDS_PER_TIME_TOKEN;
let word_tokens: Vec<u32> = sliced_tokens
.iter()
.copied()
.filter(|&t| t < special_token_begin)
.collect();
let sliced_text_tokens: &[u32] = if options.skip_special_tokens() {
&word_tokens
} else {
sliced_tokens
};
let slice_text = tokenizer.decode(sliced_text_tokens, false)?;
segments.push(
TranscriptionSegment::new()
.with_id(all_segments_count + segments.len())
.with_seek(seek)
.with_start(time_offset + start_seconds)
.with_end(time_offset + end_seconds)
.with_text(slice_text)
.with_tokens(sliced_tokens)
.with_token_log_probs(sliced_log_probs)
.with_temperature(decoding.temperature())
.with_avg_logprob(decoding.avg_logprob())
.with_compression_ratio(decoding.compression_ratio())
.with_no_speech_prob(decoding.no_speech_prob()),
);
last_slice_start = current_slice_end;
}
if no_timestamp_ending {
seek += segment_size;
} else {
let last_index = last_slice_start - usize::from(single_timestamp_ending);
let last_timestamp_token = current_tokens[last_index] - time_token;
let last_timestamp_seconds = last_timestamp_token as f32 * SECONDS_PER_TIME_TOKEN;
seek += (last_timestamp_seconds * SAMPLE_RATE as f32) as usize;
}
}
Ok((seek, Some(segments)))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DtwPath {
text_indices: Vec<isize>,
time_indices: Vec<isize>,
}
impl DtwPath {
#[inline(always)]
pub fn text_indices_slice(&self) -> &[isize] {
self.text_indices.as_slice()
}
#[inline(always)]
pub fn time_indices_slice(&self) -> &[isize] {
self.time_indices.as_slice()
}
}
fn min_cost_and_trace(diagonal: f64, up: f64, left: f64, value: f64) -> (f64, i8) {
let c0 = diagonal + value;
let c1 = up + value;
let c2 = left + value;
if c0 < c1 && c0 < c2 {
(c0, 0)
} else if c1 < c0 && c1 < c2 {
(c1, 1)
} else {
(c2, 2)
}
}
pub fn dynamic_time_warping(matrix: &AlignmentView<'_>) -> Result<DtwPath, SegmentError> {
let (rows, cols) = (matrix.rows(), matrix.cols());
if rows == 0 || cols == 0 {
return Err(SegmentError::InvalidAlignmentShape(
InvalidAlignmentShape::new(rows, cols, matrix.data().len()),
));
}
let width = cols + 1;
let mut cost = vec![f64::INFINITY; (rows + 1) * width];
let mut trace = vec![-1i8; (rows + 1) * width];
cost[0] = 0.0;
for cell in &mut trace[1..=cols] {
*cell = 2; }
for i in 1..=rows {
trace[i * width] = 1; }
for row in 1..=rows {
for column in 1..=cols {
let value = -f64::from(matrix.row(row - 1)[column - 1]);
let diagonal = cost[(row - 1) * width + column - 1];
let up = cost[(row - 1) * width + column];
let left = cost[row * width + column - 1];
let (best, direction) = min_cost_and_trace(diagonal, up, left, value);
cost[row * width + column] = best;
trace[row * width + column] = direction;
}
}
let (mut i, mut j) = (rows, cols);
let mut text_indices = Vec::new();
let mut time_indices = Vec::new();
while i > 0 || j > 0 {
text_indices.push(i as isize - 1);
time_indices.push(j as isize - 1);
match trace[i * width + j] {
0 => {
i -= 1;
j -= 1;
}
1 => i -= 1,
2 => j -= 1,
_ => break,
}
}
text_indices.reverse();
time_indices.reverse();
Ok(DtwPath {
text_indices,
time_indices,
})
}
fn swift_contains(haystack: &str, needle: &str) -> bool {
!needle.is_empty() && haystack.contains(needle)
}
fn trim_swift_whitespaces(s: &str) -> &str {
s.trim_matches(|c: char| c.is_separator_space() || c == '\u{0009}')
}
pub fn merge_punctuations(
alignment: &[WordTiming],
prepended: &str,
appended: &str,
) -> Vec<WordTiming> {
if alignment.is_empty() {
return Vec::new();
}
let mut prepended_alignment: Vec<WordTiming> = Vec::new();
if !swift_contains(prepended, trim_swift_whitespaces(alignment[0].word())) {
prepended_alignment.push(alignment[0].clone());
}
for pair in alignment.windows(2) {
let previous = &pair[0];
let current = &pair[1];
let previous_starts_with_whitespace = previous
.word()
.chars()
.next()
.is_some_and(|c| c.is_separator_space() || c == '\u{0009}');
if previous_starts_with_whitespace
&& swift_contains(prepended, trim_swift_whitespaces(previous.word()))
{
let mut word = previous.word().to_string();
word.push_str(current.word());
let mut tokens = previous.tokens_slice().to_vec();
tokens.extend_from_slice(current.tokens_slice());
let merged = WordTiming::new(
word,
tokens,
current.start(),
current.end(),
current.probability(),
);
if prepended_alignment.is_empty() {
prepended_alignment.push(merged);
} else {
let last = prepended_alignment.len() - 1;
prepended_alignment[last] = merged;
}
} else {
prepended_alignment.push(current.clone());
}
}
let mut appended_alignment: Vec<WordTiming> = Vec::new();
if let Some(first) = prepended_alignment.first() {
appended_alignment.push(first.clone());
}
for pair in prepended_alignment.windows(2) {
let previous = &pair[0];
let current = &pair[1];
if !previous.word().ends_with(' ')
&& swift_contains(appended, trim_swift_whitespaces(current.word()))
{
let mut word = previous.word().to_string();
word.push_str(current.word());
let mut tokens = previous.tokens_slice().to_vec();
tokens.extend_from_slice(current.tokens_slice());
let merged = WordTiming::new(
word,
tokens,
previous.start(),
previous.end(),
previous.probability(),
);
let last = appended_alignment.len() - 1;
appended_alignment[last] = merged;
} else {
appended_alignment.push(current.clone());
}
}
appended_alignment
.into_iter()
.filter(|w| {
!w.word().is_empty()
&& !swift_contains(appended, w.word())
&& !swift_contains(prepended, w.word())
})
.collect()
}
pub fn find_alignment(
word_token_ids: &[u32],
alignment: &AlignmentView<'_>,
token_log_probs: &[f32],
tokenizer: &WhisperTokenizer,
language_code: &str,
grouping: WordGrouping,
) -> Result<Vec<WordTiming>, SegmentError> {
Ok(
find_alignment_spanned(
word_token_ids,
alignment,
token_log_probs,
tokenizer,
language_code,
grouping,
)?
.into_iter()
.map(|(word, _)| word)
.collect(),
)
}
pub(crate) type SpannedWord = (WordTiming, (usize, usize));
pub(crate) fn find_alignment_spanned(
word_token_ids: &[u32],
alignment: &AlignmentView<'_>,
token_log_probs: &[f32],
tokenizer: &WhisperTokenizer,
language_code: &str,
grouping: WordGrouping,
) -> Result<Vec<SpannedWord>, SegmentError> {
Ok(
find_alignment_timed(
word_token_ids,
alignment,
token_log_probs,
tokenizer,
language_code,
grouping,
)?
.words,
)
}
pub(crate) struct TimedAlignment {
words: Vec<SpannedWord>,
starts: Vec<f32>,
ends: Vec<f32>,
}
fn find_alignment_timed(
word_token_ids: &[u32],
alignment: &AlignmentView<'_>,
token_log_probs: &[f32],
tokenizer: &WhisperTokenizer,
language_code: &str,
grouping: WordGrouping,
) -> Result<TimedAlignment, SegmentError> {
let path = dynamic_time_warping(alignment)?;
let text_indices = path.text_indices_slice();
let time_indices = path.time_indices_slice();
let word_tokens = tokenizer.split_to_word_tokens(word_token_ids, language_code, grouping)?;
if word_tokens.len() <= 1 {
return Ok(TimedAlignment {
words: Vec::new(),
starts: Vec::new(),
ends: Vec::new(),
});
}
let mut start_times: Vec<f32> = vec![0.0];
let mut end_times: Vec<f32> = Vec::new();
let mut current_text_index = text_indices.first().copied().unwrap_or(0);
for (index, &text_index) in text_indices.iter().enumerate() {
if text_index != current_text_index {
current_text_index = text_index;
let time = time_indices[index] as f32 * SECONDS_PER_TIME_TOKEN;
start_times.push(time);
end_times.push(time);
}
}
end_times.push(time_indices.last().copied().unwrap_or(1500) as f32 * SECONDS_PER_TIME_TOKEN);
let mut word_timings: Vec<SpannedWord> = Vec::with_capacity(word_tokens.len());
let mut current_token_index = 0usize;
for (word, tokens) in word_tokens {
let start_index = current_token_index;
let word_start_time = start_times[current_token_index];
current_token_index += tokens.len() - 1;
let word_end_time = end_times[current_token_index];
current_token_index += 1;
let probs = &token_log_probs[start_index..current_token_index];
let mean_log_prob = probs.iter().sum::<f32>() / probs.len() as f32;
word_timings.push((
WordTiming::new(
word,
tokens,
word_start_time,
word_end_time,
mean_log_prob.exp(),
),
(start_index, current_token_index),
));
}
Ok(TimedAlignment {
words: word_timings,
starts: start_times,
ends: end_times,
})
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct WordDurationConstraints {
median: f32,
max: f32,
}
impl WordDurationConstraints {
#[inline(always)]
pub const fn median(&self) -> f32 {
self.median
}
#[inline(always)]
pub const fn max_duration(&self) -> f32 {
self.max
}
}
pub fn calculate_word_duration_constraints(alignment: &[WordTiming]) -> WordDurationConstraints {
let mut durations: Vec<f32> = alignment
.iter()
.map(WordTiming::duration)
.filter(|&duration| duration > 0.0)
.collect();
durations.sort_by(f32::total_cmp);
let raw_median = durations.get(durations.len() / 2).copied().unwrap_or(0.0);
let median = raw_median.min(0.7);
let max = median * 2.0;
WordDurationConstraints { median, max }
}
const SENTENCE_END_MARKS: [&str; 6] = [".", "。", "!", "!", "?", "?"];
pub fn truncate_long_words_at_sentence_boundaries(
mut alignment: Vec<WordTiming>,
max_duration: f32,
) -> Vec<WordTiming> {
for i in 1..alignment.len() {
if alignment[i].duration() > max_duration {
if SENTENCE_END_MARKS.contains(&alignment[i].word()) {
let start = alignment[i].start();
alignment[i].set_end(start + max_duration);
} else if SENTENCE_END_MARKS.contains(&alignment[i - 1].word()) {
let end = alignment[i].end();
alignment[i].set_start(end - max_duration);
}
}
}
alignment
}
pub fn update_segments_with_word_timings(
segments: &[TranscriptionSegment],
merged_alignment: &[WordTiming],
seek: usize,
last_speech_timestamp: f32,
constrained_median_duration: f32,
max_duration: f32,
tokenizer: &WhisperTokenizer,
) -> Result<Vec<TranscriptionSegment>, SegmentError> {
let time_offset = seek as f32 / SAMPLE_RATE as f32;
let mut word_index = 0usize;
let mut last_speech_timestamp = last_speech_timestamp;
let mut updated_segments: Vec<TranscriptionSegment> = Vec::with_capacity(segments.len());
for segment in segments {
let (updated_segment, consumed) = update_segment_with_word_timings(
segment,
&merged_alignment[word_index.min(merged_alignment.len())..],
updated_segments.last().map(TranscriptionSegment::end),
time_offset,
&mut last_speech_timestamp,
constrained_median_duration,
max_duration,
tokenizer,
)?;
word_index += consumed;
updated_segments.push(updated_segment);
}
Ok(updated_segments)
}
#[allow(clippy::too_many_arguments)] fn update_segment_with_word_timings(
segment: &TranscriptionSegment,
alignment: &[WordTiming],
previous_end: Option<f32>,
time_offset: f32,
last_speech_timestamp: &mut f32,
constrained_median_duration: f32,
max_duration: f32,
tokenizer: &WhisperTokenizer,
) -> Result<(TranscriptionSegment, usize), SegmentError> {
let special_begin = tokenizer.special_tokens().special_token_begin();
let mut saved_tokens = 0usize;
let text_token_count = segment
.tokens_slice()
.iter()
.filter(|&&token| token < special_begin)
.count();
let mut words_in_segment: Vec<WordTiming> = Vec::new();
let mut consumed = 0usize;
for timing in alignment {
if saved_tokens >= text_token_count {
break;
}
consumed += 1;
let timing_tokens: Vec<u32> = timing
.tokens_slice()
.iter()
.copied()
.filter(|&token| token < special_begin)
.collect();
if timing_tokens.is_empty() {
continue;
}
let timing_tokens_len = timing_tokens.len();
let word = if timing_tokens_len < timing.tokens_slice().len() {
tokenizer.decode(&timing_tokens, false)?
} else {
timing.word().to_string()
};
let mut start = rounded_to_places(time_offset + timing.start(), 2);
let end = rounded_to_places(time_offset + timing.end(), 2);
if end - start < constrained_median_duration / 4.0 {
if let Some(previous) = words_in_segment.last() {
let previous_end = previous.end();
if start > previous_end {
let space_available = start - previous_end;
let desired_duration = space_available.min(constrained_median_duration / 2.0);
start = rounded_to_places(start - desired_duration, 2);
}
} else if let Some(previous_end) = previous_end
&& start > previous_end
{
let space_available = start - previous_end;
let desired_duration = space_available.min(constrained_median_duration / 2.0);
start = rounded_to_places(start - desired_duration, 2);
}
}
let probability = rounded_to_places(timing.probability(), 2);
words_in_segment.push(WordTiming::new(
word,
timing_tokens,
start,
end,
probability,
));
saved_tokens += timing_tokens_len;
}
let mut updated_segment = segment.clone();
if !words_in_segment.is_empty() {
let pause_length = words_in_segment[0].end() - *last_speech_timestamp;
let first_word_too_long = words_in_segment[0].duration() > max_duration;
let both_words_too_long = words_in_segment.len() > 1
&& words_in_segment[1].end() - words_in_segment[0].start() > max_duration * 2.0;
if pause_length > constrained_median_duration * 4.0
&& (first_word_too_long || both_words_too_long)
{
if words_in_segment.len() > 1 && words_in_segment[1].duration() > max_duration {
let w1_end = words_in_segment[1].end();
let boundary = (w1_end / 2.0).max(w1_end - max_duration);
words_in_segment[0].set_end(boundary);
words_in_segment[1].set_start(boundary);
}
let w0_end = words_in_segment[0].end();
words_in_segment[0].set_start(last_speech_timestamp.max(w0_end - max_duration));
}
let w0_start = words_in_segment[0].start();
let w0_end = words_in_segment[0].end();
if segment.start() < w0_end && segment.start() - 0.5 > w0_start {
let clamped = (w0_end - constrained_median_duration)
.min(segment.start())
.max(0.0);
words_in_segment[0].set_start(clamped);
} else {
updated_segment.set_start(words_in_segment[0].start());
}
let last_index = words_in_segment.len() - 1;
let last_start = words_in_segment[last_index].start();
let last_end = words_in_segment[last_index].end();
if updated_segment.end() > last_start && segment.end() + 0.5 < last_end {
let clamped = (last_start + constrained_median_duration).max(segment.end());
words_in_segment[last_index].set_end(clamped);
} else {
updated_segment.set_end(last_end);
}
*last_speech_timestamp = updated_segment.end();
}
updated_segment.set_words(words_in_segment);
Ok((updated_segment, consumed))
}
pub(crate) fn rounded_to_places(value: f32, decimal_places: i32) -> f32 {
let divisor = 10f32.powi(decimal_places);
(value * divisor).round() / divisor
}
pub(crate) fn coreml_f16_row_pitch(rows: usize, cols: usize) -> Result<usize, SegmentError> {
let probe = MultiArray::f16_surface(&[rows, cols]).map_err(|source| {
SegmentError::AlignmentPitchUnavailable(AlignmentPitchUnavailable::new(rows, cols, source))
})?;
row_pitch_of(&probe, rows, cols)
}
fn row_pitch_of(probe: &MultiArray, rows: usize, cols: usize) -> Result<usize, SegmentError> {
let strides = probe.strides();
match *strides {
[pitch, 1] if pitch >= cols => Ok(pitch),
_ => Err(SegmentError::AlignmentPitchUnexpectedLayout(
AlignmentPitchUnexpectedLayout::new(rows, cols, strides.to_vec()),
)),
}
}
type RowPitchProbe<'a> = &'a dyn Fn(usize, usize) -> Result<usize, SegmentError>;
fn gather_swift_parity_into(
out: &mut [f32],
alignment: &AlignmentView<'_>,
needed: usize,
cols: usize,
swift_source_rows: usize,
row_pitch: RowPitchProbe<'_>,
) -> Result<(), SegmentError> {
let src_rows = alignment.rows().min(swift_source_rows);
let src_pitch = if src_rows == 0 {
cols
} else {
row_pitch(swift_source_rows, cols)?
};
let dst_pitch = row_pitch(needed, cols)?;
gather_swift_rows(
out,
alignment.data(),
src_rows,
needed,
cols,
src_pitch,
dst_pitch,
);
Ok(())
}
fn gather_swift_rows(
out: &mut [f32],
source: &[f32],
src_rows: usize,
needed: usize,
cols: usize,
src_pitch: usize,
dst_pitch: usize,
) {
debug_assert!(src_pitch >= cols && dst_pitch >= cols && cols > 0);
let copied = needed * cols;
for (row, out_row) in out.chunks_mut(cols).enumerate().take(needed) {
let mut offset = row * dst_pitch;
let end = (offset + cols).min(copied);
let mut column = 0usize;
while offset < end {
let source_row = offset / src_pitch;
if source_row >= src_rows {
break;
}
let source_column = offset % src_pitch;
let run = if source_column < cols {
let run = (cols - source_column).min(end - offset);
let start = source_row * cols + source_column;
out_row[column..column + run].copy_from_slice(&source[start..start + run]);
run
} else {
(src_pitch - source_column).min(end - offset)
};
offset += run;
column += run;
}
}
}
#[allow(clippy::too_many_arguments)] pub fn add_word_timestamps(
segments: &[TranscriptionSegment],
alignment: &AlignmentView<'_>,
tokenizer: &WhisperTokenizer,
language_code: &str,
grouping: WordGrouping,
gather: AlignmentGather,
swift_source_rows: usize,
seek: usize,
prepended: &str,
appended: &str,
last_speech_timestamp: f32,
) -> Result<Vec<TranscriptionSegment>, SegmentError> {
let aligned = aligned_window(
segments,
alignment,
tokenizer,
language_code,
grouping,
gather,
swift_source_rows,
)?;
visible_words(
segments,
aligned.into_iter().map(|(word, _)| word).collect(),
tokenizer,
seek,
prepended,
appended,
last_speech_timestamp,
)
}
#[allow(clippy::too_many_arguments)] pub(crate) fn derive_visible_words(
window: &[TranscriptionSegment],
kept: &[KeptSegment],
alignment: &AlignmentView<'_>,
tokenizer: &WhisperTokenizer,
language_code: &str,
grouping: WordGrouping,
gather: AlignmentGather,
swift_source_rows: usize,
seek: usize,
prepended: &str,
appended: &str,
last_speech_timestamp: f32,
) -> Result<Vec<TranscriptionSegment>, SegmentError> {
if kept.is_empty() {
return Ok(Vec::new());
}
let timed = aligned_window_timed(
window,
alignment,
tokenizer,
language_code,
grouping,
gather,
swift_source_rows,
)?;
let (durations, own_words) = kept_words(window, kept, &timed, |tokens| {
Ok(tokenizer.decode(tokens, false)?)
})?;
derive_per_segment(
kept.iter().map(|(segment, _, _)| segment),
own_words,
durations,
tokenizer,
seek,
prepended,
appended,
last_speech_timestamp,
)
}
fn kept_words<D>(
window: &[TranscriptionSegment],
kept: &[KeptSegment],
timed: &TimedAlignment,
decode: D,
) -> Result<(WordDurationConstraints, Vec<Vec<WordTiming>>), SegmentError>
where
D: Fn(&[u32]) -> Result<String, SegmentError>,
{
let words: Vec<WordTiming> = timed.words.iter().map(|(word, _)| word.clone()).collect();
let durations = calculate_word_duration_constraints(&words);
let logged: Vec<Option<f32>> = window
.iter()
.flat_map(|segment| {
let pairs = segment.token_log_probs_slice();
segment
.tokens_slice()
.iter()
.enumerate()
.map(move |(index, &token)| match pairs.get(index) {
Some(&(logged, log_prob)) if logged == token => Some(log_prob),
_ => None,
})
})
.collect();
let mut owner: Vec<Option<(usize, u32)>> = vec![None; logged.len()];
for (index, (segment, _, positions)) in kept.iter().enumerate() {
for (&position, &token) in positions.iter().zip(segment.tokens_slice()) {
if let Some(slot) = owner.get_mut(position) {
*slot = Some((index, token));
}
}
}
let mut surviving: Vec<(usize, WordTiming)> = Vec::new();
for (word, (from, to)) in &timed.words {
let mut parts: Vec<KeptPart> = Vec::new();
for (position, &aligned) in (*from..*to).zip(word.tokens_slice()) {
let Some((index, token)) = owner.get(position).copied().flatten() else {
continue;
};
let restated = token != aligned;
let sampled = if restated {
None
} else {
logged.get(position).copied().flatten()
};
match parts.last_mut() {
Some(part) if part.segment == index => {
part.tokens.push(token);
part.positions.push(position);
part.restated |= restated;
part.own.extend(sampled);
}
_ => parts.push(KeptPart {
segment: index,
tokens: vec![token],
positions: vec![position],
restated,
own: sampled.into_iter().collect(),
}),
}
}
for part in parts {
if !part.restated && part.tokens.as_slice() == word.tokens_slice() {
surviving.push((part.segment, word.clone()));
} else {
let text = decode(&part.tokens)?;
surviving.push((part.segment, word_part(word, text, &part, timed)));
}
}
}
Ok((durations, segment_words(surviving, kept.len(), durations)))
}
struct KeptPart {
segment: usize,
tokens: Vec<u32>,
positions: Vec<usize>,
restated: bool,
own: Vec<f32>,
}
fn word_part(
word: &WordTiming,
text: String,
part: &KeptPart,
timed: &TimedAlignment,
) -> WordTiming {
let start = part
.positions
.first()
.and_then(|&first| timed.starts.get(first).copied())
.unwrap_or(word.start());
let end = part
.positions
.last()
.and_then(|&last| timed.ends.get(last).copied())
.unwrap_or(word.end());
let probability = if !part.own.is_empty() {
(part.own.iter().sum::<f32>() / part.own.len() as f32).exp()
} else if part.restated {
0.0
} else {
word.probability()
};
WordTiming::new(text, part.tokens.clone(), start, end, probability)
}
fn segment_words(
surviving: Vec<(usize, WordTiming)>,
segments: usize,
durations: WordDurationConstraints,
) -> Vec<Vec<WordTiming>> {
let (owners, words): (Vec<usize>, Vec<WordTiming>) = surviving.into_iter().unzip();
let words = truncate_long_words_at_sentence_boundaries(words, durations.max_duration());
let mut own_words: Vec<Vec<WordTiming>> = vec![Vec::new(); segments];
for (owner, word) in owners.into_iter().zip(words) {
if let Some(own) = own_words.get_mut(owner) {
own.push(word);
}
}
own_words
}
#[allow(clippy::too_many_arguments)] fn derive_per_segment<'a>(
segments: impl IntoIterator<Item = &'a TranscriptionSegment>,
own_words: Vec<Vec<WordTiming>>,
durations: WordDurationConstraints,
tokenizer: &WhisperTokenizer,
seek: usize,
prepended: &str,
appended: &str,
last_speech_timestamp: f32,
) -> Result<Vec<TranscriptionSegment>, SegmentError> {
let time_offset = seek as f32 / SAMPLE_RATE as f32;
let mut last_speech_timestamp = last_speech_timestamp;
let mut derived: Vec<TranscriptionSegment> = Vec::with_capacity(own_words.len());
for (segment, mut words) in segments.into_iter().zip(own_words) {
if !words.is_empty() {
words = merge_punctuations(&words, prepended, appended);
}
let (mut updated, _) = update_segment_with_word_timings(
segment,
&words,
derived.last().map(TranscriptionSegment::end),
time_offset,
&mut last_speech_timestamp,
durations.median(),
durations.max_duration(),
tokenizer,
)?;
let (start, end) = (segment.start(), segment.end());
updated.set_start(start).set_end(end);
let hold = |time: f32| time.max(start).min(end);
for word in updated.words_slice_mut() {
let (from, to) = (word.start(), word.end());
word.set_start(hold(from)).set_end(hold(to));
}
if !updated.words_slice().is_empty() {
last_speech_timestamp = end;
}
derived.push(updated);
}
Ok(derived)
}
pub(crate) fn attribute_window(
segments: &[TranscriptionSegment],
alignment: &AlignmentView<'_>,
rows: &AlignmentRows,
tokenizer: &WhisperTokenizer,
language_code: &str,
) -> Vec<Vec<RawWord>> {
match text_units(segments, alignment, rows, tokenizer, language_code) {
Ok(units) => own_raw_words(segments, &units),
Err(_) => Vec::new(),
}
}
pub(crate) fn token_rows(
segments: &[TranscriptionSegment],
rows: &AlignmentRows,
) -> Vec<Option<usize>> {
let count = segments
.iter()
.map(|segment| segment.tokens_slice().len())
.sum();
(0..count).map(|index| rows.row_of(index)).collect()
}
pub(crate) type TextUnit = ((usize, usize), usize, usize);
fn text_units(
segments: &[TranscriptionSegment],
alignment: &AlignmentView<'_>,
rows: &AlignmentRows,
tokenizer: &WhisperTokenizer,
language_code: &str,
) -> Result<Vec<TextUnit>, SegmentError> {
let cols = alignment.cols();
if cols == 0 {
return Err(SegmentError::InvalidAlignmentShape(
InvalidAlignmentShape::new(alignment.rows(), cols, alignment.data().len()),
));
}
let token_rows = token_rows(segments, rows);
let mut tokens: Vec<u32> = Vec::new();
let mut positions: Vec<usize> = Vec::new();
let mut data: Vec<f32> = Vec::new();
let flattened = segments
.iter()
.flat_map(|segment| segment.tokens_slice().iter().copied());
for (position, token) in flattened.enumerate() {
if let Some(row) = token_rows[position]
&& row < alignment.rows()
{
tokens.push(token);
positions.push(position);
data.extend_from_slice(alignment.row(row));
}
}
let matrix = AlignmentMatrix::new(data, tokens.len(), cols);
let (starts, ends) = token_frames(&dynamic_time_warping(&matrix.view())?);
let special_begin = tokenizer.special_tokens().special_token_begin();
let unit = |from: usize, to: usize| -> TextUnit {
(
(positions[from], positions[to - 1] + 1),
starts[from],
ends[to - 1],
)
};
let mut units = Vec::new();
let mut at = 0usize;
for (_, group) in
tokenizer.split_to_word_tokens(&tokens, language_code, WordGrouping::FineGrained)?
{
let end = (at + group.len()).min(tokens.len());
let mut from: Option<usize> = None;
for index in at..end {
let text = tokens[index] < special_begin;
if let Some(start) = from
&& (!text || positions[index] != positions[index - 1] + 1)
{
units.push(unit(start, index));
from = None;
}
if text && from.is_none() {
from = Some(index);
}
}
if let Some(start) = from {
units.push(unit(start, end));
}
at = end;
}
Ok(units)
}
fn token_frames(path: &DtwPath) -> (Vec<usize>, Vec<usize>) {
let text_indices = path.text_indices_slice();
let time_indices = path.time_indices_slice();
let frame = |time: isize| usize::try_from(time).unwrap_or(0);
let mut starts = vec![0usize];
let mut ends = Vec::new();
let mut current = text_indices.first().copied().unwrap_or(0);
for (index, &text_index) in text_indices.iter().enumerate() {
if text_index != current {
current = text_index;
starts.push(frame(time_indices[index]));
ends.push(frame(time_indices[index]));
}
}
ends.push(time_indices.last().map_or(1500, |&time| frame(time)));
(starts, ends)
}
fn aligned_window(
segments: &[TranscriptionSegment],
alignment: &AlignmentView<'_>,
tokenizer: &WhisperTokenizer,
language_code: &str,
grouping: WordGrouping,
gather: AlignmentGather,
swift_source_rows: usize,
) -> Result<Vec<SpannedWord>, SegmentError> {
Ok(
aligned_window_timed(
segments,
alignment,
tokenizer,
language_code,
grouping,
gather,
swift_source_rows,
)?
.words,
)
}
fn aligned_window_timed(
segments: &[TranscriptionSegment],
alignment: &AlignmentView<'_>,
tokenizer: &WhisperTokenizer,
language_code: &str,
grouping: WordGrouping,
gather: AlignmentGather,
swift_source_rows: usize,
) -> Result<TimedAlignment, SegmentError> {
let mut word_token_ids: Vec<u32> = Vec::new();
let mut filtered_log_probs: Vec<f32> = Vec::new();
for segment in segments {
let log_probs = segment.token_log_probs_slice();
for (index, &token) in segment.tokens_slice().iter().enumerate() {
word_token_ids.push(token);
if let Some(&(logged_token, log_prob)) = log_probs.get(index)
&& logged_token == token
{
filtered_log_probs.push(log_prob);
}
}
}
let needed = word_token_ids.len();
let cols = alignment.cols();
if cols == 0 {
return Err(SegmentError::InvalidAlignmentShape(
InvalidAlignmentShape::new(alignment.rows(), cols, alignment.data().len()),
));
}
let mut data = vec![0.0f32; needed * cols];
if gather == AlignmentGather::SwiftParity && needed > 0 {
gather_swift_parity_into(
&mut data,
alignment,
needed,
cols,
swift_source_rows,
&coreml_f16_row_pitch,
)?;
} else {
for (row_index, row) in data
.chunks_mut(cols)
.enumerate()
.take(alignment.rows().min(needed))
{
row.copy_from_slice(alignment.row(row_index));
}
}
let filtered = AlignmentMatrix::new(data, needed, cols);
find_alignment_timed(
&word_token_ids,
&filtered.view(),
&filtered_log_probs,
tokenizer,
language_code,
grouping,
)
}
fn visible_words(
segments: &[TranscriptionSegment],
mut merged: Vec<WordTiming>,
tokenizer: &WhisperTokenizer,
seek: usize,
prepended: &str,
appended: &str,
last_speech_timestamp: f32,
) -> Result<Vec<TranscriptionSegment>, SegmentError> {
let word_durations = calculate_word_duration_constraints(&merged);
merged = truncate_long_words_at_sentence_boundaries(merged, word_durations.max_duration());
if !merged.is_empty() {
merged = merge_punctuations(&merged, prepended, appended);
}
update_segments_with_word_timings(
segments,
&merged,
seek,
last_speech_timestamp,
word_durations.median(),
word_durations.max_duration(),
tokenizer,
)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct RawWord {
span: (usize, usize),
start: usize,
end: usize,
}
impl RawWord {
pub(crate) const fn new(span: (usize, usize), start: usize, end: usize) -> Self {
Self { span, start, end }
}
pub(crate) const fn span(&self) -> (usize, usize) {
self.span
}
pub(crate) const fn start(&self) -> usize {
self.start
}
pub(crate) const fn end(&self) -> usize {
self.end
}
}
const SAMPLES_PER_FRAME: usize = SAMPLE_RATE as usize / 50;
pub(crate) type KeptSegment = (TranscriptionSegment, Vec<RawWord>, Vec<usize>);
pub(crate) fn own_raw_words(
segments: &[TranscriptionSegment],
units: &[TextUnit],
) -> Vec<Vec<RawWord>> {
let mut ranges = Vec::with_capacity(segments.len());
let mut at = 0usize;
for segment in segments {
let len = segment.tokens_slice().len();
ranges.push(at..at + len);
at += len;
}
let mut owned: Vec<Vec<RawWord>> = segments.iter().map(|_| Vec::new()).collect();
for &((from, to), start, end) in units {
let Some(index) = ranges.iter().position(|range| range.contains(&from)) else {
continue;
};
let offset = ranges[index].start;
owned[index].push(RawWord::new(
(from - offset, to.saturating_sub(offset)),
start * SAMPLES_PER_FRAME,
end * SAMPLES_PER_FRAME,
));
}
owned
}
#[cfg(test)]
mod tests;