use std::{
num::NonZeroUsize,
path::{Path, PathBuf},
};
use crate::ComputeUnits;
use crate::audio::whisper::constants::MAX_TOKEN_CONTEXT;
#[derive(
Debug, Default, Clone, Copy, PartialEq, Eq, Hash, derive_more::Display, derive_more::IsVariant,
)]
#[display("{}", self.as_str())]
#[non_exhaustive]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(rename_all = "snake_case"))]
pub enum Task {
#[default]
Transcribe,
Translate,
}
impl Task {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
Self::Transcribe => "transcribe",
Self::Translate => "translate",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error("unknown task name")]
pub struct ParseTaskError(());
impl core::str::FromStr for Task {
type Err = ParseTaskError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(match s {
"transcribe" => Self::Transcribe,
"translate" => Self::Translate,
_ => return Err(ParseTaskError(())),
})
}
}
#[derive(
Debug, Default, Clone, Copy, PartialEq, Eq, Hash, derive_more::Display, derive_more::IsVariant,
)]
#[display("{}", self.as_str())]
#[non_exhaustive]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(rename_all = "snake_case"))]
pub enum ChunkingStrategy {
#[default]
#[cfg_attr(feature = "serde", serde(rename = "none"))]
Disabled,
Vad,
}
impl ChunkingStrategy {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
Self::Disabled => "none",
Self::Vad => "vad",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error("unknown chunking strategy name")]
pub struct ParseChunkingStrategyError(());
impl core::str::FromStr for ChunkingStrategy {
type Err = ParseChunkingStrategyError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(match s {
"none" => Self::Disabled,
"vad" => Self::Vad,
_ => return Err(ParseChunkingStrategyError(())),
})
}
}
#[derive(
Debug, Default, Clone, Copy, PartialEq, Eq, Hash, derive_more::Display, derive_more::IsVariant,
)]
#[display("{}", self.as_str())]
#[non_exhaustive]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(rename_all = "snake_case"))]
pub enum WordGrouping {
FineGrained,
#[default]
SwiftParity,
}
impl WordGrouping {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
Self::FineGrained => "fine_grained",
Self::SwiftParity => "swift_parity",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error("unknown word grouping name")]
pub struct ParseWordGroupingError(());
impl core::str::FromStr for WordGrouping {
type Err = ParseWordGroupingError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(match s {
"fine_grained" => Self::FineGrained,
"swift_parity" => Self::SwiftParity,
_ => return Err(ParseWordGroupingError(())),
})
}
}
#[derive(
Debug, Default, Clone, Copy, PartialEq, Eq, Hash, derive_more::Display, derive_more::IsVariant,
)]
#[display("{}", self.as_str())]
#[non_exhaustive]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(rename_all = "snake_case"))]
pub enum AlignmentGather {
SwiftParity,
#[default]
Complete,
}
impl AlignmentGather {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
Self::SwiftParity => "swift_parity",
Self::Complete => "complete",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error("unknown alignment gather name")]
pub struct ParseAlignmentGatherError(());
impl core::str::FromStr for AlignmentGather {
type Err = ParseAlignmentGatherError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(match s {
"swift_parity" => Self::SwiftParity,
"complete" => Self::Complete,
_ => return Err(ParseAlignmentGatherError(())),
})
}
}
pub const DEFAULT_TEMPERATURE: f32 = 0.0;
pub const DEFAULT_TEMPERATURE_INCREMENT_ON_FALLBACK: f32 = 0.2;
pub const DEFAULT_TEMPERATURE_FALLBACK_COUNT: usize = 5;
pub const DEFAULT_SAMPLE_LENGTH: usize = MAX_TOKEN_CONTEXT;
pub const DEFAULT_TOP_K: usize = 5;
pub const DEFAULT_WINDOW_CLIP_TIME: f32 = 1.0;
pub const DEFAULT_COMPRESSION_RATIO_THRESHOLD: f32 = 2.4;
pub const DEFAULT_LOGPROB_THRESHOLD: f32 = -1.0;
pub const DEFAULT_FIRST_TOKEN_LOGPROB_THRESHOLD: f32 = -1.5;
pub const DEFAULT_NO_SPEECH_THRESHOLD: f32 = 0.6;
pub const DEFAULT_USE_PREFILL_PROMPT: bool = true;
pub const DEFAULT_DROP_BLANK_AUDIO: bool = true;
pub const DEFAULT_CONCURRENT_WORKER_COUNT: NonZeroUsize = NonZeroUsize::new(16).unwrap();
#[cfg(feature = "serde")]
fn default_temperature_increment_on_fallback() -> f32 {
DEFAULT_TEMPERATURE_INCREMENT_ON_FALLBACK
}
#[cfg(feature = "serde")]
fn default_temperature_fallback_count() -> usize {
DEFAULT_TEMPERATURE_FALLBACK_COUNT
}
#[cfg(feature = "serde")]
fn default_sample_length() -> usize {
DEFAULT_SAMPLE_LENGTH
}
#[cfg(feature = "serde")]
fn default_top_k() -> usize {
DEFAULT_TOP_K
}
#[cfg(feature = "serde")]
fn default_window_clip_time() -> f32 {
DEFAULT_WINDOW_CLIP_TIME
}
#[cfg(feature = "serde")]
fn default_concurrent_worker_count() -> NonZeroUsize {
DEFAULT_CONCURRENT_WORKER_COUNT
}
#[cfg(feature = "serde")]
fn default_compression_ratio_threshold() -> Option<f32> {
Some(DEFAULT_COMPRESSION_RATIO_THRESHOLD)
}
#[cfg(feature = "serde")]
fn default_logprob_threshold() -> Option<f32> {
Some(DEFAULT_LOGPROB_THRESHOLD)
}
#[cfg(feature = "serde")]
fn default_first_token_logprob_threshold() -> Option<f32> {
Some(DEFAULT_FIRST_TOKEN_LOGPROB_THRESHOLD)
}
#[cfg(feature = "serde")]
fn default_no_speech_threshold() -> Option<f32> {
Some(DEFAULT_NO_SPEECH_THRESHOLD)
}
#[cfg(feature = "serde")]
fn default_use_prefill_prompt() -> bool {
DEFAULT_USE_PREFILL_PROMPT
}
#[cfg(feature = "serde")]
fn default_drop_blank_audio() -> bool {
DEFAULT_DROP_BLANK_AUDIO
}
#[cfg(feature = "serde")]
pub(crate) const NON_FINITE_FLOAT_MSG: &str = "non-finite float (NaN or infinity) is not \
representable in JSON and is rejected to keep the serde round trip lossless (matches Swift's \
JSONEncoder, which throws on non-finite by default)";
#[cfg(feature = "serde")]
pub(crate) mod finite_f32 {
use serde::{Deserialize, Deserializer, Serialize, Serializer};
pub(crate) fn serialize<S: Serializer>(value: &f32, serializer: S) -> Result<S::Ok, S::Error> {
if !value.is_finite() {
return Err(serde::ser::Error::custom(super::NON_FINITE_FLOAT_MSG));
}
value.serialize(serializer)
}
pub(crate) fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<f32, D::Error> {
let value = f32::deserialize(deserializer)?;
if !value.is_finite() {
return Err(serde::de::Error::custom(super::NON_FINITE_FLOAT_MSG));
}
Ok(value)
}
}
#[cfg(feature = "serde")]
pub(crate) mod finite_f32_option {
use serde::{Deserialize, Deserializer, Serialize, Serializer};
pub(crate) fn serialize<S: Serializer>(
value: &Option<f32>,
serializer: S,
) -> Result<S::Ok, S::Error> {
if matches!(value, Some(v) if !v.is_finite()) {
return Err(serde::ser::Error::custom(super::NON_FINITE_FLOAT_MSG));
}
value.serialize(serializer)
}
pub(crate) fn deserialize<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Option<f32>, D::Error> {
let value = Option::<f32>::deserialize(deserializer)?;
if matches!(value, Some(v) if !v.is_finite()) {
return Err(serde::de::Error::custom(super::NON_FINITE_FLOAT_MSG));
}
Ok(value)
}
}
#[cfg(feature = "serde")]
pub(crate) mod finite_f32_vec {
use serde::{Deserialize, Deserializer, Serialize, Serializer};
pub(crate) fn serialize<S: Serializer>(value: &[f32], serializer: S) -> Result<S::Ok, S::Error> {
if value.iter().any(|v| !v.is_finite()) {
return Err(serde::ser::Error::custom(super::NON_FINITE_FLOAT_MSG));
}
value.serialize(serializer)
}
pub(crate) fn deserialize<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Vec<f32>, D::Error> {
let value = Vec::<f32>::deserialize(deserializer)?;
if value.iter().any(|v| !v.is_finite()) {
return Err(serde::de::Error::custom(super::NON_FINITE_FLOAT_MSG));
}
Ok(value)
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct DecodingOptions {
#[cfg_attr(feature = "serde", serde(default))]
task: Task,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "String::is_empty")
)]
language: String,
#[cfg_attr(feature = "serde", serde(default, with = "finite_f32"))]
temperature: f32,
#[cfg_attr(
feature = "serde",
serde(
default = "default_temperature_increment_on_fallback",
with = "finite_f32"
)
)]
temperature_increment_on_fallback: f32,
#[cfg_attr(
feature = "serde",
serde(default = "default_temperature_fallback_count")
)]
temperature_fallback_count: usize,
#[cfg_attr(feature = "serde", serde(default = "default_sample_length"))]
sample_length: usize,
#[cfg_attr(feature = "serde", serde(default = "default_top_k"))]
top_k: usize,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Option::is_none")
)]
seed: Option<u64>,
#[cfg_attr(feature = "serde", serde(default = "default_use_prefill_prompt"))]
use_prefill_prompt: bool,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Option::is_none")
)]
detect_language: Option<bool>,
#[cfg_attr(feature = "serde", serde(default))]
skip_special_tokens: bool,
#[cfg_attr(feature = "serde", serde(default))]
without_timestamps: bool,
#[cfg_attr(feature = "serde", serde(default))]
word_timestamps: bool,
#[cfg_attr(
feature = "serde",
serde(
default,
skip_serializing_if = "Option::is_none",
with = "finite_f32_option"
)
)]
max_initial_timestamp: Option<f32>,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Option::is_none")
)]
max_window_seek: Option<usize>,
#[cfg_attr(
feature = "serde",
serde(
default,
skip_serializing_if = "Vec::is_empty",
with = "finite_f32_vec"
)
)]
clip_timestamps: Vec<f32>,
#[cfg_attr(
feature = "serde",
serde(default = "default_window_clip_time", with = "finite_f32")
)]
window_clip_time: f32,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Vec::is_empty")
)]
prompt_tokens: Vec<u32>,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Vec::is_empty")
)]
prefix_tokens: Vec<u32>,
#[cfg_attr(feature = "serde", serde(default))]
suppress_blank: bool,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Vec::is_empty")
)]
suppress_tokens: Vec<u32>,
#[cfg_attr(
feature = "serde",
serde(
default = "default_compression_ratio_threshold",
with = "finite_f32_option"
)
)]
compression_ratio_threshold: Option<f32>,
#[cfg_attr(
feature = "serde",
serde(default = "default_logprob_threshold", with = "finite_f32_option")
)]
logprob_threshold: Option<f32>,
#[cfg_attr(
feature = "serde",
serde(
default = "default_first_token_logprob_threshold",
with = "finite_f32_option"
)
)]
first_token_logprob_threshold: Option<f32>,
#[cfg_attr(
feature = "serde",
serde(default = "default_no_speech_threshold", with = "finite_f32_option")
)]
no_speech_threshold: Option<f32>,
#[cfg_attr(feature = "serde", serde(default = "default_concurrent_worker_count"))]
concurrent_worker_count: NonZeroUsize,
#[cfg_attr(feature = "serde", serde(default))]
chunking_strategy: ChunkingStrategy,
#[cfg_attr(feature = "serde", serde(default))]
verbose: bool,
#[cfg_attr(feature = "serde", serde(default = "default_drop_blank_audio"))]
drop_blank_audio: bool,
#[cfg_attr(feature = "serde", serde(default))]
word_grouping: WordGrouping,
#[cfg_attr(feature = "serde", serde(default))]
alignment_gather: AlignmentGather,
}
macro_rules! decoding_option_field_names {
($($field:ident),+ $(,)?) => {
#[cfg(test)]
#[allow(dead_code)] pub(crate) const DECODING_OPTION_FIELD_NAMES: &[&str] = &[$(stringify!($field)),+];
#[cfg(test)]
#[allow(dead_code)]
fn _decoding_options_field_exhaustiveness_guard(options: DecodingOptions) {
let DecodingOptions { $($field: _),+ } = options;
}
};
}
decoding_option_field_names!(
task,
language,
temperature,
temperature_increment_on_fallback,
temperature_fallback_count,
sample_length,
top_k,
seed,
use_prefill_prompt,
detect_language,
skip_special_tokens,
without_timestamps,
word_timestamps,
max_initial_timestamp,
max_window_seek,
clip_timestamps,
window_clip_time,
prompt_tokens,
prefix_tokens,
suppress_blank,
suppress_tokens,
compression_ratio_threshold,
logprob_threshold,
first_token_logprob_threshold,
no_speech_threshold,
concurrent_worker_count,
chunking_strategy,
verbose,
drop_blank_audio,
word_grouping,
alignment_gather,
);
impl Default for DecodingOptions {
fn default() -> Self {
Self::new()
}
}
impl DecodingOptions {
pub const fn new() -> Self {
Self {
task: Task::Transcribe,
language: String::new(),
temperature: DEFAULT_TEMPERATURE,
temperature_increment_on_fallback: DEFAULT_TEMPERATURE_INCREMENT_ON_FALLBACK,
temperature_fallback_count: DEFAULT_TEMPERATURE_FALLBACK_COUNT,
sample_length: DEFAULT_SAMPLE_LENGTH,
top_k: DEFAULT_TOP_K,
seed: None,
use_prefill_prompt: DEFAULT_USE_PREFILL_PROMPT,
detect_language: None,
skip_special_tokens: false,
without_timestamps: false,
word_timestamps: false,
max_initial_timestamp: None,
max_window_seek: None,
clip_timestamps: Vec::new(),
window_clip_time: DEFAULT_WINDOW_CLIP_TIME,
prompt_tokens: Vec::new(),
prefix_tokens: Vec::new(),
suppress_blank: false,
suppress_tokens: Vec::new(),
compression_ratio_threshold: Some(DEFAULT_COMPRESSION_RATIO_THRESHOLD),
logprob_threshold: Some(DEFAULT_LOGPROB_THRESHOLD),
first_token_logprob_threshold: Some(DEFAULT_FIRST_TOKEN_LOGPROB_THRESHOLD),
no_speech_threshold: Some(DEFAULT_NO_SPEECH_THRESHOLD),
concurrent_worker_count: DEFAULT_CONCURRENT_WORKER_COUNT,
chunking_strategy: ChunkingStrategy::Disabled,
verbose: false,
drop_blank_audio: DEFAULT_DROP_BLANK_AUDIO,
word_grouping: WordGrouping::SwiftParity,
alignment_gather: AlignmentGather::Complete,
}
}
#[inline(always)]
pub const fn task(&self) -> Task {
self.task
}
#[must_use]
#[inline(always)]
pub const fn with_task(mut self, task: Task) -> Self {
self.set_task(task);
self
}
#[inline(always)]
pub const fn set_task(&mut self, task: Task) -> &mut Self {
self.task = task;
self
}
#[inline(always)]
pub fn language(&self) -> &str {
self.language.as_str()
}
#[must_use]
#[inline(always)]
pub fn with_language(mut self, language: impl Into<String>) -> Self {
self.set_language(language);
self
}
#[inline(always)]
pub fn set_language(&mut self, language: impl Into<String>) -> &mut Self {
self.language = language.into();
self
}
#[inline(always)]
pub const fn temperature(&self) -> f32 {
self.temperature
}
#[must_use]
#[inline(always)]
pub const fn with_temperature(mut self, temperature: f32) -> Self {
self.set_temperature(temperature);
self
}
#[inline(always)]
pub const fn set_temperature(&mut self, temperature: f32) -> &mut Self {
self.temperature = temperature;
self
}
#[inline(always)]
pub const fn temperature_increment_on_fallback(&self) -> f32 {
self.temperature_increment_on_fallback
}
#[must_use]
#[inline(always)]
pub const fn with_temperature_increment_on_fallback(
mut self,
temperature_increment_on_fallback: f32,
) -> Self {
self.set_temperature_increment_on_fallback(temperature_increment_on_fallback);
self
}
#[inline(always)]
pub const fn set_temperature_increment_on_fallback(
&mut self,
temperature_increment_on_fallback: f32,
) -> &mut Self {
self.temperature_increment_on_fallback = temperature_increment_on_fallback;
self
}
#[inline(always)]
pub const fn temperature_fallback_count(&self) -> usize {
self.temperature_fallback_count
}
#[must_use]
#[inline(always)]
pub const fn with_temperature_fallback_count(
mut self,
temperature_fallback_count: usize,
) -> Self {
self.set_temperature_fallback_count(temperature_fallback_count);
self
}
#[inline(always)]
pub const fn set_temperature_fallback_count(
&mut self,
temperature_fallback_count: usize,
) -> &mut Self {
self.temperature_fallback_count = temperature_fallback_count;
self
}
#[inline(always)]
pub const fn sample_length(&self) -> usize {
self.sample_length
}
#[must_use]
#[inline(always)]
pub const fn with_sample_length(mut self, sample_length: usize) -> Self {
self.set_sample_length(sample_length);
self
}
#[inline(always)]
pub const fn set_sample_length(&mut self, sample_length: usize) -> &mut Self {
self.sample_length = sample_length;
self
}
#[inline(always)]
pub const fn top_k(&self) -> usize {
self.top_k
}
#[must_use]
#[inline(always)]
pub const fn with_top_k(mut self, top_k: usize) -> Self {
self.set_top_k(top_k);
self
}
#[inline(always)]
pub const fn set_top_k(&mut self, top_k: usize) -> &mut Self {
self.top_k = top_k;
self
}
#[inline(always)]
pub const fn seed(&self) -> Option<u64> {
self.seed
}
#[must_use]
#[inline(always)]
pub const fn with_seed(mut self, seed: u64) -> Self {
self.set_seed(seed);
self
}
#[inline(always)]
pub const fn set_seed(&mut self, seed: u64) -> &mut Self {
self.seed = Some(seed);
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_seed(mut self, seed: Option<u64>) -> Self {
self.update_seed(seed);
self
}
#[inline(always)]
pub const fn update_seed(&mut self, seed: Option<u64>) -> &mut Self {
self.seed = seed;
self
}
#[inline(always)]
pub const fn clear_seed(&mut self) -> &mut Self {
self.seed = None;
self
}
#[inline(always)]
pub const fn use_prefill_prompt(&self) -> bool {
self.use_prefill_prompt
}
#[must_use]
#[inline(always)]
pub const fn with_use_prefill_prompt(mut self) -> Self {
self.set_use_prefill_prompt();
self
}
#[inline(always)]
pub const fn set_use_prefill_prompt(&mut self) -> &mut Self {
self.use_prefill_prompt = true;
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_use_prefill_prompt(mut self, use_prefill_prompt: bool) -> Self {
self.update_use_prefill_prompt(use_prefill_prompt);
self
}
#[inline(always)]
pub const fn update_use_prefill_prompt(&mut self, use_prefill_prompt: bool) -> &mut Self {
self.use_prefill_prompt = use_prefill_prompt;
self
}
#[inline(always)]
pub const fn clear_use_prefill_prompt(&mut self) -> &mut Self {
self.use_prefill_prompt = false;
self
}
#[inline(always)]
pub const fn detect_language(&self) -> bool {
match self.detect_language {
Some(explicit) => explicit,
None => !self.use_prefill_prompt,
}
}
#[must_use]
#[inline(always)]
pub const fn with_detect_language(mut self) -> Self {
self.set_detect_language();
self
}
#[inline(always)]
pub const fn set_detect_language(&mut self) -> &mut Self {
self.detect_language = Some(true);
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_detect_language(mut self, detect_language: bool) -> Self {
self.update_detect_language(detect_language);
self
}
#[inline(always)]
pub const fn update_detect_language(&mut self, detect_language: bool) -> &mut Self {
self.detect_language = Some(detect_language);
self
}
#[inline(always)]
pub const fn clear_detect_language(&mut self) -> &mut Self {
self.detect_language = Some(false);
self
}
#[inline(always)]
pub const fn skip_special_tokens(&self) -> bool {
self.skip_special_tokens
}
#[must_use]
#[inline(always)]
pub const fn with_skip_special_tokens(mut self) -> Self {
self.set_skip_special_tokens();
self
}
#[inline(always)]
pub const fn set_skip_special_tokens(&mut self) -> &mut Self {
self.skip_special_tokens = true;
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_skip_special_tokens(mut self, skip_special_tokens: bool) -> Self {
self.update_skip_special_tokens(skip_special_tokens);
self
}
#[inline(always)]
pub const fn update_skip_special_tokens(&mut self, skip_special_tokens: bool) -> &mut Self {
self.skip_special_tokens = skip_special_tokens;
self
}
#[inline(always)]
pub const fn clear_skip_special_tokens(&mut self) -> &mut Self {
self.skip_special_tokens = false;
self
}
#[inline(always)]
pub const fn without_timestamps(&self) -> bool {
self.without_timestamps
}
#[must_use]
#[inline(always)]
pub const fn with_without_timestamps(mut self) -> Self {
self.set_without_timestamps();
self
}
#[inline(always)]
pub const fn set_without_timestamps(&mut self) -> &mut Self {
self.without_timestamps = true;
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_without_timestamps(mut self, without_timestamps: bool) -> Self {
self.update_without_timestamps(without_timestamps);
self
}
#[inline(always)]
pub const fn update_without_timestamps(&mut self, without_timestamps: bool) -> &mut Self {
self.without_timestamps = without_timestamps;
self
}
#[inline(always)]
pub const fn clear_without_timestamps(&mut self) -> &mut Self {
self.without_timestamps = false;
self
}
#[inline(always)]
pub const fn word_timestamps(&self) -> bool {
self.word_timestamps
}
#[must_use]
#[inline(always)]
pub const fn with_word_timestamps(mut self) -> Self {
self.set_word_timestamps();
self
}
#[inline(always)]
pub const fn set_word_timestamps(&mut self) -> &mut Self {
self.word_timestamps = true;
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_word_timestamps(mut self, word_timestamps: bool) -> Self {
self.update_word_timestamps(word_timestamps);
self
}
#[inline(always)]
pub const fn update_word_timestamps(&mut self, word_timestamps: bool) -> &mut Self {
self.word_timestamps = word_timestamps;
self
}
#[inline(always)]
pub const fn clear_word_timestamps(&mut self) -> &mut Self {
self.word_timestamps = false;
self
}
#[inline(always)]
pub const fn max_initial_timestamp(&self) -> Option<f32> {
self.max_initial_timestamp
}
#[must_use]
#[inline(always)]
pub const fn with_max_initial_timestamp(mut self, max_initial_timestamp: f32) -> Self {
self.set_max_initial_timestamp(max_initial_timestamp);
self
}
#[inline(always)]
pub const fn set_max_initial_timestamp(&mut self, max_initial_timestamp: f32) -> &mut Self {
self.max_initial_timestamp = Some(max_initial_timestamp);
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_max_initial_timestamp(mut self, max_initial_timestamp: Option<f32>) -> Self {
self.update_max_initial_timestamp(max_initial_timestamp);
self
}
#[inline(always)]
pub const fn update_max_initial_timestamp(
&mut self,
max_initial_timestamp: Option<f32>,
) -> &mut Self {
self.max_initial_timestamp = max_initial_timestamp;
self
}
#[inline(always)]
pub const fn clear_max_initial_timestamp(&mut self) -> &mut Self {
self.max_initial_timestamp = None;
self
}
#[inline(always)]
pub const fn max_window_seek(&self) -> Option<usize> {
self.max_window_seek
}
#[must_use]
#[inline(always)]
pub const fn with_max_window_seek(mut self, max_window_seek: usize) -> Self {
self.set_max_window_seek(max_window_seek);
self
}
#[inline(always)]
pub const fn set_max_window_seek(&mut self, max_window_seek: usize) -> &mut Self {
self.max_window_seek = Some(max_window_seek);
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_max_window_seek(mut self, max_window_seek: Option<usize>) -> Self {
self.update_max_window_seek(max_window_seek);
self
}
#[inline(always)]
pub const fn update_max_window_seek(&mut self, max_window_seek: Option<usize>) -> &mut Self {
self.max_window_seek = max_window_seek;
self
}
#[inline(always)]
pub const fn clear_max_window_seek(&mut self) -> &mut Self {
self.max_window_seek = None;
self
}
#[inline(always)]
pub const fn clip_timestamps_slice(&self) -> &[f32] {
self.clip_timestamps.as_slice()
}
#[must_use]
#[inline(always)]
pub fn with_clip_timestamps(mut self, clip_timestamps: impl Into<Vec<f32>>) -> Self {
self.set_clip_timestamps(clip_timestamps);
self
}
#[inline(always)]
pub fn set_clip_timestamps(&mut self, clip_timestamps: impl Into<Vec<f32>>) -> &mut Self {
self.clip_timestamps = clip_timestamps.into();
self
}
#[inline(always)]
pub const fn window_clip_time(&self) -> f32 {
self.window_clip_time
}
#[must_use]
#[inline(always)]
pub const fn with_window_clip_time(mut self, window_clip_time: f32) -> Self {
self.set_window_clip_time(window_clip_time);
self
}
#[inline(always)]
pub const fn set_window_clip_time(&mut self, window_clip_time: f32) -> &mut Self {
self.window_clip_time = window_clip_time;
self
}
#[inline(always)]
pub const fn prompt_tokens_slice(&self) -> &[u32] {
self.prompt_tokens.as_slice()
}
#[must_use]
#[inline(always)]
pub fn with_prompt_tokens(mut self, prompt_tokens: impl Into<Vec<u32>>) -> Self {
self.set_prompt_tokens(prompt_tokens);
self
}
#[inline(always)]
pub fn set_prompt_tokens(&mut self, prompt_tokens: impl Into<Vec<u32>>) -> &mut Self {
self.prompt_tokens = prompt_tokens.into();
self
}
#[inline(always)]
pub const fn prefix_tokens_slice(&self) -> &[u32] {
self.prefix_tokens.as_slice()
}
#[must_use]
#[inline(always)]
pub fn with_prefix_tokens(mut self, prefix_tokens: impl Into<Vec<u32>>) -> Self {
self.set_prefix_tokens(prefix_tokens);
self
}
#[inline(always)]
pub fn set_prefix_tokens(&mut self, prefix_tokens: impl Into<Vec<u32>>) -> &mut Self {
self.prefix_tokens = prefix_tokens.into();
self
}
#[inline(always)]
pub const fn suppress_blank(&self) -> bool {
self.suppress_blank
}
#[must_use]
#[inline(always)]
pub const fn with_suppress_blank(mut self) -> Self {
self.set_suppress_blank();
self
}
#[inline(always)]
pub const fn set_suppress_blank(&mut self) -> &mut Self {
self.suppress_blank = true;
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_suppress_blank(mut self, suppress_blank: bool) -> Self {
self.update_suppress_blank(suppress_blank);
self
}
#[inline(always)]
pub const fn update_suppress_blank(&mut self, suppress_blank: bool) -> &mut Self {
self.suppress_blank = suppress_blank;
self
}
#[inline(always)]
pub const fn clear_suppress_blank(&mut self) -> &mut Self {
self.suppress_blank = false;
self
}
#[inline(always)]
pub const fn suppress_tokens_slice(&self) -> &[u32] {
self.suppress_tokens.as_slice()
}
#[must_use]
#[inline(always)]
pub fn with_suppress_tokens(mut self, suppress_tokens: impl Into<Vec<u32>>) -> Self {
self.set_suppress_tokens(suppress_tokens);
self
}
#[inline(always)]
pub fn set_suppress_tokens(&mut self, suppress_tokens: impl Into<Vec<u32>>) -> &mut Self {
self.suppress_tokens = suppress_tokens.into();
self
}
#[inline(always)]
pub const fn compression_ratio_threshold(&self) -> Option<f32> {
self.compression_ratio_threshold
}
#[must_use]
#[inline(always)]
pub const fn with_compression_ratio_threshold(
mut self,
compression_ratio_threshold: f32,
) -> Self {
self.set_compression_ratio_threshold(compression_ratio_threshold);
self
}
#[inline(always)]
pub const fn set_compression_ratio_threshold(
&mut self,
compression_ratio_threshold: f32,
) -> &mut Self {
self.compression_ratio_threshold = Some(compression_ratio_threshold);
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_compression_ratio_threshold(
mut self,
compression_ratio_threshold: Option<f32>,
) -> Self {
self.update_compression_ratio_threshold(compression_ratio_threshold);
self
}
#[inline(always)]
pub const fn update_compression_ratio_threshold(
&mut self,
compression_ratio_threshold: Option<f32>,
) -> &mut Self {
self.compression_ratio_threshold = compression_ratio_threshold;
self
}
#[inline(always)]
pub const fn clear_compression_ratio_threshold(&mut self) -> &mut Self {
self.compression_ratio_threshold = None;
self
}
#[inline(always)]
pub const fn logprob_threshold(&self) -> Option<f32> {
self.logprob_threshold
}
#[must_use]
#[inline(always)]
pub const fn with_logprob_threshold(mut self, logprob_threshold: f32) -> Self {
self.set_logprob_threshold(logprob_threshold);
self
}
#[inline(always)]
pub const fn set_logprob_threshold(&mut self, logprob_threshold: f32) -> &mut Self {
self.logprob_threshold = Some(logprob_threshold);
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_logprob_threshold(mut self, logprob_threshold: Option<f32>) -> Self {
self.update_logprob_threshold(logprob_threshold);
self
}
#[inline(always)]
pub const fn update_logprob_threshold(&mut self, logprob_threshold: Option<f32>) -> &mut Self {
self.logprob_threshold = logprob_threshold;
self
}
#[inline(always)]
pub const fn clear_logprob_threshold(&mut self) -> &mut Self {
self.logprob_threshold = None;
self
}
#[inline(always)]
pub const fn first_token_logprob_threshold(&self) -> Option<f32> {
self.first_token_logprob_threshold
}
#[must_use]
#[inline(always)]
pub const fn with_first_token_logprob_threshold(
mut self,
first_token_logprob_threshold: f32,
) -> Self {
self.set_first_token_logprob_threshold(first_token_logprob_threshold);
self
}
#[inline(always)]
pub const fn set_first_token_logprob_threshold(
&mut self,
first_token_logprob_threshold: f32,
) -> &mut Self {
self.first_token_logprob_threshold = Some(first_token_logprob_threshold);
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_first_token_logprob_threshold(
mut self,
first_token_logprob_threshold: Option<f32>,
) -> Self {
self.update_first_token_logprob_threshold(first_token_logprob_threshold);
self
}
#[inline(always)]
pub const fn update_first_token_logprob_threshold(
&mut self,
first_token_logprob_threshold: Option<f32>,
) -> &mut Self {
self.first_token_logprob_threshold = first_token_logprob_threshold;
self
}
#[inline(always)]
pub const fn clear_first_token_logprob_threshold(&mut self) -> &mut Self {
self.first_token_logprob_threshold = None;
self
}
#[inline(always)]
pub const fn no_speech_threshold(&self) -> Option<f32> {
self.no_speech_threshold
}
#[must_use]
#[inline(always)]
pub const fn with_no_speech_threshold(mut self, no_speech_threshold: f32) -> Self {
self.set_no_speech_threshold(no_speech_threshold);
self
}
#[inline(always)]
pub const fn set_no_speech_threshold(&mut self, no_speech_threshold: f32) -> &mut Self {
self.no_speech_threshold = Some(no_speech_threshold);
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_no_speech_threshold(mut self, no_speech_threshold: Option<f32>) -> Self {
self.update_no_speech_threshold(no_speech_threshold);
self
}
#[inline(always)]
pub const fn update_no_speech_threshold(
&mut self,
no_speech_threshold: Option<f32>,
) -> &mut Self {
self.no_speech_threshold = no_speech_threshold;
self
}
#[inline(always)]
pub const fn clear_no_speech_threshold(&mut self) -> &mut Self {
self.no_speech_threshold = None;
self
}
#[inline(always)]
pub const fn concurrent_worker_count(&self) -> NonZeroUsize {
self.concurrent_worker_count
}
#[must_use]
#[inline(always)]
pub const fn with_concurrent_worker_count(
mut self,
concurrent_worker_count: NonZeroUsize,
) -> Self {
self.set_concurrent_worker_count(concurrent_worker_count);
self
}
#[inline(always)]
pub const fn set_concurrent_worker_count(
&mut self,
concurrent_worker_count: NonZeroUsize,
) -> &mut Self {
self.concurrent_worker_count = concurrent_worker_count;
self
}
#[inline(always)]
pub const fn chunking_strategy(&self) -> ChunkingStrategy {
self.chunking_strategy
}
#[must_use]
#[inline(always)]
pub const fn with_chunking_strategy(mut self, chunking_strategy: ChunkingStrategy) -> Self {
self.set_chunking_strategy(chunking_strategy);
self
}
#[inline(always)]
pub const fn set_chunking_strategy(&mut self, chunking_strategy: ChunkingStrategy) -> &mut Self {
self.chunking_strategy = chunking_strategy;
self
}
#[inline(always)]
pub const fn verbose(&self) -> bool {
self.verbose
}
#[must_use]
#[inline(always)]
pub const fn with_verbose(mut self) -> Self {
self.set_verbose();
self
}
#[inline(always)]
pub const fn set_verbose(&mut self) -> &mut Self {
self.verbose = true;
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_verbose(mut self, verbose: bool) -> Self {
self.update_verbose(verbose);
self
}
#[inline(always)]
pub const fn update_verbose(&mut self, verbose: bool) -> &mut Self {
self.verbose = verbose;
self
}
#[inline(always)]
pub const fn clear_verbose(&mut self) -> &mut Self {
self.verbose = false;
self
}
#[inline(always)]
pub const fn drop_blank_audio(&self) -> bool {
self.drop_blank_audio
}
#[must_use]
#[inline(always)]
pub const fn with_drop_blank_audio(mut self) -> Self {
self.set_drop_blank_audio();
self
}
#[inline(always)]
pub const fn set_drop_blank_audio(&mut self) -> &mut Self {
self.drop_blank_audio = true;
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_drop_blank_audio(mut self, drop_blank_audio: bool) -> Self {
self.update_drop_blank_audio(drop_blank_audio);
self
}
#[inline(always)]
pub const fn update_drop_blank_audio(&mut self, drop_blank_audio: bool) -> &mut Self {
self.drop_blank_audio = drop_blank_audio;
self
}
#[inline(always)]
pub const fn clear_drop_blank_audio(&mut self) -> &mut Self {
self.drop_blank_audio = false;
self
}
#[inline(always)]
pub const fn word_grouping(&self) -> WordGrouping {
self.word_grouping
}
#[must_use]
#[inline(always)]
pub const fn with_word_grouping(mut self, word_grouping: WordGrouping) -> Self {
self.set_word_grouping(word_grouping);
self
}
#[inline(always)]
pub const fn set_word_grouping(&mut self, word_grouping: WordGrouping) -> &mut Self {
self.word_grouping = word_grouping;
self
}
#[inline(always)]
pub const fn alignment_gather(&self) -> AlignmentGather {
self.alignment_gather
}
#[must_use]
#[inline(always)]
pub const fn with_alignment_gather(mut self, alignment_gather: AlignmentGather) -> Self {
self.set_alignment_gather(alignment_gather);
self
}
#[inline(always)]
pub const fn set_alignment_gather(&mut self, alignment_gather: AlignmentGather) -> &mut Self {
self.alignment_gather = alignment_gather;
self
}
}
pub const DEFAULT_MEL_COMPUTE_UNITS: ComputeUnits = ComputeUnits::CpuAndGpu;
pub const DEFAULT_ENCODER_COMPUTE_UNITS: ComputeUnits = ComputeUnits::CpuAndNeuralEngine;
pub const DEFAULT_DECODER_COMPUTE_UNITS: ComputeUnits = ComputeUnits::CpuAndNeuralEngine;
#[cfg(feature = "serde")]
fn default_mel_compute_units() -> ComputeUnits {
DEFAULT_MEL_COMPUTE_UNITS
}
#[cfg(feature = "serde")]
fn default_encoder_compute_units() -> ComputeUnits {
DEFAULT_ENCODER_COMPUTE_UNITS
}
#[cfg(feature = "serde")]
fn default_decoder_compute_units() -> ComputeUnits {
DEFAULT_DECODER_COMPUTE_UNITS
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct ComputeOptions {
#[cfg_attr(feature = "serde", serde(default = "default_mel_compute_units"))]
mel: ComputeUnits,
#[cfg_attr(feature = "serde", serde(default = "default_encoder_compute_units"))]
encoder: ComputeUnits,
#[cfg_attr(feature = "serde", serde(default = "default_decoder_compute_units"))]
decoder: ComputeUnits,
}
impl Default for ComputeOptions {
fn default() -> Self {
Self::new()
}
}
impl ComputeOptions {
pub const fn new() -> Self {
Self {
mel: DEFAULT_MEL_COMPUTE_UNITS,
encoder: DEFAULT_ENCODER_COMPUTE_UNITS,
decoder: DEFAULT_DECODER_COMPUTE_UNITS,
}
}
#[inline(always)]
pub const fn mel(&self) -> ComputeUnits {
self.mel
}
#[must_use]
#[inline(always)]
pub const fn with_mel(mut self, mel: ComputeUnits) -> Self {
self.set_mel(mel);
self
}
#[inline(always)]
pub const fn set_mel(&mut self, mel: ComputeUnits) -> &mut Self {
self.mel = mel;
self
}
#[inline(always)]
pub const fn encoder(&self) -> ComputeUnits {
self.encoder
}
#[must_use]
#[inline(always)]
pub const fn with_encoder(mut self, encoder: ComputeUnits) -> Self {
self.set_encoder(encoder);
self
}
#[inline(always)]
pub const fn set_encoder(&mut self, encoder: ComputeUnits) -> &mut Self {
self.encoder = encoder;
self
}
#[inline(always)]
pub const fn decoder(&self) -> ComputeUnits {
self.decoder
}
#[must_use]
#[inline(always)]
pub const fn with_decoder(mut self, decoder: ComputeUnits) -> Self {
self.set_decoder(decoder);
self
}
#[inline(always)]
pub const fn set_decoder(&mut self, decoder: ComputeUnits) -> &mut Self {
self.decoder = decoder;
self
}
}
pub const DEFAULT_PREWARM: bool = false;
pub const DEFAULT_LOAD: bool = true;
#[cfg(feature = "serde")]
fn default_prewarm() -> bool {
DEFAULT_PREWARM
}
#[cfg(feature = "serde")]
fn default_load() -> bool {
DEFAULT_LOAD
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Options {
model_folder: PathBuf,
tokenizer_folder: PathBuf,
#[cfg_attr(feature = "serde", serde(default))]
compute: ComputeOptions,
#[cfg_attr(feature = "serde", serde(default = "default_prewarm"))]
prewarm: bool,
#[cfg_attr(feature = "serde", serde(default = "default_load"))]
load: bool,
}
impl Options {
pub fn new(model_folder: impl Into<PathBuf>, tokenizer_folder: impl Into<PathBuf>) -> Self {
Self {
model_folder: model_folder.into(),
tokenizer_folder: tokenizer_folder.into(),
compute: ComputeOptions::new(),
prewarm: DEFAULT_PREWARM,
load: DEFAULT_LOAD,
}
}
#[inline(always)]
pub fn model_folder(&self) -> &Path {
self.model_folder.as_path()
}
#[must_use]
#[inline(always)]
pub fn with_model_folder(mut self, model_folder: impl Into<PathBuf>) -> Self {
self.set_model_folder(model_folder);
self
}
#[inline(always)]
pub fn set_model_folder(&mut self, model_folder: impl Into<PathBuf>) -> &mut Self {
self.model_folder = model_folder.into();
self
}
#[inline(always)]
pub fn tokenizer_folder(&self) -> &Path {
self.tokenizer_folder.as_path()
}
#[must_use]
#[inline(always)]
pub fn with_tokenizer_folder(mut self, tokenizer_folder: impl Into<PathBuf>) -> Self {
self.set_tokenizer_folder(tokenizer_folder);
self
}
#[inline(always)]
pub fn set_tokenizer_folder(&mut self, tokenizer_folder: impl Into<PathBuf>) -> &mut Self {
self.tokenizer_folder = tokenizer_folder.into();
self
}
#[inline(always)]
pub const fn compute(&self) -> ComputeOptions {
self.compute
}
#[must_use]
#[inline(always)]
pub const fn with_compute(mut self, compute: ComputeOptions) -> Self {
self.set_compute(compute);
self
}
#[inline(always)]
pub const fn set_compute(&mut self, compute: ComputeOptions) -> &mut Self {
self.compute = compute;
self
}
#[inline(always)]
pub const fn prewarm(&self) -> bool {
self.prewarm
}
#[must_use]
#[inline(always)]
pub const fn with_prewarm(mut self) -> Self {
self.set_prewarm();
self
}
#[inline(always)]
pub const fn set_prewarm(&mut self) -> &mut Self {
self.prewarm = true;
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_prewarm(mut self, prewarm: bool) -> Self {
self.update_prewarm(prewarm);
self
}
#[inline(always)]
pub const fn update_prewarm(&mut self, prewarm: bool) -> &mut Self {
self.prewarm = prewarm;
self
}
#[inline(always)]
pub const fn clear_prewarm(&mut self) -> &mut Self {
self.prewarm = false;
self
}
#[inline(always)]
pub const fn load(&self) -> bool {
self.load
}
#[must_use]
#[inline(always)]
pub const fn with_load(mut self) -> Self {
self.set_load();
self
}
#[inline(always)]
pub const fn set_load(&mut self) -> &mut Self {
self.load = true;
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_load(mut self, load: bool) -> Self {
self.update_load(load);
self
}
#[inline(always)]
pub const fn update_load(&mut self, load: bool) -> &mut Self {
self.load = load;
self
}
#[inline(always)]
pub const fn clear_load(&mut self) -> &mut Self {
self.load = false;
self
}
}
#[cfg(test)]
mod tests;