use core::{
num::NonZeroU64,
sync::atomic::{AtomicU64, Ordering},
};
use std::sync::Arc;
use mediatime::TimeRange;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use smol_str::SmolStr;
use crate::types::{ChunkId, Lang, Word, WorkFailure};
#[derive(Clone, Debug)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct AsrParams {
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Option::is_none")
)]
language_hint: Option<Lang>,
#[cfg_attr(feature = "serde", serde(default))]
strategy: SamplingStrategy,
#[cfg_attr(feature = "serde", serde(default = "default_initial_temperature"))]
initial_temperature: f32,
#[cfg_attr(feature = "serde", serde(default = "default_temperature_increment"))]
temperature_increment: f32,
#[cfg_attr(
feature = "serde",
serde(
default = "default_max_attempts",
deserialize_with = "deserialize_nonzero_max_attempts"
)
)]
max_attempts: u8,
#[cfg_attr(feature = "serde", serde(default = "default_log_prob_threshold"))]
log_prob_threshold: f32,
#[cfg_attr(
feature = "serde",
serde(default = "default_compression_ratio_threshold")
)]
compression_ratio_threshold: f32,
#[cfg_attr(feature = "serde", serde(default = "default_no_speech_threshold"))]
no_speech_threshold: f32,
#[cfg_attr(feature = "serde", serde(default = "default_no_context"))]
no_context: bool,
#[cfg_attr(feature = "serde", serde(default = "default_suppress_blank"))]
suppress_blank: bool,
#[cfg_attr(feature = "serde", serde(default))]
suppress_non_speech_tokens: bool,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Option::is_none")
)]
initial_prompt: Option<SmolStr>,
#[cfg_attr(
feature = "serde",
serde(
default = "default_n_threads",
deserialize_with = "deserialize_positive_n_threads"
)
)]
n_threads: i32,
}
#[cfg(feature = "serde")]
const fn default_initial_temperature() -> f32 {
0.0
}
#[cfg(feature = "serde")]
const fn default_temperature_increment() -> f32 {
0.2
}
#[cfg(feature = "serde")]
const fn default_max_attempts() -> u8 {
6
}
#[cfg(feature = "serde")]
fn deserialize_nonzero_max_attempts<'de, D>(deserializer: D) -> Result<u8, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::Error as _;
let v = u8::deserialize(deserializer)?;
if v == 0 {
return Err(D::Error::custom(
"max_attempts must be > 0; use 1 for a single attempt with no retries",
));
}
Ok(v)
}
#[cfg(feature = "serde")]
const fn default_log_prob_threshold() -> f32 {
-1.0
}
#[cfg(feature = "serde")]
const fn default_compression_ratio_threshold() -> f32 {
2.4
}
#[cfg(feature = "serde")]
const fn default_no_speech_threshold() -> f32 {
0.6
}
#[cfg(feature = "serde")]
const fn default_no_context() -> bool {
true
}
#[cfg(feature = "serde")]
const fn default_suppress_blank() -> bool {
true
}
#[cfg(feature = "serde")]
const fn default_n_threads() -> i32 {
1
}
#[cfg(feature = "serde")]
fn deserialize_positive_n_threads<'de, D>(deserializer: D) -> Result<i32, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::Error as _;
let v = i32::deserialize(deserializer)?;
if v < 1 {
return Err(D::Error::custom(format!(
"n_threads must be >= 1 (got {v}); whisper.cpp would underflow / abort otherwise"
)));
}
Ok(v)
}
impl AsrParams {
pub const fn new() -> Self {
Self {
language_hint: None,
strategy: SamplingStrategy::BeamSearch {
beam_size: 5,
patience: -1.0,
},
initial_temperature: 0.0,
temperature_increment: 0.2,
max_attempts: 6,
log_prob_threshold: -1.0,
compression_ratio_threshold: 2.4,
no_speech_threshold: 0.6,
no_context: true,
suppress_blank: true,
suppress_non_speech_tokens: false,
initial_prompt: None,
n_threads: 1,
}
}
pub const fn language_hint(&self) -> Option<&Lang> {
self.language_hint.as_ref()
}
pub const fn strategy(&self) -> SamplingStrategy {
self.strategy
}
pub const fn initial_temperature(&self) -> f32 {
self.initial_temperature
}
pub const fn temperature_increment(&self) -> f32 {
self.temperature_increment
}
pub const fn max_attempts(&self) -> u8 {
self.max_attempts
}
pub const fn log_prob_threshold(&self) -> f32 {
self.log_prob_threshold
}
pub const fn compression_ratio_threshold(&self) -> f32 {
self.compression_ratio_threshold
}
pub const fn no_speech_threshold(&self) -> f32 {
self.no_speech_threshold
}
pub const fn no_context(&self) -> bool {
self.no_context
}
pub const fn suppress_blank(&self) -> bool {
self.suppress_blank
}
pub const fn suppress_non_speech_tokens(&self) -> bool {
self.suppress_non_speech_tokens
}
pub const fn initial_prompt(&self) -> Option<&SmolStr> {
self.initial_prompt.as_ref()
}
pub const fn n_threads(&self) -> i32 {
self.n_threads
}
pub fn set_language_hint(&mut self, value: Option<Lang>) {
self.language_hint = value;
}
pub const fn set_strategy(&mut self, value: SamplingStrategy) {
self.strategy = value;
}
pub const fn set_initial_temperature(&mut self, value: f32) {
self.initial_temperature = value;
}
pub const fn set_temperature_increment(&mut self, value: f32) {
self.temperature_increment = value;
}
pub const fn set_max_attempts(&mut self, value: u8) {
assert!(
value > 0,
"max_attempts must be > 0 (got 0); use 1 for a single attempt with no retries"
);
self.max_attempts = value;
}
pub const fn set_log_prob_threshold(&mut self, value: f32) {
self.log_prob_threshold = value;
}
pub const fn set_compression_ratio_threshold(&mut self, value: f32) {
self.compression_ratio_threshold = value;
}
pub const fn set_no_speech_threshold(&mut self, value: f32) {
self.no_speech_threshold = value;
}
pub const fn set_no_context(&mut self, value: bool) {
self.no_context = value;
}
pub const fn set_suppress_blank(&mut self, value: bool) {
self.suppress_blank = value;
}
pub const fn set_suppress_non_speech_tokens(&mut self, value: bool) {
self.suppress_non_speech_tokens = value;
}
pub fn set_initial_prompt(&mut self, value: Option<SmolStr>) {
self.initial_prompt = value;
}
pub const fn set_n_threads(&mut self, value: i32) {
assert!(
value >= 1,
"n_threads must be >= 1; whisper.cpp would underflow / abort otherwise"
);
self.n_threads = value;
}
pub fn with_language_hint(mut self, value: Option<Lang>) -> Self {
self.language_hint = value;
self
}
pub const fn with_strategy(mut self, value: SamplingStrategy) -> Self {
self.strategy = value;
self
}
pub const fn with_initial_temperature(mut self, value: f32) -> Self {
self.initial_temperature = value;
self
}
pub const fn with_temperature_increment(mut self, value: f32) -> Self {
self.temperature_increment = value;
self
}
pub const fn with_max_attempts(mut self, value: u8) -> Self {
assert!(
value > 0,
"max_attempts must be > 0 (got 0); use 1 for a single attempt with no retries"
);
self.max_attempts = value;
self
}
pub const fn with_log_prob_threshold(mut self, value: f32) -> Self {
self.log_prob_threshold = value;
self
}
pub const fn with_compression_ratio_threshold(mut self, value: f32) -> Self {
self.compression_ratio_threshold = value;
self
}
pub const fn with_no_speech_threshold(mut self, value: f32) -> Self {
self.no_speech_threshold = value;
self
}
pub const fn with_no_context(mut self, value: bool) -> Self {
self.no_context = value;
self
}
pub const fn with_suppress_blank(mut self, value: bool) -> Self {
self.suppress_blank = value;
self
}
pub const fn with_suppress_non_speech_tokens(mut self, value: bool) -> Self {
self.suppress_non_speech_tokens = value;
self
}
pub fn with_initial_prompt(mut self, value: Option<SmolStr>) -> Self {
self.initial_prompt = value;
self
}
pub const fn with_n_threads(mut self, value: i32) -> Self {
assert!(
value >= 1,
"n_threads must be >= 1; whisper.cpp would underflow / abort otherwise"
);
self.n_threads = value;
self
}
}
impl Default for AsrParams {
fn default() -> Self {
Self::new()
}
}
#[derive(Copy, Clone, Debug)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
#[cfg_attr(feature = "serde", serde(rename_all = "snake_case"))]
pub enum SamplingStrategy {
Greedy {
best_of: i32,
},
BeamSearch {
beam_size: i32,
patience: f32,
},
}
impl Default for SamplingStrategy {
fn default() -> Self {
Self::BeamSearch {
beam_size: 5,
patience: -1.0,
}
}
}
#[derive(Clone, Debug)]
pub struct AsrResult {
text: SmolStr,
language: Lang,
avg_logprob: f32,
no_speech_prob: f32,
temperature: f32,
runs: Vec<crate::align::Run>,
}
impl AsrResult {
pub fn new(
text: SmolStr,
language: Lang,
avg_logprob: f32,
no_speech_prob: f32,
temperature: f32,
) -> Self {
Self {
text,
language,
avg_logprob,
no_speech_prob,
temperature,
runs: Vec::new(),
}
}
pub fn text(&self) -> &SmolStr {
&self.text
}
pub fn language(&self) -> &Lang {
&self.language
}
pub const fn avg_logprob(&self) -> f32 {
self.avg_logprob
}
pub const fn no_speech_prob(&self) -> f32 {
self.no_speech_prob
}
pub const fn temperature(&self) -> f32 {
self.temperature
}
pub fn runs(&self) -> &[crate::align::Run] {
&self.runs
}
#[must_use]
pub fn with_runs(mut self, runs: Vec<crate::align::Run>) -> Self {
self.runs = runs;
self
}
pub fn set_runs(&mut self, runs: Vec<crate::align::Run>) {
self.runs = runs;
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum AlignmentUnit {
Whole,
Run(usize),
}
#[derive(Clone, Debug)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize), serde(transparent))]
pub struct AlignedWords(
#[cfg_attr(
feature = "serde",
serde(deserialize_with = "deserialize_aligned_words")
)]
Vec<Word>,
);
#[cfg(feature = "serde")]
fn deserialize_aligned_words<'de, D>(deserializer: D) -> Result<Vec<Word>, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::Error as _;
let words = Vec::<Word>::deserialize(deserializer)?;
AlignedWords::new(words)
.map(AlignedWords::into_words)
.ok_or_else(|| D::Error::custom("aligned words are never empty"))
}
impl AlignedWords {
#[must_use]
pub fn new(mut words: Vec<Word>) -> Option<Self> {
if words.is_empty() {
None
} else {
sort_words_by_pts(&mut words);
Some(Self(words))
}
}
#[must_use]
pub fn words(&self) -> &[Word] {
&self.0
}
#[must_use]
pub fn into_words(self) -> Vec<Word> {
self.0
}
pub(crate) fn map(self, f: impl FnMut(Word) -> Word) -> Self {
Self(self.0.into_iter().map(f).collect())
}
}
#[derive(Clone, Debug)]
#[cfg_attr(
feature = "serde",
derive(Serialize, Deserialize),
serde(rename_all = "snake_case")
)]
pub enum UnitAlignment {
Aligned(AlignedWords),
Unaligned(UnalignedCause),
}
impl UnitAlignment {
#[must_use]
pub fn words(&self) -> &[Word] {
match self {
Self::Aligned(words) => words.words(),
Self::Unaligned(_) => &[],
}
}
#[must_use]
pub const fn cause(&self) -> Option<&UnalignedCause> {
match self {
Self::Aligned(_) => None,
Self::Unaligned(cause) => Some(cause),
}
}
pub(crate) fn from_words(words: Vec<Word>) -> Self {
AlignedWords::new(words).map_or(
Self::Unaligned(UnalignedCause::NoSurvivingWords),
Self::Aligned,
)
}
}
#[derive(Debug, Default)]
pub(crate) struct Abandoned(std::sync::Mutex<Vec<(ChunkId, NonZeroU64)>>);
impl Abandoned {
fn report(&self, chunk_id: ChunkId, ticket: NonZeroU64) {
self
.0
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.push((chunk_id, ticket));
}
pub(crate) fn take(&self) -> Vec<(ChunkId, NonZeroU64)> {
core::mem::take(
&mut *self
.0
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner),
)
}
}
#[derive(Debug)]
pub(crate) struct AlignmentTicket {
id: NonZeroU64,
chunk_id: ChunkId,
transcriber: NonZeroU64,
issuer: Option<Arc<Abandoned>>,
}
impl AlignmentTicket {
pub(crate) fn mint(
chunk_id: ChunkId,
transcriber: NonZeroU64,
issuer: Option<Arc<Abandoned>>,
) -> Self {
static COUNTER: AtomicU64 = AtomicU64::new(1);
let raw = COUNTER.fetch_add(1, Ordering::Relaxed);
Self {
id: NonZeroU64::new(raw).expect("AlignmentTicket counter overflowed u64"),
chunk_id,
transcriber,
issuer,
}
}
pub(crate) fn settle(mut self) {
self.issuer = None;
}
pub(crate) const fn chunk_id(&self) -> ChunkId {
self.chunk_id
}
pub(crate) const fn id(&self) -> NonZeroU64 {
self.id
}
pub(crate) const fn transcriber(&self) -> NonZeroU64 {
self.transcriber
}
}
impl Drop for AlignmentTicket {
fn drop(&mut self) {
if let Some(issuer) = self.issuer.take() {
issuer.report(self.chunk_id, self.id);
}
}
}
#[derive(Debug)]
#[must_use = "a unit is answered only by an aligner consuming its job, or by skipping it"]
pub struct UnitJob {
ticket: NonZeroU64,
unit: AlignmentUnit,
text: SmolStr,
language: Lang,
samples: Arc<[f32]>,
window: core::ops::Range<usize>,
#[cfg(any(feature = "alignment", feature = "emissions"))]
place: UnitPlace,
}
#[cfg(any(feature = "alignment", feature = "emissions"))]
#[derive(Debug)]
pub(crate) struct UnitPlace {
pub(crate) first_sample_in_stream: u64,
pub(crate) sub_segments: Vec<TimeRange>,
pub(crate) output_tb: mediatime::Timebase,
pub(crate) base_pts_out_anchor: i64,
}
impl UnitJob {
#[must_use]
pub const fn unit(&self) -> AlignmentUnit {
self.unit
}
#[must_use]
pub fn text(&self) -> &str {
&self.text
}
#[must_use]
pub const fn language(&self) -> &Lang {
&self.language
}
#[must_use]
pub fn samples(&self) -> &[f32] {
&self.samples[self.window.clone()]
}
pub fn skip(self) -> UnitOutcome {
self.answer(UnitAlignment::Unaligned(UnalignedCause::Skipped))
}
pub(crate) fn answer(self, alignment: UnitAlignment) -> UnitOutcome {
let alignment = match (self.unit, alignment) {
(AlignmentUnit::Run(_), UnitAlignment::Aligned(words)) => {
let language = &self.language;
UnitAlignment::Aligned(words.map(|word| word.with_language(Some(language.clone()))))
}
(_, alignment) => alignment,
};
UnitOutcome {
ticket: self.ticket,
unit: self.unit,
alignment,
}
}
pub(crate) const fn ticket(&self) -> NonZeroU64 {
self.ticket
}
#[cfg(any(feature = "alignment", feature = "emissions"))]
pub(crate) const fn place(&self) -> &UnitPlace {
&self.place
}
#[cfg(feature = "alignment")]
pub(crate) fn samples_to_output_range(&self) -> Arc<dyn Fn(u64, u64) -> TimeRange + Send + Sync> {
crate::core::buffer::SampleBuffer::samples_to_output_range_fn_at(
self.place.output_tb,
self.place.base_pts_out_anchor,
)
}
}
#[derive(Debug)]
pub struct UnitOutcome {
ticket: NonZeroU64,
unit: AlignmentUnit,
alignment: UnitAlignment,
}
impl UnitOutcome {
#[must_use]
pub const fn unit(&self) -> AlignmentUnit {
self.unit
}
#[must_use]
pub const fn alignment(&self) -> &UnitAlignment {
&self.alignment
}
}
#[cfg(any(feature = "alignment", feature = "emissions"))]
#[derive(Debug)]
pub(crate) struct ChunkContext {
pub(crate) first_sample: u64,
pub(crate) sub_segments_samples: Vec<(u64, u64)>,
pub(crate) output_tb: mediatime::Timebase,
pub(crate) base_pts_out_anchor: i64,
}
#[derive(Debug)]
#[must_use = "an alignment command is answered only through its request"]
pub struct AlignmentRequest {
ticket: AlignmentTicket,
samples: Arc<[f32]>,
sub_segments: Vec<TimeRange>,
text: SmolStr,
language: Lang,
runs: Vec<crate::align::Run>,
units_taken: bool,
#[cfg(any(feature = "alignment", feature = "emissions"))]
context: ChunkContext,
}
impl AlignmentRequest {
pub(crate) fn new(
ticket: AlignmentTicket,
samples: Arc<[f32]>,
sub_segments: Vec<TimeRange>,
text: SmolStr,
language: Lang,
runs: Vec<crate::align::Run>,
#[cfg(any(feature = "alignment", feature = "emissions"))] context: ChunkContext,
) -> Self {
Self {
ticket,
samples,
sub_segments,
text,
language,
runs,
units_taken: false,
#[cfg(any(feature = "alignment", feature = "emissions"))]
context,
}
}
#[must_use]
pub const fn chunk_id(&self) -> ChunkId {
self.ticket.chunk_id()
}
#[must_use]
pub const fn samples(&self) -> &Arc<[f32]> {
&self.samples
}
#[must_use]
pub fn sub_segments(&self) -> &[TimeRange] {
&self.sub_segments
}
#[must_use]
pub const fn text(&self) -> &SmolStr {
&self.text
}
#[must_use]
pub const fn language(&self) -> &Lang {
&self.language
}
#[must_use]
pub fn runs(&self) -> &[crate::align::Run] {
&self.runs
}
#[must_use]
pub fn units(&self) -> Vec<AlignmentUnit> {
alignment_units(self.runs.len())
}
pub fn take_units(&mut self) -> Vec<UnitJob> {
if core::mem::replace(&mut self.units_taken, true) {
return Vec::new();
}
alignment_units(self.runs.len())
.into_iter()
.map(|unit| self.unit_job(unit))
.collect()
}
fn unit_job(&self, unit: AlignmentUnit) -> UnitJob {
let (text, language, window) = match unit {
AlignmentUnit::Whole => (
self.text.clone(),
self.language.clone(),
0..self.samples.len(),
),
AlignmentUnit::Run(index) => {
let run = &self.runs[index];
let (lo, hi) = run_audio_slice(run, self.samples.len(), 0);
(SmolStr::new(run.text()), run.language().clone(), lo..hi)
}
};
#[cfg(any(feature = "alignment", feature = "emissions"))]
let place = {
let chunk_local = self.chunk_local_sub_segments();
let sub_segments = match unit {
AlignmentUnit::Whole => chunk_local,
AlignmentUnit::Run(_) => {
clip_sub_segments(&chunk_local, window.start, window.end, &language).unwrap_or_default()
}
};
UnitPlace {
first_sample_in_stream: self
.context
.first_sample
.saturating_add(window.start as u64),
sub_segments,
output_tb: self.context.output_tb,
base_pts_out_anchor: self.context.base_pts_out_anchor,
}
};
UnitJob {
ticket: self.ticket.id(),
unit,
text,
language,
samples: self.samples.clone(),
window,
#[cfg(any(feature = "alignment", feature = "emissions"))]
place,
}
}
pub fn aligned(
self,
outcomes: Vec<UnitOutcome>,
) -> Result<AlignmentCompletion, UnaccountedOutcomes> {
let expected = self.units();
let own = |outcome: &UnitOutcome| outcome.ticket == self.ticket.id();
let accounted = outcomes.len() == expected.len()
&& outcomes
.iter()
.zip(&expected)
.all(|(outcome, unit)| own(outcome) && outcome.unit == *unit);
if !accounted {
let error = crate::types::UnaccountedAlignment::new(
self.chunk_id(),
expected,
outcomes.iter().map(UnitOutcome::unit).collect(),
outcomes.iter().filter(|outcome| !own(outcome)).count(),
);
return Err(UnaccountedOutcomes(Box::new(Refused {
error,
request: self,
outcomes,
})));
}
let mut alignments = outcomes.into_iter().map(|outcome| outcome.alignment);
let report = match (self.runs.is_empty(), alignments.next()) {
(true, Some(whole)) => AlignmentReport::Whole(whole),
(_, first) => AlignmentReport::Runs(first.into_iter().chain(alignments).collect()),
};
Ok(AlignmentCompletion {
ticket: self.ticket,
answer: Answer::Aligned(report),
})
}
pub fn align_units<E: crate::types::IntoWorkFailure>(
mut self,
mut align: impl FnMut(UnitJob) -> Result<UnitOutcome, E>,
) -> AlignmentCompletion {
let jobs = self.take_units();
let mut outcomes = Vec::with_capacity(jobs.len());
for job in jobs {
let language = job.language().clone();
match std::panic::catch_unwind(core::panic::AssertUnwindSafe(|| align(job))) {
Ok(Ok(outcome)) => outcomes.push(outcome),
Ok(Err(error)) => return self.failed(error.into_work_failure(&language)),
Err(panic) => return self.failed(panic_failure(panic.as_ref(), language)),
}
}
self.aligned(outcomes).unwrap_or_else(|refused| {
let message = smol_str::format_smolstr!("{refused}");
let (request, _) = refused.into_parts();
let language = request.language().clone();
request.failed(WorkFailure::Alignment(
crate::types::AlignmentError::Tokenization(crate::types::AlignmentFailure::new(
message, language,
)),
))
})
}
pub fn failed(self, failure: WorkFailure) -> AlignmentCompletion {
AlignmentCompletion {
ticket: self.ticket,
answer: Answer::Failed(failure),
}
}
#[cfg(feature = "alignment")]
pub(crate) const fn chunk_first_sample(&self) -> u64 {
self.context.first_sample
}
#[cfg(any(feature = "alignment", feature = "emissions"))]
pub(crate) fn chunk_local_sub_segments(&self) -> Vec<TimeRange> {
use crate::time::{ANALYSIS_TIMEBASE, sample_pts};
let first = self.context.first_sample;
let local = |sample: u64| sample_pts(sample.saturating_sub(first), ANALYSIS_TIMEBASE, 0);
self
.context
.sub_segments_samples
.iter()
.map(|&(start, end)| TimeRange::new(local(start), local(end), ANALYSIS_TIMEBASE))
.collect()
}
#[cfg(feature = "alignment")]
pub(crate) fn samples_to_output_range(&self) -> Arc<dyn Fn(u64, u64) -> TimeRange + Send + Sync> {
crate::core::buffer::SampleBuffer::samples_to_output_range_fn_at(
self.context.output_tb,
self.context.base_pts_out_anchor,
)
}
#[cfg(test)]
pub(crate) fn for_test(
chunk_id: ChunkId,
transcriber: NonZeroU64,
samples: Arc<[f32]>,
text: SmolStr,
language: Lang,
runs: Vec<crate::align::Run>,
) -> Self {
Self::new(
AlignmentTicket::mint(chunk_id, transcriber, None),
samples,
Vec::new(),
text,
language,
runs,
#[cfg(any(feature = "alignment", feature = "emissions"))]
ChunkContext {
first_sample: 0,
sub_segments_samples: Vec::new(),
output_tb: mediatime::Timebase::new(
1,
core::num::NonZeroI32::new(16_000).expect("16000 != 0"),
),
base_pts_out_anchor: 0,
},
)
}
}
pub(crate) fn run_audio_slice(
run: &crate::align::Run,
samples_len: usize,
_chunk_first_sample_in_stream: u64,
) -> (usize, usize) {
use crate::align::BoundsSource;
if matches!(run.bounds_source(), BoundsSource::Wholeclip) {
return (0, samples_len);
}
let t0 = run.audio_t0_ms();
let t1 = run.audio_t1_ms();
if t0 < 0 || t1 <= t0 {
return (0, 0);
}
let lo_u64 = (t0 as u64).saturating_mul(16);
let hi_u64 = (t1 as u64).saturating_mul(16);
if lo_u64 >= samples_len as u64 {
eprintln!(
"asry alignment Run bounds appear out-of-chunk: \
audio_t0_ms={t0} audio_t1_ms={t1} chunk_samples_len={samples_len}; \
check your AsrSource — Run::audio_t*_ms must be chunk-local ms, not stream-absolute"
);
return (samples_len, samples_len);
}
let lo = lo_u64.min(samples_len as u64) as usize;
let hi = hi_u64.min(samples_len as u64) as usize;
if hi <= lo {
return (lo, lo);
}
(lo, hi)
}
#[cfg(any(feature = "alignment", feature = "emissions"))]
pub(crate) fn clip_sub_segments(
subs: &[TimeRange],
slice_lo: usize,
slice_hi: usize,
language: &Lang,
) -> Result<Vec<TimeRange>, WorkFailure> {
use core::num::NonZeroI32;
let tb = mediatime::Timebase::new(1, NonZeroI32::new(16_000).unwrap());
let mut out = Vec::with_capacity(subs.len());
let lo_i = slice_lo as i64;
let hi_i = slice_hi as i64;
for sub in subs {
let actual_tb = sub.timebase();
if actual_tb.num() != 1 || actual_tb.den().get() != 16_000 {
return Err(WorkFailure::Alignment(
crate::types::AlignmentError::ModelInference(crate::types::AlignmentFailure::new(
smol_str::format_smolstr!(
"sub_segments must be in 1/16000 (chunk-local sample-index) timebase; got \
{}/{}. Convert via `Transcriber::chunk_first_sample` + a 1/16000 timebase \
before passing to the aligner.",
actual_tb.num(),
actual_tb.den().get(),
),
language.clone(),
)),
));
}
let s = sub.start_pts().max(lo_i);
let e = sub.end_pts().min(hi_i);
if e > s {
out.push(TimeRange::new(s - lo_i, e - lo_i, tb));
}
}
Ok(out)
}
pub(crate) fn panic_failure(panic: &(dyn core::any::Any + Send), language: Lang) -> WorkFailure {
let message = panic
.downcast_ref::<&str>()
.copied()
.or_else(|| panic.downcast_ref::<String>().map(String::as_str))
.unwrap_or("a panic with no message");
WorkFailure::Alignment(crate::types::AlignmentError::ModelInference(
crate::types::AlignmentFailure::new(
smol_str::format_smolstr!("the alignment job panicked: {message}"),
language,
),
))
}
#[derive(Debug, thiserror::Error)]
#[error("{}", .0.error)]
pub struct UnaccountedOutcomes(Box<Refused>);
#[derive(Debug)]
struct Refused {
error: crate::types::UnaccountedAlignment,
request: AlignmentRequest,
outcomes: Vec<UnitOutcome>,
}
impl UnaccountedOutcomes {
#[must_use]
pub fn error(&self) -> &crate::types::UnaccountedAlignment {
&self.0.error
}
pub fn into_parts(self) -> (AlignmentRequest, Vec<UnitOutcome>) {
let Refused {
request, outcomes, ..
} = *self.0;
(request, outcomes)
}
}
#[derive(Debug)]
pub(crate) enum Answer {
Aligned(AlignmentReport),
Failed(WorkFailure),
}
#[derive(Debug)]
#[must_use = "a completion answers its command only when Transcriber::complete takes it"]
pub struct AlignmentCompletion {
ticket: AlignmentTicket,
answer: Answer,
}
impl AlignmentCompletion {
#[must_use]
pub const fn chunk_id(&self) -> ChunkId {
self.ticket.chunk_id()
}
#[must_use]
pub const fn report(&self) -> Option<&AlignmentReport> {
match &self.answer {
Answer::Aligned(report) => Some(report),
Answer::Failed(_) => None,
}
}
#[must_use]
pub const fn failure(&self) -> Option<&WorkFailure> {
match &self.answer {
Answer::Aligned(_) => None,
Answer::Failed(failure) => Some(failure),
}
}
pub(crate) const fn ticket(&self) -> &AlignmentTicket {
&self.ticket
}
pub(crate) fn into_parts(self) -> (AlignmentTicket, Answer) {
(self.ticket, self.answer)
}
}
#[derive(Debug, thiserror::Error)]
#[error("{}", .0.error)]
pub struct RefusedCompletion(Box<Refusal>);
#[derive(Debug)]
struct Refusal {
error: crate::types::TranscriberError,
completion: AlignmentCompletion,
}
impl RefusedCompletion {
pub(crate) fn new(
error: crate::types::TranscriberError,
completion: AlignmentCompletion,
) -> Self {
Self(Box::new(Refusal { error, completion }))
}
#[must_use]
pub fn error(&self) -> &crate::types::TranscriberError {
&self.0.error
}
pub fn into_completion(self) -> AlignmentCompletion {
self.0.completion
}
#[must_use = "the completion is dropped; keep the refusal at least"]
pub fn discard_completion(self) -> crate::types::TranscriberError {
self.0.error
}
}
#[derive(Clone, Debug)]
#[cfg_attr(
feature = "serde",
derive(Serialize, Deserialize),
serde(rename_all = "snake_case")
)]
pub enum AlignmentReport {
NotAttempted,
Whole(UnitAlignment),
Runs(Vec<UnitAlignment>),
}
impl AlignmentReport {
fn outcomes(&self) -> &[UnitAlignment] {
match self {
Self::NotAttempted => &[],
Self::Whole(outcome) => core::slice::from_ref(outcome),
Self::Runs(outcomes) => outcomes,
}
}
pub fn units(&self) -> impl ExactSizeIterator<Item = (AlignmentUnit, &UnitAlignment)> + '_ {
let whole = matches!(self, Self::Whole(_));
self
.outcomes()
.iter()
.enumerate()
.map(move |(index, outcome)| {
let unit = if whole {
AlignmentUnit::Whole
} else {
AlignmentUnit::Run(index)
};
(unit, outcome)
})
}
pub fn unaligned(&self) -> impl Iterator<Item = (AlignmentUnit, &UnalignedCause)> + '_ {
self
.units()
.filter_map(|(unit, outcome)| outcome.cause().map(|cause| (unit, cause)))
}
pub fn words(&self) -> impl ExactSizeIterator<Item = &Word> + '_ {
TimeOrdered::new(self.outcomes())
}
}
struct TimeOrdered<'a> {
units: smallvec::SmallVec<[&'a [Word]; 4]>,
remaining: usize,
}
impl<'a> TimeOrdered<'a> {
fn new(outcomes: &'a [UnitAlignment]) -> Self {
let units: smallvec::SmallVec<[&'a [Word]; 4]> = outcomes
.iter()
.map(UnitAlignment::words)
.filter(|words| !words.is_empty())
.collect();
let remaining = units.iter().map(|words| words.len()).sum();
Self { units, remaining }
}
}
impl<'a> Iterator for TimeOrdered<'a> {
type Item = &'a Word;
fn next(&mut self) -> Option<&'a Word> {
let key = |word: &Word| {
let range = word.range();
(range.start_pts(), range.end_pts())
};
let mut first: Option<usize> = None;
for (index, words) in self.units.iter().enumerate() {
let Some(word) = words.first() else {
continue;
};
if first.is_none_or(|best| key(word) < key(&self.units[best][0])) {
first = Some(index);
}
}
let index = first?;
let words: &'a [Word] = self.units[index];
let (word, rest) = words.split_first()?;
self.units[index] = rest;
self.remaining -= 1;
Some(word)
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.remaining, Some(self.remaining))
}
}
impl ExactSizeIterator for TimeOrdered<'_> {}
pub(crate) fn alignment_units(runs: usize) -> Vec<AlignmentUnit> {
if runs == 0 {
vec![AlignmentUnit::Whole]
} else {
(0..runs).map(AlignmentUnit::Run).collect()
}
}
pub(crate) fn sort_words_by_pts(words: &mut [Word]) {
words.sort_by_key(|word| {
let range = word.range();
(range.start_pts(), range.end_pts())
});
}
#[derive(Clone, Debug)]
#[cfg_attr(
feature = "serde",
derive(Serialize, Deserialize),
serde(rename_all = "snake_case")
)]
#[non_exhaustive]
pub enum UnalignedCause {
Skipped,
Refused,
NoAlignableText,
NoSurvivingWords,
Failed(crate::types::AlignmentError),
}
#[derive(Debug)]
pub enum Command {
Asr {
chunk_id: ChunkId,
samples: Arc<[f32]>,
sample_rate: u32,
params: AsrParams,
},
Alignment(AlignmentRequest),
}
#[derive(Clone, Debug, Default)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct AsrParamsOverride {
#[cfg_attr(
feature = "serde",
serde(
default,
skip_serializing_if = "Option::is_none",
deserialize_with = "deserialize_double_option_lang"
)
)]
language_hint: Option<Option<Lang>>,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Option::is_none")
)]
strategy: Option<SamplingStrategy>,
#[cfg_attr(
feature = "serde",
serde(default, skip_serializing_if = "Option::is_none")
)]
initial_temperature: Option<f32>,
#[cfg_attr(
feature = "serde",
serde(
default,
skip_serializing_if = "Option::is_none",
deserialize_with = "deserialize_double_option_smolstr"
)
)]
initial_prompt: Option<Option<SmolStr>>,
}
#[cfg(feature = "serde")]
fn deserialize_double_option_lang<'de, D>(d: D) -> Result<Option<Option<Lang>>, D::Error>
where
D: serde::Deserializer<'de>,
{
Ok(Some(Option::<Lang>::deserialize(d)?))
}
#[cfg(feature = "serde")]
fn deserialize_double_option_smolstr<'de, D>(d: D) -> Result<Option<Option<SmolStr>>, D::Error>
where
D: serde::Deserializer<'de>,
{
Ok(Some(Option::<SmolStr>::deserialize(d)?))
}
impl AsrParamsOverride {
pub const fn new() -> Self {
Self {
language_hint: None,
strategy: None,
initial_temperature: None,
initial_prompt: None,
}
}
pub fn apply_to(&self, base: &AsrParams) -> AsrParams {
let mut out = base.clone();
if let Some(opt_lang) = &self.language_hint {
out.set_language_hint(opt_lang.clone());
}
if let Some(strategy) = self.strategy {
out.set_strategy(strategy);
}
if let Some(t) = self.initial_temperature {
out.set_initial_temperature(t);
}
if let Some(prompt) = &self.initial_prompt {
out.set_initial_prompt(prompt.clone());
}
out
}
pub const fn language_hint(&self) -> Option<&Option<Lang>> {
self.language_hint.as_ref()
}
pub const fn strategy(&self) -> Option<SamplingStrategy> {
self.strategy
}
pub const fn initial_temperature(&self) -> Option<f32> {
self.initial_temperature
}
pub const fn initial_prompt(&self) -> Option<&Option<SmolStr>> {
self.initial_prompt.as_ref()
}
pub fn set_language_hint(&mut self, value: Option<Option<Lang>>) {
self.language_hint = value;
}
pub const fn set_strategy(&mut self, value: Option<SamplingStrategy>) {
self.strategy = value;
}
pub const fn set_initial_temperature(&mut self, value: Option<f32>) {
self.initial_temperature = value;
}
pub fn set_initial_prompt(&mut self, value: Option<Option<SmolStr>>) {
self.initial_prompt = value;
}
pub fn with_language_hint(mut self, value: Option<Option<Lang>>) -> Self {
self.language_hint = value;
self
}
pub const fn with_strategy(mut self, value: Option<SamplingStrategy>) -> Self {
self.strategy = value;
self
}
pub const fn with_initial_temperature(mut self, value: Option<f32>) -> Self {
self.initial_temperature = value;
self
}
pub fn with_initial_prompt(mut self, value: Option<Option<SmolStr>>) -> Self {
self.initial_prompt = value;
self
}
}
#[allow(dead_code)] pub(crate) type ChunkAudio = Arc<[f32]>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_report_gives_each_unit_exactly_one_outcome() {
use core::num::NonZeroI32;
use crate::types::{AlignmentError, AlignmentFailure};
let word = |text: &str, start: i64| {
Word::new(
SmolStr::new(text),
TimeRange::new(
start,
start + 10,
mediatime::Timebase::new(1, NonZeroI32::new(1_000).expect("1000 != 0")),
),
0.9,
)
};
assert!(AlignedWords::new(Vec::new()).is_none(), "never empty");
assert!(matches!(
UnitAlignment::from_words(Vec::new()),
UnitAlignment::Unaligned(UnalignedCause::NoSurvivingWords)
));
let failed =
AlignmentError::NoAlignmentPath(AlignmentFailure::new(SmolStr::new("too short"), Lang::En));
let report = AlignmentReport::Runs(vec![
UnitAlignment::from_words(vec![word("b", 20)]),
UnitAlignment::Unaligned(UnalignedCause::Skipped),
UnitAlignment::Unaligned(UnalignedCause::Refused),
UnitAlignment::Unaligned(UnalignedCause::NoAlignableText),
UnitAlignment::Unaligned(UnalignedCause::NoSurvivingWords),
UnitAlignment::Unaligned(UnalignedCause::Failed(failed)),
UnitAlignment::from_words(vec![word("a", 0)]),
]);
assert_eq!(
report.units().map(|(unit, _)| unit).collect::<Vec<_>>(),
(0..7).map(AlignmentUnit::Run).collect::<Vec<_>>()
);
let unaligned: Vec<AlignmentUnit> = report.unaligned().map(|(unit, _)| unit).collect();
assert_eq!(
unaligned,
(1..6).map(AlignmentUnit::Run).collect::<Vec<_>>()
);
assert!(matches!(
report.unaligned().last(),
Some((
_,
UnalignedCause::Failed(AlignmentError::NoAlignmentPath(_))
))
));
assert_eq!(
report.words().map(Word::text).collect::<Vec<_>>(),
["a", "b"],
"time order"
);
assert_eq!(report.words().len(), 2);
let report = AlignmentReport::Runs(vec![
UnitAlignment::from_words(vec![word("late", 30), word("tie-0", 10)]),
UnitAlignment::Unaligned(UnalignedCause::NoAlignableText),
UnitAlignment::from_words(vec![word("tie-2", 10), word("first", 0)]),
]);
assert_eq!(
report.words().map(Word::text).collect::<Vec<_>>(),
["first", "tie-0", "tie-2", "late"]
);
assert_eq!(report.words().len(), 4);
assert_eq!(AlignmentReport::NotAttempted.words().len(), 0);
assert_eq!(AlignmentReport::NotAttempted.units().len(), 0);
let whole = AlignmentReport::Whole(UnitAlignment::Unaligned(UnalignedCause::Refused));
assert_eq!(
whole.units().map(|(unit, _)| unit).collect::<Vec<_>>(),
[AlignmentUnit::Whole]
);
}
#[test]
fn a_run_outcome_carries_the_run_language() {
let transcriber = NonZeroU64::new(1).expect("1 != 0");
let tb = mediatime::Timebase::new(1, core::num::NonZeroI32::new(16_000).expect("16000 != 0"));
let words = || {
AlignedWords::new(vec![Word::new(
SmolStr::new("hello"),
TimeRange::new(0, 10, tb),
0.9,
)])
.expect("a word")
};
let run = |language: Lang, text: &str| {
crate::align::Run::new(
language,
SmolStr::new(text),
0,
50,
0,
crate::align::BoundsSource::Segment,
)
};
let mut request = AlignmentRequest::for_test(
ChunkId::from_raw(1),
transcriber,
Arc::from(vec![0.0_f32; 1_600]),
SmolStr::new("hello 안녕"),
Lang::En,
vec![run(Lang::En, "hello"), run(Lang::Ko, " 안녕")],
);
let outcomes: Vec<UnitOutcome> = request
.take_units()
.into_iter()
.map(|job| job.answer(UnitAlignment::Aligned(words())))
.collect();
let languages: Vec<Option<&Lang>> = outcomes
.iter()
.map(|outcome| outcome.alignment().words()[0].language())
.collect();
assert_eq!(languages, [Some(&Lang::En), Some(&Lang::Ko)]);
let mut whole = AlignmentRequest::for_test(
ChunkId::from_raw(2),
transcriber,
Arc::from(vec![0.0_f32; 1_600]),
SmolStr::new("hello"),
Lang::En,
Vec::new(),
);
let outcome = whole
.take_units()
.pop()
.expect("the whole text's job")
.answer(UnitAlignment::Aligned(words()));
assert_eq!(outcome.alignment().words()[0].language(), None);
}
#[test]
fn a_request_hands_out_one_job_per_unit_and_answers_once() {
let transcriber = NonZeroU64::new(1).expect("1 != 0");
let audio: Arc<[f32]> = (0..1_600).map(|i| i as f32).collect();
let run = |language: Lang, text: &str, t0_ms: i64, t1_ms: i64| {
crate::align::Run::new(
language,
SmolStr::new(text),
t0_ms,
t1_ms,
0,
crate::align::BoundsSource::Segment,
)
};
let request = |runs| {
AlignmentRequest::for_test(
ChunkId::from_raw(3),
transcriber,
audio.clone(),
SmolStr::new("hello 세계"),
Lang::En,
runs,
)
};
for (runs, units, texts, languages, windows) in [
(
Vec::new(),
vec![AlignmentUnit::Whole],
vec!["hello 세계"],
vec![Lang::En],
vec![(0, 1_600)],
),
(
vec![
run(Lang::En, "hello", 0, 50),
run(Lang::Ko, " 세계", 50, 100),
],
vec![AlignmentUnit::Run(0), AlignmentUnit::Run(1)],
vec!["hello", " 세계"],
vec![Lang::En, Lang::Ko],
vec![(0, 800), (800, 1_600)],
),
] {
let mut request = request(runs);
assert_eq!(request.chunk_id(), ChunkId::from_raw(3));
assert_eq!(request.units(), units);
let jobs = request.take_units();
assert!(request.take_units().is_empty(), "the jobs are taken once");
assert_eq!(jobs.iter().map(UnitJob::unit).collect::<Vec<_>>(), units);
assert_eq!(jobs.iter().map(UnitJob::text).collect::<Vec<_>>(), texts);
assert_eq!(
jobs
.iter()
.map(|job| job.language().clone())
.collect::<Vec<_>>(),
languages
);
for (job, (lo, hi)) in jobs.iter().zip(windows) {
assert_eq!(job.samples(), &audio[lo..hi], "{:?}", job.unit());
}
let outcomes: Vec<UnitOutcome> = jobs
.into_iter()
.map(|job| job.answer(UnitAlignment::Unaligned(UnalignedCause::NoSurvivingWords)))
.collect();
assert_eq!(
outcomes.iter().map(UnitOutcome::unit).collect::<Vec<_>>(),
units
);
let completion = request.aligned(outcomes).expect("its own units, in order");
assert_eq!(completion.chunk_id(), ChunkId::from_raw(3));
assert!(completion.failure().is_none());
let report = completion.report().expect("aligned");
assert_eq!(
report.units().map(|(unit, _)| unit).collect::<Vec<_>>(),
units
);
assert!(matches!(
(units.len(), report),
(1, AlignmentReport::Whole(_)) | (2, AlignmentReport::Runs(_))
));
}
let failure = WorkFailure::LanguageUnsupported(
crate::types::LanguageUnsupportedForAlignment::new(Lang::En),
);
let completion = request(Vec::new()).failed(failure);
assert!(completion.report().is_none());
assert!(matches!(
completion.failure(),
Some(WorkFailure::LanguageUnsupported(_))
));
}
#[test]
fn asr_params_defaults_match_spec() {
let p = AsrParams::default();
match p.strategy {
SamplingStrategy::BeamSearch {
beam_size,
patience,
} => {
assert_eq!(beam_size, 5);
assert!((patience - -1.0).abs() < 1e-9);
}
_ => panic!("default should be BeamSearch"),
}
assert!((p.initial_temperature - 0.0).abs() < 1e-9);
assert!((p.temperature_increment - 0.2).abs() < 1e-9);
assert_eq!(p.max_attempts, 6);
assert!(p.no_context);
}
#[cfg(feature = "serde")]
#[test]
fn asr_params_serde_round_trip() {
let mut p = AsrParams::default();
p.set_initial_temperature(0.7);
p.set_max_attempts(3);
let json = serde_json::to_string(&p).expect("serialize");
let back: AsrParams = serde_json::from_str(&json).expect("deserialize");
assert!((back.initial_temperature() - 0.7).abs() < 1e-9);
assert_eq!(back.max_attempts(), 3);
}
#[test]
#[should_panic(expected = "max_attempts must be > 0")]
fn set_max_attempts_zero_panics() {
let mut p = AsrParams::default();
p.set_max_attempts(0);
}
#[test]
#[should_panic(expected = "max_attempts must be > 0")]
fn with_max_attempts_zero_panics() {
let _ = AsrParams::default().with_max_attempts(0);
}
#[test]
#[should_panic(expected = "n_threads must be >= 1")]
fn set_n_threads_zero_panics() {
let mut p = AsrParams::default();
p.set_n_threads(0);
}
#[test]
#[should_panic(expected = "n_threads must be >= 1")]
fn set_n_threads_negative_panics() {
let mut p = AsrParams::default();
p.set_n_threads(-3);
}
#[test]
#[should_panic(expected = "n_threads must be >= 1")]
fn with_n_threads_zero_panics() {
let _ = AsrParams::default().with_n_threads(0);
}
#[cfg(feature = "serde")]
#[test]
fn deserialize_rejects_zero_n_threads() {
let json = r#"{"n_threads": 0}"#;
let res: Result<AsrParams, _> = serde_json::from_str(json);
assert!(res.is_err(), "n_threads=0 must be rejected");
let err = res.err().unwrap().to_string();
assert!(err.contains("n_threads must be >= 1"), "got {err:?}");
}
#[cfg(feature = "serde")]
#[test]
fn deserialize_rejects_negative_n_threads() {
let json = r#"{"n_threads": -2}"#;
let res: Result<AsrParams, _> = serde_json::from_str(json);
assert!(res.is_err(), "n_threads=-2 must be rejected");
}
#[cfg(feature = "serde")]
#[test]
fn deserialize_rejects_zero_max_attempts() {
let json = r#"{"max_attempts": 0}"#;
let res: Result<AsrParams, _> = serde_json::from_str(json);
assert!(res.is_err(), "max_attempts=0 must be rejected");
let err = res.err().unwrap().to_string();
assert!(
err.contains("max_attempts must be > 0"),
"expected diagnostic, got {err:?}"
);
}
#[cfg(feature = "serde")]
#[test]
fn asr_params_serde_empty_yields_defaults() {
let p: AsrParams = serde_json::from_str("{}").expect("deserialize empty");
assert_eq!(
p.initial_temperature(),
AsrParams::default().initial_temperature()
);
assert_eq!(p.max_attempts(), AsrParams::default().max_attempts());
assert_eq!(p.no_context(), AsrParams::default().no_context());
assert!(p.language_hint().is_none());
assert!(p.initial_prompt().is_none());
}
#[cfg(feature = "serde")]
#[test]
fn asr_params_override_serde_absent_means_no_override() {
let ovr: AsrParamsOverride = serde_json::from_str("{}").expect("deserialize empty");
assert!(
ovr.language_hint().is_none(),
"absent field must mean None (no override)"
);
assert!(
ovr.initial_prompt().is_none(),
"absent field must mean None (no override)"
);
}
#[cfg(feature = "serde")]
#[test]
fn asr_params_override_serde_null_means_clear() {
let ovr: AsrParamsOverride =
serde_json::from_str(r#"{"language_hint": null}"#).expect("deserialize null");
match ovr.language_hint() {
Some(None) => {}
other => panic!("JSON null on language_hint must produce Some(None) (clear); got {other:?}"),
}
let ovr: AsrParamsOverride =
serde_json::from_str(r#"{"initial_prompt": null}"#).expect("deserialize null");
match ovr.initial_prompt() {
Some(None) => {}
other => panic!("JSON null on initial_prompt must produce Some(None) (clear); got {other:?}"),
}
}
#[cfg(feature = "serde")]
#[test]
fn asr_params_override_serde_value_means_set() {
let ovr: AsrParamsOverride =
serde_json::from_str(r#"{"language_hint": "EN"}"#).expect("deserialize value");
match ovr.language_hint() {
Some(Some(Lang::En)) => {}
other => panic!("expected Some(Some(Lang::En)); got {other:?}"),
}
let ovr: AsrParamsOverride =
serde_json::from_str(r#"{"initial_prompt": "hint"}"#).expect("deserialize value");
match ovr.initial_prompt() {
Some(Some(s)) if s.as_str() == "hint" => {}
other => panic!("expected Some(Some(\"hint\")); got {other:?}"),
}
}
#[cfg(feature = "serde")]
#[test]
fn asr_params_override_serde_round_trips_three_states() {
let mut ovr_absent = AsrParamsOverride::new();
ovr_absent.set_initial_temperature(Some(0.7)); let json = serde_json::to_string(&ovr_absent).unwrap();
assert!(
!json.contains("language_hint") && !json.contains("initial_prompt"),
"absent fields must skip-serialize; got {json}"
);
let back: AsrParamsOverride = serde_json::from_str(&json).unwrap();
assert!(back.language_hint().is_none());
assert!(back.initial_prompt().is_none());
let ovr_clear = AsrParamsOverride::new()
.with_language_hint(Some(None))
.with_initial_prompt(Some(None));
let json = serde_json::to_string(&ovr_clear).unwrap();
assert!(json.contains("\"language_hint\":null"), "got {json}");
assert!(json.contains("\"initial_prompt\":null"), "got {json}");
let back: AsrParamsOverride = serde_json::from_str(&json).unwrap();
assert!(matches!(back.language_hint(), Some(None)));
assert!(matches!(back.initial_prompt(), Some(None)));
let ovr_set = AsrParamsOverride::new()
.with_language_hint(Some(Some(Lang::En)))
.with_initial_prompt(Some(Some(SmolStr::new("hint"))));
let json = serde_json::to_string(&ovr_set).unwrap();
let back: AsrParamsOverride = serde_json::from_str(&json).unwrap();
assert!(matches!(back.language_hint(), Some(Some(Lang::En))));
assert!(
matches!(back.initial_prompt(), Some(Some(s)) if s.as_str() == "hint"),
"got {:?}",
back.initial_prompt()
);
}
#[cfg(feature = "serde")]
#[test]
fn sampling_strategy_serde_uses_snake_case() {
let strat = SamplingStrategy::Greedy { best_of: 1 };
let json = serde_json::to_string(&strat).expect("serialize");
assert!(
json.contains("greedy"),
"external rep must be snake_case; got {json}"
);
let back: SamplingStrategy = serde_json::from_str(&json).expect("deserialize");
match back {
SamplingStrategy::Greedy { best_of } => assert_eq!(best_of, 1),
_ => panic!("expected Greedy"),
}
}
#[cfg(any(feature = "alignment", feature = "emissions"))]
#[test]
fn a_sub_segment_is_an_offset_within_its_chunk_wherever_the_chunk_sits() {
let mut request = AlignmentRequest::for_test(
ChunkId::from_raw(5),
NonZeroU64::new(1).expect("1 != 0"),
Arc::from(vec![0.0_f32; 1_600]),
SmolStr::new("hello"),
Lang::En,
Vec::new(),
);
let first = (1_u64 << 63) - 5;
request.context.first_sample = first;
request.context.sub_segments_samples = vec![(first, first + 5), (first + 15, first + 30)];
let local: Vec<(i64, i64, mediatime::Timebase)> = request
.chunk_local_sub_segments()
.iter()
.map(|range| (range.start_pts(), range.end_pts(), range.timebase()))
.collect();
let analysis = crate::time::ANALYSIS_TIMEBASE;
assert_eq!(local, [(0, 5, analysis), (15, 30, analysis)]);
}
}