use std::{
cell::Cell,
sync::{
Mutex, PoisonError,
atomic::{AtomicBool, Ordering},
},
time::{Duration, Instant},
};
use unicode_categories::UnicodeCategories;
use crate::audio::whisper::{
audio::{self, chunker},
backend::{AlignmentMatrix, InferenceBackend, coreml::CoreMlBackend},
constants::{
APPEND_PUNCTUATION, BLANK_AUDIO_MARKER, DEFAULT_LANGUAGE_CODE, PREPEND_PUNCTUATION, SAMPLE_RATE,
},
decode::{
self, TranscriptionProgressCallback,
sampler::{self, GreedyTokenSampler},
},
error::{DecodeError, InvalidState, ModelError, TranscribeError, VadError},
model::{
ModelVariant, detect_variant,
manager::{ModelLoadTimings, ModelManager},
},
options::{DecodingOptions, Options},
result::{
DecodingResult, TranscriptionProgress, TranscriptionResult, TranscriptionSegment,
TranscriptionTimings, merge_transcription_results_with_options, needs_fallback,
},
segment,
stream::{AudioStreamTranscriber, agreement::LocalAgreementTranscriber},
task_facts::{SpanKnowledge, TaskFacts},
tokenizer::WhisperTokenizer,
};
#[cfg(test)]
mod tests;
fn trim_swift_whitespaces(s: &str) -> &str {
s.trim_matches(|c: char| c.is_separator_space() || c == '\u{0009}')
}
pub type SegmentDiscoveryCallback<'a> = &'a (dyn Fn(&[TranscriptionSegment]) + Sync);
pub struct TranscribeTask<'ctx, B> {
backend: &'ctx B,
tokenizer: &'ctx WhisperTokenizer,
segment_callback: Option<SegmentDiscoveryCallback<'ctx>>,
progress_callback: Option<TranscriptionProgressCallback<'ctx>>,
window_id_offset: usize,
facts_sink: Option<&'ctx Mutex<TaskFacts>>,
}
impl<'ctx, B> TranscribeTask<'ctx, B> {
pub const fn new(backend: &'ctx B, tokenizer: &'ctx WhisperTokenizer) -> Self {
Self {
backend,
tokenizer,
segment_callback: None,
progress_callback: None,
window_id_offset: 0,
facts_sink: None,
}
}
#[must_use]
#[inline(always)]
pub(crate) const fn with_facts_sink(mut self, sink: &'ctx Mutex<TaskFacts>) -> Self {
self.facts_sink = Some(sink);
self
}
#[must_use]
#[inline(always)]
pub const fn with_segment_callback(
mut self,
segment_callback: SegmentDiscoveryCallback<'ctx>,
) -> Self {
self.set_segment_callback(segment_callback);
self
}
#[inline(always)]
pub const fn set_segment_callback(
&mut self,
segment_callback: SegmentDiscoveryCallback<'ctx>,
) -> &mut Self {
self.segment_callback = Some(segment_callback);
self
}
#[must_use]
#[inline(always)]
pub const fn with_progress_callback(
mut self,
progress_callback: TranscriptionProgressCallback<'ctx>,
) -> Self {
self.set_progress_callback(progress_callback);
self
}
#[inline(always)]
pub const fn set_progress_callback(
&mut self,
progress_callback: TranscriptionProgressCallback<'ctx>,
) -> &mut Self {
self.progress_callback = Some(progress_callback);
self
}
#[must_use]
#[inline(always)]
pub const fn with_window_id_offset(mut self, window_id_offset: usize) -> Self {
self.set_window_id_offset(window_id_offset);
self
}
#[inline(always)]
pub const fn set_window_id_offset(&mut self, window_id_offset: usize) -> &mut Self {
self.window_id_offset = window_id_offset;
self
}
}
pub(crate) fn last_speech_timestamp_seed(previous_seek: usize) -> f32 {
(previous_seek as f64 / f64::from(SAMPLE_RATE)) as f32
}
impl<B> TranscribeTask<'_, B>
where
B: InferenceBackend,
{
pub fn run(
&self,
audio: &[f32],
options: &DecodingOptions,
) -> Result<TranscriptionResult, TranscribeError> {
let pipeline_start = Instant::now();
let mut timings = TranscriptionTimings::new();
let content_frames = audio.len();
let clip_start_seconds = options
.clip_timestamps_slice()
.first()
.copied()
.unwrap_or(0.0);
timings.set_input_audio_seconds(
content_frames as f64 / f64::from(SAMPLE_RATE) - f64::from(clip_start_seconds),
);
let mut all_segments: Vec<TranscriptionSegment> = Vec::new();
let mut decoded_segment_span = 0usize;
let mut allocated_ordinals = 0usize;
let mut detected_language: Option<String> = None;
let mut observed_language: Option<String> = None;
let mut language_observations: Vec<LanguageObservation> = Vec::new();
let local_facts = Mutex::new(TaskFacts::observed_clean());
let facts_sink: &Mutex<TaskFacts> = self.facts_sink.unwrap_or(&local_facts);
let decoder_init_start = Instant::now();
let mut state = self
.backend
.new_decoder_state()
.map_err(DecodeError::from)?;
timings.set_decoding_init(decoder_init_start.elapsed().as_secs_f64());
let mut initial_prompt: Vec<u32> =
vec![self.tokenizer.special_tokens().start_of_transcript_token()];
if options.use_prefill_prompt() {
initial_prompt = decode::prefill_tokens(
options,
self.tokenizer,
self.backend.dims().is_multilingual(),
);
}
let seek_clips = chunker::prepare_seek_clips(options.clip_timestamps_slice(), content_frames)?;
let window_padding = (options.window_clip_time() * SAMPLE_RATE as f32) as usize;
let window_samples = self.backend.dims().window_samples();
let mut window_index: u64 = 0;
let decode_loop_start = Instant::now();
for (seek_clip_start, seek_clip_end) in seek_clips {
if seek_clip_end <= window_padding {
continue;
}
let clip_guard = seek_clip_end - window_padding;
let mut seek = seek_clip_start;
while seek < clip_guard {
let segment_size = window_samples
.min(content_frames.saturating_sub(seek))
.min(seek_clip_end.saturating_sub(seek));
if segment_size == 0 {
break;
}
let audio_processing_start = Instant::now();
let padded = audio::pad_or_trim(&audio[seek..seek + segment_size], window_samples);
timings.set_audio_processing(
timings.audio_processing() + audio_processing_start.elapsed().as_secs_f64(),
);
timings.set_total_audio_processing_runs(timings.total_audio_processing_runs() + 1.0);
let logmel_start = Instant::now();
let features = self
.backend
.extract_features(&padded)
.map_err(DecodeError::from)?;
timings.set_logmels(timings.logmels() + logmel_start.elapsed().as_secs_f64());
timings.set_total_logmel_runs(timings.total_logmel_runs() + 1.0);
let encoder_start = Instant::now();
let encoder_output = self.backend.encode(&features).map_err(DecodeError::from)?;
timings.set_encoding(timings.encoding() + encoder_start.elapsed().as_secs_f64());
timings.set_total_encoding_runs(timings.total_encoding_runs() + 1.0);
let mut window_probe: Option<LanguageDetection> = None;
let (decoding_result, captured_alignment) = self.decode_with_fallback(
&encoder_output,
&mut state,
&mut initial_prompt,
&mut detected_language,
&mut observed_language,
&mut window_probe,
facts_sink,
options,
&mut timings,
window_index,
)?;
window_index += 1;
if let Some(detection) = window_probe {
let (start, end) = segment::window_span(seek, segment_size);
language_observations.push(LanguageObservation::new(start, end, detection));
}
let windowing_start = Instant::now();
let previous_seek = seek;
let (new_seek, mut current_segments) = segment::find_seek_point_and_segments(
&decoding_result,
options,
decoded_segment_span,
seek,
segment_size,
self.tokenizer,
)?;
let allocated_this_window = current_segments.as_ref().map_or(0, Vec::len);
seek = seek.max(new_seek);
if options.word_timestamps()
&& let Some(matrix) = &captured_alignment
{
let word_timestamps_start = Instant::now();
let language = detected_language
.as_deref()
.unwrap_or(DEFAULT_LANGUAGE_CODE);
let with_words = segment::add_word_timestamps(
current_segments.as_deref().unwrap_or(&[]), &matrix.view(),
self.tokenizer,
language,
options.word_grouping(), options.alignment_gather(), self.backend.dims().max_token_context(),
previous_seek,
PREPEND_PUNCTUATION,
APPEND_PUNCTUATION,
last_speech_timestamp_seed(previous_seek), )?;
timings.set_decoding_word_timestamps(
timings.decoding_word_timestamps() + word_timestamps_start.elapsed().as_secs_f64(),
);
timings
.set_total_timestamp_alignment_runs(timings.total_timestamp_alignment_runs() + 1.0);
let filtered: Vec<TranscriptionSegment> = with_words
.into_iter()
.filter(|segment| segment.end() > segment.start())
.collect();
if let Some(last_end) = filtered.last().map(TranscriptionSegment::end) {
seek = seek.max((last_end * SAMPLE_RATE as f32) as usize);
}
current_segments = Some(filtered);
}
decoded_segment_span = decoded_segment_span.saturating_add(if options.drop_blank_audio() {
allocated_this_window
} else {
current_segments.as_ref().map_or(0, Vec::len)
});
allocated_ordinals = allocated_ordinals.saturating_add(allocated_this_window);
if let Some(max_window_seek) = options.max_window_seek() {
seek = seek.min(previous_seek.saturating_add(max_window_seek.max(1)));
}
let Some(current_segments) = current_segments else {
continue;
};
if let Some(callback) = self.segment_callback {
callback(¤t_segments);
}
all_segments.extend(current_segments);
timings.set_decoding_windowing(
timings.decoding_windowing() + windowing_start.elapsed().as_secs_f64(),
);
timings.set_total_decoding_windows(timings.total_decoding_windows() + 1.0);
self.backend.reset_decoder_state(&mut state);
}
}
timings.set_decoding_loop(decode_loop_start.elapsed().as_secs_f64());
let special_token_begin = self.tokenizer.special_tokens().special_token_begin();
if options.drop_blank_audio() {
self.drop_blank_audio_segments(&mut all_segments, special_token_begin)?;
}
let word_tokens: Vec<u32> = all_segments
.iter()
.flat_map(|segment| segment.tokens_slice().iter().copied())
.filter(|&token| token < special_token_begin)
.collect();
let text = self.tokenizer.decode(&word_tokens, false)?;
let trimmed_text = trim_swift_whitespaces(&text);
timings.set_full_pipeline(pipeline_start.elapsed().as_secs_f64());
Ok(
TranscriptionResult::new(
trimmed_text,
all_segments,
detected_language
.clone()
.unwrap_or_else(|| DEFAULT_LANGUAGE_CODE.to_string()),
timings,
)
.with_language_observations(language_observations)
.with_task_facts({
let facts = facts_sink.lock().unwrap_or_else(PoisonError::into_inner);
facts
.clone()
.with_observed_language(observed_language)
.with_worker(self.window_id_offset)
.with_decoded_span(SpanKnowledge::Exact(allocated_ordinals))
}),
)
}
fn drop_blank_audio_segments(
&self,
segments: &mut Vec<TranscriptionSegment>,
special_token_begin: u32,
) -> Result<(), TranscribeError> {
let mut blank = Vec::with_capacity(segments.len());
for segment in segments.iter() {
let clean_tokens: Vec<u32> = segment
.tokens_slice()
.iter()
.copied()
.filter(|&token| token < special_token_begin)
.collect();
let clean_text = self.tokenizer.decode(&clean_tokens, false)?;
blank.push(trim_swift_whitespaces(&clean_text) == BLANK_AUDIO_MARKER);
}
if !blank.contains(&true) {
return Ok(());
}
let survivors: Vec<TranscriptionSegment> = segments
.drain(..)
.zip(blank)
.filter_map(|(segment, is_blank)| (!is_blank).then_some(segment))
.collect();
*segments = survivors;
Ok(())
}
#[allow(clippy::too_many_arguments)] fn decode_with_fallback(
&self,
encoder_output: &B::EncoderOutput,
state: &mut B::DecoderState,
initial_prompt: &mut Vec<u32>,
detected_language: &mut Option<String>,
observed_language: &mut Option<String>,
window_probe: &mut Option<LanguageDetection>,
facts_sink: &Mutex<TaskFacts>,
options: &DecodingOptions,
timings: &mut TranscriptionTimings,
window_index: u64,
) -> Result<(DecodingResult, Option<AlignmentMatrix>), TranscribeError> {
let special = *self.tokenizer.special_tokens();
let window_id = self.window_id_offset
+ (timings.total_decoding_windows() - timings.total_decoding_fallbacks()).max(0.0) as usize;
let stamped = self.progress_callback.map(|callback| {
move |progress: &TranscriptionProgress| -> Option<bool> {
let mut with_id = progress.clone();
with_id.set_window_id(window_id);
callback(&with_id)
}
});
let window_callback: Option<TranscriptionProgressCallback<'_>> = stamped
.as_ref()
.map(|wrapper| wrapper as &(dyn Fn(&TranscriptionProgress) -> Option<bool> + Sync));
let mut decoding = None;
let mut captured_alignment: Option<AlignmentMatrix> = None;
for attempt in 0..=options.temperature_fallback_count() {
let attempt_start = Instant::now();
let temperature =
options.temperature() + attempt as f32 * options.temperature_increment_on_fallback();
let mut sampler = GreedyTokenSampler::new(temperature, special.end_token(), options);
if let Some(seed) = options.seed() {
sampler = sampler.with_seed(sampler::derive_attempt_seed(
seed,
self.window_id_offset as u64,
window_index,
attempt as u64,
));
}
let early_stop = AtomicBool::new(false);
let observed_language_token: Cell<Option<u32>> = Cell::new(None);
let mut language_probe_swallowed = false;
let mut window_options = options.clone();
if self.backend.dims().is_multilingual()
&& options.language().is_empty()
&& options.detect_language()
{
match decode::detect_language(
self.backend,
encoder_output,
state,
self.tokenizer,
&mut sampler,
timings,
) {
Ok(probe) => {
window_options.set_language(probe.language().to_string());
*detected_language = Some(probe.language().to_string());
*window_probe = Some(LanguageDetection::new(
probe.language(),
probe.language_probs_slice().to_vec(),
));
if observed_language.is_none() && !probe.language_probs_slice().is_empty() {
*observed_language = Some(probe.language().to_string());
}
}
Err(_) => {
*detected_language = None;
*window_probe = None;
language_probe_swallowed = true;
}
}
if options.use_prefill_prompt() {
*initial_prompt = decode::prefill_tokens(&window_options, self.tokenizer, true);
}
}
crate::audio::whisper::text::clear_compression_error_swallowed();
let outcome = decode::decode_text(
self.backend,
encoder_output,
state,
initial_prompt.as_slice(),
&mut sampler,
&window_options,
self.tokenizer,
timings,
&early_stop,
&observed_language_token,
window_callback,
);
let attempt_observation = observed_language.clone().or_else(|| {
observed_language_token.get().and_then(|token| {
self.tokenizer.decode(&[token], false).ok().map(|decoded| {
crate::audio::whisper::text::trim_special_token_chars(&decoded).to_string()
})
})
});
facts_sink
.lock()
.unwrap_or_else(PoisonError::into_inner)
.merge(
&TaskFacts::unknown()
.with_drew_from_rng(sampler.drew_from_rng())
.with_early_stopped(early_stop.load(Ordering::Relaxed))
.with_had_swallowed_error(
language_probe_swallowed
|| crate::audio::whisper::text::take_compression_error_swallowed(),
)
.with_observed_language(attempt_observation),
);
let result = outcome?;
if options.word_timestamps() {
captured_alignment = self
.backend
.alignment_weights(state)
.map(|view| view.to_matrix());
}
if detected_language.is_none() {
*detected_language = Some(result.language().to_string());
}
if observed_language.is_none()
&& let Some(predicted) = result.observed_language()
{
*observed_language = Some(predicted.to_string());
}
let is_first_token_log_prob_too_low = options
.first_token_logprob_threshold()
.is_some_and(|threshold| result.first_token_log_prob() < threshold);
let fallback = needs_fallback(is_first_token_log_prob_too_low, &result, options);
decoding = Some(result);
match fallback {
Some(_reason) => {
timings.set_decoding_fallback(
timings.decoding_fallback() + attempt_start.elapsed().as_secs_f64(),
);
self.backend.reset_decoder_state(state);
timings.set_total_decoding_fallbacks(attempt as f64);
}
None => break,
}
}
Ok((
decoding.expect("the loop runs at least once (0..=count always yields >= 1 attempt)"),
captured_alignment,
))
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct LanguageDetection {
language: String,
probs: Vec<(String, f32)>,
}
impl LanguageDetection {
pub fn new(language: impl Into<String>, probs: impl Into<Vec<(String, f32)>>) -> Self {
Self {
language: language.into(),
probs: probs.into(),
}
}
#[inline(always)]
pub fn language(&self) -> &str {
self.language.as_str()
}
#[inline(always)]
pub const fn probs_slice(&self) -> &[(String, f32)] {
self.probs.as_slice()
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct LanguageObservation {
start: f32,
end: f32,
detection: LanguageDetection,
}
impl LanguageObservation {
pub const fn new(start: f32, end: f32, detection: LanguageDetection) -> Self {
Self {
start,
end,
detection,
}
}
#[inline(always)]
pub const fn start(&self) -> f32 {
self.start
}
#[inline(always)]
pub const fn set_start(&mut self, start: f32) -> &mut Self {
self.start = start;
self
}
#[inline(always)]
pub const fn end(&self) -> f32 {
self.end
}
#[inline(always)]
pub const fn set_end(&mut self, end: f32) -> &mut Self {
self.end = end;
self
}
#[inline(always)]
pub const fn detection(&self) -> &LanguageDetection {
&self.detection
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
struct LoadTimings {
model_loading: Duration,
prewarm_load_time: Duration,
encoder_load: Duration,
decoder_load: Duration,
encoder_specialization: Duration,
decoder_specialization: Duration,
tokenizer_load_time: Duration,
}
impl LoadTimings {
fn stamp(&self, timings: &mut TranscriptionTimings) {
timings
.set_model_loading(self.model_loading.as_secs_f64())
.set_prewarm_load_time(self.prewarm_load_time.as_secs_f64())
.set_encoder_load_time(self.encoder_load.as_secs_f64())
.set_decoder_load_time(self.decoder_load.as_secs_f64())
.set_encoder_specialization_time(self.encoder_specialization.as_secs_f64())
.set_decoder_specialization_time(self.decoder_specialization.as_secs_f64())
.set_tokenizer_load_time(self.tokenizer_load_time.as_secs_f64());
}
}
pub struct WhisperKit<B> {
backend: B,
tokenizer: WhisperTokenizer,
variant: Option<ModelVariant>,
vad_detector: Box<dyn audio::vad::VoiceActivityDetector + Send + Sync>,
load_timings: LoadTimings,
}
impl WhisperKit<CoreMlBackend> {
pub fn new(options: &Options) -> Result<Self, TranscribeError> {
if !options.load() {
return Err(
ModelError::InvalidState(InvalidState::new(
"load = true (WhisperKit::new always loads at construction)",
"load = false",
))
.into(),
);
}
let mut manager = ModelManager::new(options.model_folder(), options.compute());
let prewarm_load_time = if options.prewarm() {
let start = Instant::now();
manager.prewarm()?;
start.elapsed()
} else {
Duration::ZERO
};
let model_load_start = Instant::now();
let (models, model_splits): (_, ModelLoadTimings) = manager.into_loaded()?;
let tokenizer_start = Instant::now();
let tokenizer = WhisperTokenizer::from_folder(options.tokenizer_folder())?;
let tokenizer_load_time = tokenizer_start.elapsed();
let backend =
CoreMlBackend::from_loaded(models, tokenizer.vocab_size()).map_err(DecodeError::from)?;
let dims = backend.dims();
let variant = detect_variant(dims.vocab(), dims.embed_dim());
let model_loading = model_load_start.elapsed() + prewarm_load_time;
Ok(Self {
backend,
tokenizer,
variant,
vad_detector: Box::new(audio::vad::EnergyVad::new()),
load_timings: LoadTimings {
model_loading,
prewarm_load_time,
encoder_load: model_splits.encoder_load(),
decoder_load: model_splits.decoder_load(),
encoder_specialization: model_splits.encoder_specialization(),
decoder_specialization: model_splits.decoder_specialization(),
tokenizer_load_time,
},
})
}
}
impl<B> WhisperKit<B> {
pub fn with_backend(backend: B, tokenizer: WhisperTokenizer) -> Self {
Self {
backend,
tokenizer,
variant: None,
vad_detector: Box::new(audio::vad::EnergyVad::new()),
load_timings: LoadTimings::default(),
}
}
#[inline(always)]
pub const fn backend(&self) -> &B {
&self.backend
}
#[inline(always)]
pub const fn tokenizer(&self) -> &WhisperTokenizer {
&self.tokenizer
}
#[inline(always)]
pub const fn variant(&self) -> Option<ModelVariant> {
self.variant
}
#[inline(always)]
pub fn vad_detector(&self) -> &(dyn audio::vad::VoiceActivityDetector + Send + Sync) {
self.vad_detector.as_ref()
}
#[must_use]
#[inline(always)]
pub fn with_vad_detector(
mut self,
detector: Box<dyn audio::vad::VoiceActivityDetector + Send + Sync>,
) -> Self {
self.set_vad_detector(detector);
self
}
#[inline(always)]
pub fn set_vad_detector(
&mut self,
detector: Box<dyn audio::vad::VoiceActivityDetector + Send + Sync>,
) -> &mut Self {
self.vad_detector = detector;
self
}
#[inline(always)]
pub fn audio_stream_transcriber(
&self,
decoding_options: DecodingOptions,
) -> AudioStreamTranscriber<'_, B> {
AudioStreamTranscriber::new(self.backend(), self.tokenizer(), decoding_options)
}
#[inline(always)]
pub fn local_agreement_transcriber(
&self,
options: DecodingOptions,
) -> LocalAgreementTranscriber<'_, B> {
LocalAgreementTranscriber::new(self, options)
}
fn stamp_load_timings(&self, result: &mut TranscriptionResult) {
let mut timings = result.timings().clone();
self.load_timings.stamp(&mut timings);
result.set_timings(timings);
}
}
fn recover_vad_run_facts(
sink: TaskFacts,
worker_schedule: Option<Vec<usize>>,
decoded_span: SpanKnowledge,
) -> TaskFacts {
sink
.with_worker_schedule(worker_schedule)
.with_decoded_span(decoded_span)
}
impl<B> WhisperKit<B>
where
B: InferenceBackend,
{
pub fn transcribe(
&self,
audio: &[f32],
options: &DecodingOptions,
) -> Result<TranscriptionResult, TranscribeError> {
let window_samples = self.backend.dims().window_samples();
if options.chunking_strategy().is_vad() && audio.len() > window_samples {
let vad_chunker = chunker::VadChunker::new();
let clip_ranges = chunker::prepare_seek_clips(options.clip_timestamps_slice(), audio.len())?;
let detection_generation = self.vad_detector.detection_generation();
let chunks = vad_chunker.chunk_all(
self.vad_detector.as_ref(),
audio,
window_samples,
&clip_ranges,
);
if self.vad_detector.detection_generation() != detection_generation {
let source = self
.vad_detector
.last_detection_error()
.unwrap_or_else(|| "hard model inference failure during VAD chunking".into());
return Err(TranscribeError::Vad(VadError::Detection(source)));
}
let chunk_options = options.clone().with_clip_timestamps(Vec::new());
let facts_sink = Mutex::new(TaskFacts::observed_clean());
let mut chunk_results = Vec::with_capacity(chunks.len());
let mut schedule = TaskFacts::unknown().with_worker_schedule(Some(Vec::new()));
let mut any_chunk_dropped = false;
for (chunk_index, chunk) in chunks.iter().enumerate() {
let outcome = TranscribeTask::new(&self.backend, &self.tokenizer)
.with_window_id_offset(chunk_index)
.with_facts_sink(&facts_sink)
.run(chunk.samples_slice(), &chunk_options);
let coordinate = if let Ok(mut result) = outcome {
chunker::apply_result_seek_offset(&mut result, chunk.seek_offset());
chunk_results.push(result);
TaskFacts::unknown().with_worker(chunk_index)
} else {
facts_sink
.lock()
.unwrap_or_else(PoisonError::into_inner)
.merge(&TaskFacts::observed_clean().with_had_swallowed_error(true));
any_chunk_dropped = true;
TaskFacts::unknown()
};
schedule.merge(&coordinate);
}
let mut merged = merge_transcription_results_with_options(&chunk_results, options);
let sink_facts = facts_sink
.into_inner()
.unwrap_or_else(PoisonError::into_inner);
let decoded_span = if any_chunk_dropped {
SpanKnowledge::AtLeast(merged.task_facts().decoded_span().lower_bound())
} else if chunk_results.is_empty() {
SpanKnowledge::Exact(0)
} else {
merged.task_facts().decoded_span()
};
let recovered = recover_vad_run_facts(
sink_facts,
schedule.worker_schedule().map(|s| s.to_vec()),
decoded_span,
);
*merged.task_facts_mut() = recovered;
self.stamp_load_timings(&mut merged);
return Ok(merged);
}
let mut result = TranscribeTask::new(&self.backend, &self.tokenizer).run(audio, options)?;
self.stamp_load_timings(&mut result);
Ok(result)
}
pub fn transcribe_all(
&self,
audios: &[&[f32]],
options: &DecodingOptions,
) -> Vec<Result<TranscriptionResult, TranscribeError>>
where
B: Sync,
{
let batch_size = options.concurrent_worker_count().get();
let mut results = Vec::with_capacity(audios.len());
for (batch_index, batch) in audios.chunks(batch_size).enumerate() {
let batch_results = std::thread::scope(|scope| {
let handles: Vec<_> = batch
.iter()
.enumerate()
.map(|(audio_index, &audio)| {
let global_index = audio_index + batch_index * batch_size;
scope.spawn(move || {
TranscribeTask::new(&self.backend, &self.tokenizer)
.with_window_id_offset(global_index)
.run(audio, options)
})
})
.collect();
handles
.into_iter()
.map(|handle| handle.join().expect("transcribe worker thread panicked"))
.collect::<Vec<_>>()
});
results.extend(batch_results);
}
for result in results.iter_mut().flatten() {
self.stamp_load_timings(result);
}
results
}
pub fn detect_language(&self, audio: &[f32]) -> Result<LanguageDetection, TranscribeError> {
let dims = self.backend.dims();
if !dims.is_multilingual() {
return Err(
ModelError::InvalidState(InvalidState::new(
"a multilingual model",
"a monolingual (English-only) model",
))
.into(),
);
}
let padded = audio::pad_or_trim(audio, dims.window_samples());
let features = self
.backend
.extract_features(&padded)
.map_err(DecodeError::from)?;
let encoder_output = self.backend.encode(&features).map_err(DecodeError::from)?;
let mut state = self
.backend
.new_decoder_state()
.map_err(DecodeError::from)?;
let mut timings = TranscriptionTimings::new();
let special = *self.tokenizer.special_tokens();
let mut sampler =
decode::sampler::GreedyTokenSampler::new(0.0, special.end_token(), &DecodingOptions::new());
let probe = decode::detect_language(
&self.backend,
&encoder_output,
&mut state,
&self.tokenizer,
&mut sampler,
&mut timings,
)?;
Ok(LanguageDetection::new(
probe.language(),
probe.language_probs_slice().to_vec(),
))
}
}