use std::{
sync::{Arc, atomic::AtomicBool},
vec::Vec,
};
use mediatime::TimeRange;
use smol_str::{SmolStr, format_smolstr};
use core::sync::atomic::Ordering;
use std::time::Instant;
use ort::session::RunOptions;
use crate::{
align::{Run, script_dispatch::runs_reproduce_text},
core::{
AlignmentCompletion, AlignmentRequest, OovDecision, OovResolution, UnalignedCause,
UnitAlignment, UnitJob, UnitOutcome, panic_failure,
},
runner::aligner::{AlignmentFallback, AlignmentLookup, AlignmentSet},
types::{
AlignmentError, AlignmentFailure, ChunkId, Lang, LanguageUnsupportedForAlignment, WorkFailure,
WorkerHangTimeout, WorkerKind,
},
};
mod job;
#[cfg(test)]
use crate::core::{clip_sub_segments, run_audio_slice};
pub(crate) use job::JobId;
pub use job::{JobDetection, JobResolution};
pub struct AlignWorkItem {
id: JobId,
request: AlignmentRequest,
sub_segments: Vec<TimeRange>,
abort_flag: Arc<AtomicBool>,
chunk_first_sample_in_stream: u64,
samples_to_output_range: Arc<dyn Fn(u64, u64) -> TimeRange + Send + Sync>,
}
impl AlignWorkItem {
#[must_use]
pub fn new(request: AlignmentRequest, abort_flag: Arc<AtomicBool>) -> Self {
Self {
id: JobId::next(),
sub_segments: request.chunk_local_sub_segments(),
abort_flag,
chunk_first_sample_in_stream: request.chunk_first_sample(),
samples_to_output_range: request.samples_to_output_range(),
request,
}
}
pub fn failed(self, failure: WorkFailure) -> AlignmentCompletion {
self.request.failed(failure)
}
pub(crate) const fn id(&self) -> JobId {
self.id
}
#[must_use]
pub const fn chunk_id(&self) -> ChunkId {
self.request.chunk_id()
}
#[must_use]
pub const fn samples(&self) -> &Arc<[f32]> {
self.request.samples()
}
#[must_use]
pub fn sub_segments(&self) -> &[TimeRange] {
&self.sub_segments
}
#[must_use]
pub const fn text(&self) -> &SmolStr {
self.request.text()
}
#[must_use]
pub const fn language(&self) -> &Lang {
self.request.language()
}
#[must_use]
pub fn runs(&self) -> &[Run] {
self.request.runs()
}
#[must_use]
pub fn abort_flag(&self) -> &Arc<AtomicBool> {
&self.abort_flag
}
#[must_use]
pub const fn chunk_first_sample_in_stream(&self) -> u64 {
self.chunk_first_sample_in_stream
}
#[must_use]
pub fn samples_to_output_range(&self) -> &Arc<dyn Fn(u64, u64) -> TimeRange + Send + Sync> {
&self.samples_to_output_range
}
}
pub fn run_one_alignment(
set: &AlignmentSet,
job: AlignWorkItem,
resolution: JobResolution,
run_options: &RunOptions,
) -> AlignmentCompletion {
answer_job(job, |job, units| {
align_job(set, job, units, resolution, run_options)
})
}
fn answer_job(
mut job: AlignWorkItem,
align: impl FnOnce(&AlignWorkItem, Vec<UnitJob>) -> Result<Vec<UnitOutcome>, WorkFailure>,
) -> AlignmentCompletion {
let units = job.request.take_units();
let answered = std::panic::catch_unwind(core::panic::AssertUnwindSafe(|| align(&job, units)))
.unwrap_or_else(|panic| Err(panic_failure(panic.as_ref(), job.language().clone())));
let AlignWorkItem { request, .. } = job;
match answered {
Ok(outcomes) => request.aligned(outcomes).unwrap_or_else(|refused| {
let message = format_smolstr!("{refused}");
let (request, _) = refused.into_parts();
let language = request.language().clone();
request.failed(WorkFailure::Alignment(AlignmentError::Tokenization(
AlignmentFailure::new(message, language),
)))
}),
Err(failure) => request.failed(failure),
}
}
fn align_job(
set: &AlignmentSet,
job: &AlignWorkItem,
units: Vec<UnitJob>,
resolution: JobResolution,
run_options: &RunOptions,
) -> Result<Vec<UnitOutcome>, WorkFailure> {
let started_at = Instant::now();
if job.abort_flag.load(Ordering::Relaxed) {
return Err(WorkFailure::WorkerHang(WorkerHangTimeout::new(
WorkerKind::Alignment,
started_at.elapsed(),
)));
}
validate_runs_reproduce_text(job.runs(), job.text(), job.language())?;
let resolutions = resolution.into_units_for(job, set.id())?;
if units.len() != resolutions.len() {
return Err(WorkFailure::Alignment(AlignmentError::Tokenization(
AlignmentFailure::new(
format_smolstr!(
"the job's request handed out {} unit jobs for {} units; each unit is answered by \
consuming its own job, once",
units.len(),
resolutions.len(),
),
job.language().clone(),
),
)));
}
let outcome = if job.runs().is_empty() {
let mut pairs = resolutions.iter().zip(units);
let Some((resolution, unit)) = pairs.next() else {
return Ok(Vec::new());
};
align_unit(set, resolution, unit, &job.abort_flag, run_options).map(|(alignment, unit)| {
if let UnitAlignment::Unaligned(cause) = &alignment {
log_unaligned(job.chunk_id(), None, job.language(), cause);
}
vec![unit.answer(alignment)]
})
} else {
dispatch_runs(set, job, resolutions, units, run_options)
};
match outcome {
Err(WorkFailure::WorkerHang(timeout)) => Err(WorkFailure::WorkerHang(WorkerHangTimeout::new(
timeout.kind(),
started_at.elapsed(),
))),
other => other,
}
}
fn align_unit(
set: &AlignmentSet,
unit: &OovResolution,
job: UnitJob,
abort_flag: &AtomicBool,
run_options: &RunOptions,
) -> Result<(UnitAlignment, UnitJob), WorkFailure> {
let language = job.language().clone();
let language = &language;
if unit.unit() != job.unit() {
return Err(WorkFailure::Alignment(AlignmentError::Tokenization(
AlignmentFailure::new(
format_smolstr!(
"the resolution of unit {:?} was offered to the job of unit {:?}; a unit's \
decisions apply to that unit alone",
unit.unit(),
job.unit(),
),
language.clone(),
),
)));
}
let aligner = match set.lookup(language) {
AlignmentLookup::Hit { aligner, .. } | AlignmentLookup::AnyFallback { aligner } => aligner,
AlignmentLookup::Miss { fallback } => {
return resolve_not_inspected(unit, fallback, language)
.map(|cause| (UnitAlignment::Unaligned(cause), job));
}
};
let mut guard = aligner.lock().unwrap_or_else(|p| p.into_inner());
let decisions = unit.read_by(guard.id()).ok_or_else(|| {
WorkFailure::Alignment(AlignmentError::Tokenization(AlignmentFailure::new(
SmolStr::new_static(
"this unit's decisions were not detected by the aligner that now reads it: the \
registry's aligners changed between detection and dispatch. Detect the job again.",
),
language.clone(),
)))
})?;
match guard.align_job(&job, decisions, language, abort_flag, run_options) {
Ok(alignment) => Ok((alignment, job)),
Err(WorkFailure::Alignment(err)) if alignment_error_is_recoverable(&err) => {
Ok((UnitAlignment::Unaligned(UnalignedCause::Failed(err)), job))
}
Err(failure) => Err(failure),
}
}
fn resolve_not_inspected(
unit: &OovResolution,
fallback: AlignmentFallback,
language: &Lang,
) -> Result<UnalignedCause, WorkFailure> {
let decision = unit.unread_decision().ok_or_else(|| {
WorkFailure::Alignment(AlignmentError::Tokenization(AlignmentFailure::new(
SmolStr::new_static(
"no aligner can read this unit, but its decisions were detected by an aligner: the \
registry changed between detection and dispatch. Detect the job again.",
),
language.clone(),
)))
})?;
match decision {
OovDecision::FailClosed => Ok(UnalignedCause::Refused),
OovDecision::Wildcard => match fallback {
AlignmentFallback::SkipChunk => Ok(UnalignedCause::Skipped),
AlignmentFallback::Error => Err(WorkFailure::LanguageUnsupported(
LanguageUnsupportedForAlignment::new(language.clone()),
)),
},
}
}
fn log_unaligned(
chunk_id: ChunkId,
run_index: Option<usize>,
language: &Lang,
cause: &UnalignedCause,
) {
let kind = match cause {
UnalignedCause::Skipped => "skipped",
UnalignedCause::Refused => "refused",
UnalignedCause::NoAlignableText => "no_alignable_text",
UnalignedCause::NoSurvivingWords => "no_surviving_words",
UnalignedCause::Failed(AlignmentError::NoAlignmentPath(_)) => "failed:no_alignment_path",
UnalignedCause::Failed(AlignmentError::EmptyText(_)) => "failed:empty_text",
UnalignedCause::Failed(AlignmentError::SemanticOutOfVocab(_)) => "failed:semantic_out_of_vocab",
UnalignedCause::Failed(_) => "failed",
};
eprintln!(
"asry alignment unaligned chunk={chunk_id:?} run={run_index:?} language={language:?} cause={kind}"
);
}
fn validate_runs_reproduce_text(
runs: &[Run],
text: &str,
language: &Lang,
) -> Result<(), WorkFailure> {
if runs.is_empty() || runs_reproduce_text(runs, text) {
return Ok(());
}
Err(WorkFailure::Alignment(AlignmentError::Tokenization(
AlignmentFailure::new(
SmolStr::new_static(
"AlignWorkItem::runs do not reproduce AlignWorkItem::text; the per-run road \
aligns the runs' texts only, so it may be taken only when they are the text. Forward \
the runs `Command::Alignment` carries, which reproduce its text, or pass no runs to \
align the whole text.",
),
language.clone(),
),
)))
}
fn alignment_failure_is_recoverable(failure: &WorkFailure) -> bool {
matches!(failure, WorkFailure::Alignment(err) if alignment_error_is_recoverable(err))
}
pub(crate) fn alignment_error_is_recoverable(err: &AlignmentError) -> bool {
matches!(
err,
AlignmentError::NoAlignmentPath(_)
| AlignmentError::EmptyText(_)
| AlignmentError::SemanticOutOfVocab(_)
)
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub(super) struct BoundsSourceCounters {
runs_total: usize,
runs_dtw: usize,
runs_segment: usize,
runs_wholeclip: usize,
runs_unaligned: usize,
}
impl BoundsSourceCounters {
pub(super) fn observe_bounds(&mut self, source: crate::align::BoundsSource) {
self.runs_total += 1;
match source {
crate::align::BoundsSource::Dtw => self.runs_dtw += 1,
crate::align::BoundsSource::Segment => self.runs_segment += 1,
crate::align::BoundsSource::Wholeclip => self.runs_wholeclip += 1,
}
}
pub(super) const fn observe_unaligned(&mut self) {
self.runs_unaligned += 1;
}
pub(super) const fn runs_total(&self) -> usize {
self.runs_total
}
pub(super) const fn runs_dtw(&self) -> usize {
self.runs_dtw
}
pub(super) const fn runs_segment(&self) -> usize {
self.runs_segment
}
pub(super) const fn runs_wholeclip(&self) -> usize {
self.runs_wholeclip
}
pub(super) const fn runs_unaligned(&self) -> usize {
self.runs_unaligned
}
}
fn check_abort_between_runs(
abort_flag: &AtomicBool,
dispatch_started_at: Instant,
) -> Result<(), WorkFailure> {
if abort_flag.load(Ordering::Relaxed) {
return Err(WorkFailure::WorkerHang(WorkerHangTimeout::new(
WorkerKind::Alignment,
dispatch_started_at.elapsed(),
)));
}
Ok(())
}
fn dispatch_runs(
set: &AlignmentSet,
job: &AlignWorkItem,
resolutions: Vec<OovResolution>,
units: Vec<UnitJob>,
run_options: &RunOptions,
) -> Result<Vec<UnitOutcome>, WorkFailure> {
let mut counters = BoundsSourceCounters::default();
let mut outcomes: Vec<UnitOutcome> = Vec::with_capacity(job.runs().len());
let dispatch_started_at = Instant::now();
for (((run_idx, run), resolution), unit) in
job.runs().iter().enumerate().zip(&resolutions).zip(units)
{
if let Err(failure) = check_abort_between_runs(&job.abort_flag, dispatch_started_at) {
emit_telemetry(job.chunk_id(), &counters);
return Err(failure);
}
counters.observe_bounds(run.bounds_source());
let (outcome, unit) = align_unit(set, resolution, unit, &job.abort_flag, run_options)
.inspect_err(|_| emit_telemetry(job.chunk_id(), &counters))?;
if let UnitAlignment::Unaligned(cause) = &outcome {
counters.observe_unaligned();
log_unaligned(job.chunk_id(), Some(run_idx), run.language(), cause);
}
outcomes.push(unit.answer(outcome));
}
emit_telemetry(job.chunk_id(), &counters);
Ok(outcomes)
}
fn emit_telemetry(chunk_id: ChunkId, c: &BoundsSourceCounters) {
std::eprintln!(
"script_dispatch chunk={} runs={} dtw={} segment={} wholeclip={} unaligned={}",
chunk_id.as_u64(),
c.runs_total(),
c.runs_dtw(),
c.runs_segment(),
c.runs_wholeclip(),
c.runs_unaligned(),
);
}
#[cfg(test)]
mod tests;