use core::{num::NonZeroUsize, sync::atomic::AtomicBool, time::Duration};
use std::{collections::BTreeSet, path::Path};
use crate::ComputeUnits;
use asry::{
Lang, TimeRange,
emissions::{
DynTextNormalizer, EmissionsAligner, EmissionsError, EmissionsFailure, OovDecision,
OovDetection, OovResolution, OutputClock, PreparedChunk, SpeechCoverage, SpeechSpans,
UnitAlignment,
},
};
use tokenizers::{Tokenizer, models::ModelWrapper};
use crate::audio::align::{
acoustic::{AcousticContract, check_tokenization},
encode::{DEFAULT_ENCODER_COMPUTE, Encoder, EncoderInput},
error::{
AlignError, AlignerError, BlankOutOfVocabulary, InputTooLong, Refusal, RefusedOov,
ReservedSetMismatch, VocabularyMismatch,
},
vocab::Vocabulary,
};
pub const DEFAULT_MIN_SPEECH_COVERAGE: f32 = SpeechCoverage::DEFAULT.get();
pub const DEFAULT_MAX_INTRA_SILENT_RUN: Duration = asry::emissions::DEFAULT_MAX_INTRA_SILENT_RUN;
#[cfg(feature = "serde")]
fn default_min_speech_coverage() -> f32 {
DEFAULT_MIN_SPEECH_COVERAGE
}
#[cfg(feature = "serde")]
fn default_max_intra_silent_run() -> Duration {
DEFAULT_MAX_INTRA_SILENT_RUN
}
#[cfg(feature = "serde")]
fn default_compute() -> ComputeUnits {
DEFAULT_ENCODER_COMPUTE
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct AlignerOptions {
#[cfg_attr(feature = "serde", serde(default = "default_min_speech_coverage"))]
min_speech_coverage: f32,
#[cfg_attr(feature = "serde", serde(default = "default_max_intra_silent_run"))]
max_intra_silent_run: Duration,
#[cfg_attr(feature = "serde", serde(default = "default_compute"))]
compute: ComputeUnits,
}
impl Default for AlignerOptions {
fn default() -> Self {
Self::new()
}
}
impl AlignerOptions {
#[must_use]
pub const fn new() -> Self {
Self {
min_speech_coverage: DEFAULT_MIN_SPEECH_COVERAGE,
max_intra_silent_run: DEFAULT_MAX_INTRA_SILENT_RUN,
compute: DEFAULT_ENCODER_COMPUTE,
}
}
#[must_use]
pub const fn min_speech_coverage(&self) -> f32 {
self.min_speech_coverage
}
#[must_use]
pub const fn with_min_speech_coverage(mut self, coverage: f32) -> Self {
self.set_min_speech_coverage(coverage);
self
}
pub const fn set_min_speech_coverage(&mut self, coverage: f32) -> &mut Self {
self.min_speech_coverage = coverage;
self
}
#[must_use]
pub const fn max_intra_silent_run(&self) -> Duration {
self.max_intra_silent_run
}
#[must_use]
pub const fn with_max_intra_silent_run(mut self, run: Duration) -> Self {
self.set_max_intra_silent_run(run);
self
}
pub const fn set_max_intra_silent_run(&mut self, run: Duration) -> &mut Self {
self.max_intra_silent_run = run;
self
}
#[must_use]
pub const fn compute(&self) -> ComputeUnits {
self.compute
}
#[must_use]
pub const fn with_compute(mut self, compute: ComputeUnits) -> Self {
self.set_compute(compute);
self
}
pub const fn set_compute(&mut self, compute: ComputeUnits) -> &mut Self {
self.compute = compute;
self
}
}
impl core::fmt::Display for AlignerOptions {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(
f,
"min_speech_coverage={},max_intra_silent_run={},compute={}",
self.min_speech_coverage,
humantime::format_duration(self.max_intra_silent_run),
self.compute
)
}
}
fn build_seam(
language: Lang,
vocabulary: &Vocabulary,
contract: &AcousticContract,
normalizer: DynTextNormalizer,
options: &AlignerOptions,
) -> Result<EmissionsAligner, AlignerError> {
let geometry = contract.geometry();
let tokenization = contract.tokenization();
let document = vocabulary.tokenizer_json(contract);
let seam = EmissionsAligner::builder(language, &document)
.normalizer(normalizer)
.hop_samples(geometry.stride())
.receptive_field_samples(geometry.receptive_field())
.word_delimiter(tokenization.delimiter().seam_token())
.letter_case(tokenization.case().seam())
.min_speech_coverage(SpeechCoverage::clamped(options.min_speech_coverage()))
.max_intra_silent_run(options.max_intra_silent_run())
.blank_token_id(contract.blank())
.build()?;
check_reserved(
&seam,
&document,
vocabulary.non_lexical(contract.blank(), tokenization),
)?;
Ok(seam)
}
fn seam_reserved(
seam: &EmissionsAligner,
document: &[u8],
) -> Result<BTreeSet<usize>, AlignerError> {
let tokenizer = Tokenizer::from_bytes(document).map_err(|error| {
AlignerError::Seam(EmissionsError::Config(EmissionsFailure::new(
format!("the tokenizer document the seam was built from does not parse again: {error}")
.into(),
)))
})?;
let specials = tokenizer
.get_added_vocabulary()
.get_added_tokens_decoder()
.iter()
.filter(|(_, token)| token.special)
.map(|(&id, _)| id);
let unknown = if let ModelWrapper::WordLevel(model) = tokenizer.get_model() {
tokenizer.token_to_id(&model.unk_token)
} else {
None
};
Ok(
specials
.chain([seam.blank_token_id()])
.chain(tokenizer.token_to_id(seam.word_delimiter()))
.chain(unknown)
.filter_map(|id| usize::try_from(id).ok())
.collect(),
)
}
fn check_reserved(
seam: &EmissionsAligner,
document: &[u8],
declared: BTreeSet<usize>,
) -> Result<(), AlignerError> {
let reserved = seam_reserved(seam, document)?;
if reserved == declared {
Ok(())
} else {
Err(AlignerError::ReservedSetMismatch(ReservedSetMismatch::new(
declared.into_iter().collect(),
reserved.into_iter().collect(),
)))
}
}
fn check_blank(blank: u32, vocabulary: NonZeroUsize) -> Result<(), AlignerError> {
if usize::try_from(blank).is_ok_and(|blank| blank < vocabulary.get()) {
Ok(())
} else {
Err(AlignerError::BlankOutOfVocabulary(
BlankOutOfVocabulary::new(blank, vocabulary.get()),
))
}
}
fn check_vocabulary_width(
vocabulary: NonZeroUsize,
model: NonZeroUsize,
) -> Result<(), AlignerError> {
if vocabulary == model {
Ok(())
} else {
Err(AlignerError::VocabularyMismatch(VocabularyMismatch::new(
vocabulary.get(),
model.get(),
)))
}
}
fn effective_options(seam: &EmissionsAligner, requested: &AlignerOptions) -> AlignerOptions {
requested.with_min_speech_coverage(seam.min_speech_coverage().get())
}
pub struct Aligner {
encoder: crate::audio::align::encode::Encoder,
inner: EmissionsAligner,
options: AlignerOptions,
}
impl Aligner {
pub fn from_paths(
language: Lang,
model_path: &Path,
normalizer: DynTextNormalizer,
) -> Result<Self, AlignerError> {
Self::from_paths_with(language, model_path, normalizer, AlignerOptions::new())
}
pub fn from_paths_with(
language: Lang,
model_path: &Path,
normalizer: DynTextNormalizer,
options: AlignerOptions,
) -> Result<Self, AlignerError> {
Self::from_paths_with_vocabulary(
language,
model_path,
&Vocabulary::bundled(),
&AcousticContract::BASE960H,
normalizer,
options,
)
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(
name = "alignkit.aligner.load",
level = "info",
skip_all,
fields(
aligner_language = ?language,
model_path = ?model_path,
compute = ?options.compute(),
vocabulary = vocabulary.size().get(),
contract = ?contract,
),
)
)]
pub fn from_paths_with_vocabulary(
language: Lang,
model_path: &Path,
vocabulary: &Vocabulary,
contract: &AcousticContract,
normalizer: DynTextNormalizer,
options: AlignerOptions,
) -> Result<Self, AlignerError> {
check_blank(contract.blank(), vocabulary.size())?;
check_tokenization(
contract.blank(),
contract.tokenization(),
vocabulary,
normalizer.use_word_delimiter(),
)
.map_err(AlignerError::Tokenization)?;
let encoder = Encoder::load(model_path, contract, options.compute())?;
let inner = build_seam(language, vocabulary, contract, normalizer, &options)?;
check_vocabulary_width(inner.vocab_size(), encoder.vocab_size())?;
let options = effective_options(&inner, &options);
Ok(Self {
encoder,
inner,
options,
})
}
#[must_use]
pub const fn language_ref(&self) -> &Lang {
self.inner.language()
}
#[must_use]
pub const fn options(&self) -> AlignerOptions {
self.options
}
#[must_use]
pub const fn sample_rate(&self) -> u32 {
asry::time::SAMPLE_RATE_HZ
}
#[must_use]
pub const fn window_samples(&self) -> usize {
self.encoder.window_samples()
}
#[must_use]
pub const fn contract(&self) -> &AcousticContract {
self.encoder.contract()
}
pub fn detect_oov(&self, text: &str) -> Result<OovDetection, AlignError> {
Ok(self.inner.detect_oov(text)?)
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(
name = "alignkit.align_chunk",
level = "debug",
skip_all,
fields(
aligner_language = ?self.language_ref(),
samples = samples.len(),
sub_segments = sub_segments.len(),
text_bytes = text.len(),
oov_decisions = resolution.resolved().len(),
),
)
)]
pub fn align_chunk(
&self,
samples: &[f32],
sub_segments: &[TimeRange],
text: &str,
clock: OutputClock,
abort_flag: &AtomicBool,
resolution: OovResolution,
) -> Result<UnitAlignment, AlignError> {
if samples.len() > self.window_samples() {
return Err(AlignError::InputTooLong(InputTooLong::new(
samples.len(),
self.window_samples(),
)));
}
let speech = if sub_segments.is_empty() {
SpeechSpans::all_speech()
} else {
SpeechSpans::from_time_ranges(sub_segments)?
};
let refused = refused_positions(&resolution);
let prepared = self
.inner
.prepare(samples, &speech, text, resolution, clock, abort_flag)
.map_err(|err| seam_error(err, &refused))?;
check_audio(&prepared)?;
let emissions = prepared.encode_with(|buffer| {
debug_assert!(
core::ptr::eq(buffer, prepared.encoder_input()),
"asry hands the encoder the chunk's own prepared input"
);
self
.encoder
.emissions(EncoderInput::from_prepared(&prepared))
})?;
self
.inner
.finish(prepared, emissions, abort_flag)
.map_err(|err| seam_error(err, &refused))
}
}
fn check_audio(prepared: &PreparedChunk<'_>) -> Result<(), AlignError> {
if prepared.is_trivial() || prepared.real_samples() > 0 {
return Ok(());
}
Err(AlignError::NoAlignmentPath(EmissionsFailure::new(
"the chunk holds no audio, so no frame can carry its tokens".into(),
)))
}
fn refused_positions(resolution: &OovResolution) -> Vec<RefusedOov> {
resolution
.resolved()
.iter()
.filter(|resolved| resolved.decision() == OovDecision::FailClosed)
.map(|resolved| RefusedOov::detected(resolved.event()))
.collect()
}
fn seam_error(err: EmissionsError, refused: &[RefusedOov]) -> AlignError {
match err {
EmissionsError::SemanticOutOfVocab(failure) => {
if refused.is_empty() {
AlignError::Alignment(EmissionsError::SemanticOutOfVocab(failure))
} else {
AlignError::Refused(Refusal::new(refused.to_vec()))
}
}
EmissionsError::NoAlignmentPath(failure) => AlignError::NoAlignmentPath(failure),
other => AlignError::Alignment(other),
}
}
#[cfg(test)]
mod tests;