use unicode_categories::UnicodeCategories;
use crate::{
MultiArray,
audio::whisper::{
backend::{AlignmentMatrix, AlignmentView},
constants::{SAMPLE_RATE, SECONDS_PER_TIME_TOKEN},
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> {
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(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<WordTiming> = 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(),
));
}
Ok(word_timings)
}
#[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 special_begin = tokenizer.special_tokens().special_token_begin();
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_index, segment) in segments.iter().enumerate() {
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();
for timing in &merged_alignment[word_index.min(merged_alignment.len())..] {
if saved_tokens >= text_token_count {
break;
}
word_index += 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 segment_index > 0
&& updated_segments.len() > segment_index - 1
&& start > updated_segments[segment_index - 1].end()
{
let previous_end = updated_segments[segment_index - 1].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);
updated_segments.push(updated_segment);
}
Ok(updated_segments)
}
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 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);
let mut merged = find_alignment(
&word_token_ids,
&filtered.view(),
&filtered_log_probs,
tokenizer,
language_code,
grouping,
)?;
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,
)
}
#[cfg(test)]
mod tests;