use std::num::NonZeroUsize;
use super::*;
use crate::audio::whisper::{
audio::{
chunker::{AudioChunk, VadChunker, prepare_seek_clips},
vad::{EnergyVad, VoiceActivityDetector},
},
options::{AlignmentGather, ChunkingStrategy, Task, WordGrouping},
result::TranscriptionTimings,
task_facts::{SpanKnowledge, TaskFacts},
};
type OptionMutation = (&'static str, fn(DecodingOptions) -> DecodingOptions);
fn mutations() -> Vec<OptionMutation> {
vec![
("task", |o| o.with_task(Task::Translate)),
("language", |o| o.with_language("es")),
("temperature", |o| o.with_temperature(0.4)),
("temperature_increment_on_fallback", |o| {
o.with_temperature_increment_on_fallback(0.3)
}),
("temperature_fallback_count", |o| {
o.with_temperature_fallback_count(9)
}),
("sample_length", |o| o.with_sample_length(64)),
("top_k", |o| o.with_top_k(11)),
("seed", |o| o.with_seed(42)),
("use_prefill_prompt", |o| o.maybe_use_prefill_prompt(false)),
("detect_language", |o| o.maybe_detect_language(true)),
("skip_special_tokens", |o| o.with_skip_special_tokens()),
("without_timestamps", |o| o.with_without_timestamps()),
("word_timestamps", |o| o.with_word_timestamps()),
("max_initial_timestamp", |o| {
o.with_max_initial_timestamp(2.5)
}),
("max_window_seek", |o| o.with_max_window_seek(1_234)),
("clip_timestamps", |o| {
o.with_clip_timestamps(vec![0.5, 3.0])
}),
("window_clip_time", |o| o.with_window_clip_time(2.0)),
("prompt_tokens", |o| {
o.with_prompt_tokens(vec![101_u32, 102])
}),
("prefix_tokens", |o| o.with_prefix_tokens(vec![201_u32])),
("suppress_blank", |o| o.with_suppress_blank()),
("suppress_tokens", |o| o.with_suppress_tokens(vec![301_u32])),
("compression_ratio_threshold", |o| {
o.maybe_compression_ratio_threshold(None)
}),
("logprob_threshold", |o| o.maybe_logprob_threshold(None)),
("first_token_logprob_threshold", |o| {
o.maybe_first_token_logprob_threshold(None)
}),
("no_speech_threshold", |o| o.maybe_no_speech_threshold(None)),
("concurrent_worker_count", |o| {
o.with_concurrent_worker_count(NonZeroUsize::new(3).unwrap())
}),
("chunking_strategy", |o| {
o.with_chunking_strategy(ChunkingStrategy::Vad)
}),
("verbose", |o| o.with_verbose()),
("drop_blank_audio", |o| o.maybe_drop_blank_audio(false)),
("word_grouping", |o| {
o.with_word_grouping(WordGrouping::FineGrained)
}),
("alignment_gather", |o| {
o.with_alignment_gather(AlignmentGather::SwiftParity)
}),
]
}
#[test]
fn provenance_records_every_decoding_option() {
let compute = ComputeOptions::new();
let baseline = DecodingOptions::new();
let baseline_record = Provenance::from_options(&baseline, &compute, 0.0, false);
for (field, mutate) in mutations() {
let mutated = mutate(baseline.clone());
assert_ne!(
mutated, baseline,
"`{field}`'s row does not actually change the options, so its assertion \
below would prove nothing"
);
let record = Provenance::from_options(&mutated, &compute, 0.0, false);
assert_ne!(
record, baseline_record,
"`{field}` is NOT represented in the provenance record: two runs \
differing only in it leave byte-identical records"
);
assert_eq!(
record.decoding(),
&mutated,
"`{field}` must be captured verbatim, not approximated"
);
}
}
#[test]
fn drop_blank_audio_and_word_grouping_are_recorded() {
let compute = ComputeOptions::new();
let record =
|decoding: &DecodingOptions| Provenance::from_options(decoding, &compute, 0.0, false);
let dropping = DecodingOptions::new();
let emitting = DecodingOptions::new().maybe_drop_blank_audio(false);
assert!(dropping.drop_blank_audio(), "dropping is the default");
assert_ne!(
record(&dropping),
record(&emitting),
"drop_blank_audio must be legible in the record"
);
assert!(record(&dropping).decoding().drop_blank_audio());
assert!(!record(&emitting).decoding().drop_blank_audio());
let fine = DecodingOptions::new().with_word_grouping(WordGrouping::FineGrained);
let swift = DecodingOptions::new();
assert_eq!(fine.word_grouping(), WordGrouping::FineGrained);
assert_eq!(
swift.word_grouping(),
WordGrouping::SwiftParity,
"swift-parity is the #41 default"
);
assert_ne!(
record(&fine),
record(&swift),
"word_grouping must be legible in the record"
);
assert_eq!(
record(&fine).decoding().word_grouping(),
WordGrouping::FineGrained
);
}
#[test]
fn mutation_table_covers_every_decoding_option() {
let expected: std::collections::BTreeSet<&str> =
crate::audio::whisper::options::DECODING_OPTION_FIELD_NAMES
.iter()
.copied()
.collect();
let covered: std::collections::BTreeSet<&str> =
mutations().iter().map(|(field, _)| *field).collect();
assert_eq!(
covered, expected,
"the provenance mutation table has fallen out of step with \
DecodingOptions -- every knob needs a row (see `mutations`)"
);
}
type TaskFactMutation = (&'static str, fn(TranscriptionResult) -> TranscriptionResult);
fn baseline_task_result() -> TranscriptionResult {
TranscriptionResult::new(
"x",
vec![TranscriptionSegment::new().with_temperature(0.0)],
"en",
TranscriptionTimings::new(),
)
}
const DERIVED_TASK_FACT: &str = "effective_temperature";
fn task_fact_mutations() -> Vec<TaskFactMutation> {
vec![
("observed_language", |r| {
r.with_task_facts(TaskFacts::unknown().with_observed_language(Some("es".to_string())))
}),
("drew_from_rng", |r| {
r.with_task_facts(TaskFacts::unknown().with_drew_from_rng(true))
}),
("early_stopped", |r| {
r.with_task_facts(TaskFacts::unknown().with_early_stopped(true))
}),
("had_swallowed_error", |r| {
r.with_task_facts(TaskFacts::unknown().with_had_swallowed_error(true))
}),
("worker_schedule", |r| {
r.with_task_facts(TaskFacts::unknown().with_worker(3))
}),
("decoded_span", |r| {
r.with_task_facts(TaskFacts::unknown().with_decoded_span(SpanKnowledge::Exact(1)))
}),
(DERIVED_TASK_FACT, |mut r| {
r.set_segments(vec![TranscriptionSegment::new().with_temperature(0.6)]);
r
}),
]
}
#[test]
fn provenance_records_every_task_fact() {
let decoding = DecodingOptions::new();
let compute = ComputeOptions::new();
let baseline = baseline_task_result();
let baseline_record = Provenance::for_result(&decoding, &compute, &baseline);
for (field, mutate) in task_fact_mutations() {
let mutated = mutate(baseline.clone());
assert_ne!(
mutated, baseline,
"`{field}`'s row does not actually change the result, so its assertion \
below would prove nothing"
);
let record = Provenance::for_result(&decoding, &compute, &mutated);
assert_ne!(
record, baseline_record,
"task fact `{field}` is NOT recorded in the provenance: two runs \
differing only in it leave byte-identical records"
);
}
}
#[test]
fn task_fact_table_covers_every_provenance_task_fact() {
const NON_TASK_FACTS: &[&str] = &[
"decoding", "compute", "model_id",
"model_revision",
"tokenizer_id",
"tokenizer_revision",
"vad_detector",
];
let provenance_task_layer: std::collections::BTreeSet<&str> =
crate::audio::whisper::provenance::PROVENANCE_FIELD_NAMES
.iter()
.copied()
.filter(|field| !NON_TASK_FACTS.contains(field))
.collect();
let provenance_expected: std::collections::BTreeSet<&str> =
[DERIVED_TASK_FACT, "task_facts"].into_iter().collect();
assert_eq!(
provenance_task_layer, provenance_expected,
"a Provenance field is neither a non-task-fact, the derived outcome, nor the \
carried `task_facts` record -- place it in this test's partition"
);
let expected: std::collections::BTreeSet<&str> =
crate::audio::whisper::task_facts::TASK_FACTS_FIELD_NAMES
.iter()
.copied()
.chain(std::iter::once(DERIVED_TASK_FACT))
.collect();
let covered: std::collections::BTreeSet<&str> = task_fact_mutations()
.iter()
.map(|(field, _)| *field)
.collect();
assert_eq!(
covered, expected,
"the provenance task-fact table has fallen out of step with TaskFacts -- \
every carried sub-fact needs a mutation row (see `task_fact_mutations`)"
);
}
#[cfg(feature = "serde")]
#[test]
fn every_decoding_option_survives_the_provenance_round_trip() {
let compute = ComputeOptions::new();
let baseline = DecodingOptions::new();
let baseline_json =
serde_json::to_string(&Provenance::from_options(&baseline, &compute, 0.0, false))
.expect("baseline serializes");
for (field, mutate) in mutations() {
let mutated = mutate(baseline.clone());
let record = Provenance::from_options(&mutated, &compute, 0.0, false);
let json = serde_json::to_string(&record).expect("record serializes");
assert_ne!(
json, baseline_json,
"`{field}` leaves no trace in the SERIALIZED record"
);
assert_eq!(
serde_json::from_str::<Provenance>(&json).expect("record deserializes"),
record,
"`{field}` does not survive the round trip"
);
}
}
fn distinctive_decoding() -> DecodingOptions {
DecodingOptions::new()
.with_task(Task::Translate)
.with_language("es")
.with_skip_special_tokens()
.with_word_timestamps()
.with_chunking_strategy(ChunkingStrategy::Vad)
.with_temperature(0.2)
.with_temperature_increment_on_fallback(0.3)
.with_temperature_fallback_count(9)
.with_seed(42)
}
fn distinctive_compute() -> ComputeOptions {
ComputeOptions::new()
.with_mel(ComputeUnits::CpuOnly)
.with_encoder(ComputeUnits::All)
.with_decoder(ComputeUnits::CpuAndGpu)
}
#[test]
fn from_options_captures_the_options_and_invents_nothing_else() {
let decoding = distinctive_decoding();
let compute = distinctive_compute();
let provenance = Provenance::from_options(&decoding, &compute, 0.6, false);
assert_eq!(provenance.decoding(), &decoding);
assert_eq!(provenance.compute(), compute);
assert_eq!(provenance.encoder_compute_units(), ComputeUnits::All);
assert_eq!(provenance.effective_temperature(), Some(0.6));
assert_eq!(provenance.task_facts().observed_language(), None);
assert_eq!(provenance.model_id(), None);
assert_eq!(provenance.model_revision(), None);
assert_eq!(provenance.tokenizer_id(), None);
assert_eq!(provenance.tokenizer_revision(), None);
assert_eq!(provenance.vad_detector(), None);
}
#[test]
fn detect_language_reads_back_resolved_not_raw() {
let compute = ComputeOptions::new();
let prefilled = DecodingOptions::new();
assert!(prefilled.use_prefill_prompt());
assert!(
!Provenance::from_options(&prefilled, &compute, 0.0, false)
.decoding()
.detect_language()
);
let mut no_prefill = DecodingOptions::new();
no_prefill.clear_use_prefill_prompt();
let provenance = Provenance::from_options(&no_prefill, &compute, 0.0, false);
assert!(
provenance.decoding().detect_language(),
"the resolved coupling, not the unset raw tri-state"
);
assert!(!provenance.decoding().use_prefill_prompt());
}
#[test]
fn for_segment_reads_the_effective_temperature_off_the_segment() {
let decoding = DecodingOptions::new();
let compute = ComputeOptions::new();
let segment = TranscriptionSegment::new().with_temperature(0.4);
let provenance = Provenance::for_segment(&decoding, &compute, &segment, false);
assert_eq!(provenance.effective_temperature(), Some(0.4));
assert_eq!(
provenance.decoding().temperature(),
0.0,
"the BASE temperature stays what was configured"
);
assert_eq!(
provenance,
Provenance::from_options(&decoding, &compute, 0.4, false),
"for_segment is from_options with the segment's temperature and draw fact"
);
}
fn result_at(language: &str, temperatures: &[f32]) -> TranscriptionResult {
TranscriptionResult::new(
"Hello world.",
temperatures
.iter()
.map(|&t| TranscriptionSegment::new().with_temperature(t))
.collect::<Vec<_>>(),
language,
TranscriptionTimings::new(),
)
.with_task_facts(
TaskFacts::unknown()
.with_observed_language((!language.is_empty()).then(|| language.to_string()))
.with_early_stopped(false)
.with_had_swallowed_error(false)
.with_drew_from_rng(temperatures.iter().any(|&t| t != 0.0)),
)
}
#[test]
fn for_result_records_the_detected_language_not_the_configured_one() {
let decoding = DecodingOptions::new();
let compute = ComputeOptions::new();
assert_eq!(decoding.language(), "", "auto-detect is the default");
let provenance = Provenance::for_result(&decoding, &compute, &result_at("es", &[0.0]));
assert_eq!(
provenance.decoding().language(),
"",
"the CONFIGURED language, verbatim"
);
assert_eq!(
provenance.task_facts().observed_language(),
Some("es"),
"the DETECTED language — the fact the record exists to carry"
);
assert_eq!(
Provenance::from_options(&decoding, &compute, 0.0, false)
.task_facts()
.observed_language(),
None
);
}
#[test]
fn for_result_detected_language_is_absent_when_no_window_observed_one() {
let compute = ComputeOptions::new();
let opts = DecodingOptions::new();
let unobserved = TranscriptionResult::new("", Vec::new(), "en", TranscriptionTimings::new())
.with_task_facts(TaskFacts::unknown().with_observed_language(None));
assert_eq!(
unobserved.language(),
"en",
"the Swift-compat display fallback is kept on the result"
);
assert_eq!(
Provenance::for_result(&opts, &compute, &unobserved)
.task_facts()
.observed_language(),
None,
"nothing was observed -- the detected language is absent, not fabricated"
);
let dropped_english = TranscriptionResult::new("", Vec::new(), "en", TranscriptionTimings::new())
.with_task_facts(TaskFacts::unknown().with_observed_language(Some("en".to_string())));
assert!(dropped_english.segments_slice().is_empty());
assert_eq!(
Provenance::for_result(&opts, &compute, &dropped_english)
.task_facts()
.observed_language(),
Some("en"),
"a genuinely observed language must survive its segments being dropped"
);
}
#[test]
fn for_result_temperature_is_some_only_when_every_segment_agrees() {
let decoding = DecodingOptions::new();
let compute = ComputeOptions::new();
let for_result = |temperatures: &[f32]| {
Provenance::for_result(&decoding, &compute, &result_at("en", temperatures))
.effective_temperature()
};
assert_eq!(
for_result(&[0.0, 0.0, 0.0]),
Some(0.0),
"no fallback anywhere"
);
assert_eq!(
for_result(&[0.4, 0.4]),
Some(0.4),
"the same rung throughout"
);
assert_eq!(
for_result(&[0.0, 0.2]),
None,
"the ladder split the segments: no single temperature describes this"
);
assert_eq!(
for_result(&[]),
None,
"no segments (silence, once drop_blank_audio empties it) landed anywhere"
);
}
#[test]
fn is_reproducible_requires_greedy_or_a_seed() {
let compute = ComputeOptions::new();
let unseeded = DecodingOptions::new();
let seeded = DecodingOptions::new().with_seed(7);
let repro = |decoding: &DecodingOptions, temps: &[f32]| {
let result = result_at("en", temps);
let facts = result.task_facts().clone().with_worker(0);
Provenance::for_result(decoding, &compute, &result.with_task_facts(facts)).is_reproducible()
};
assert!(repro(&unseeded, &[0.0]));
assert!(repro(&seeded, &[0.0]));
assert!(
!repro(&unseeded, &[0.2]),
"a fallback climb with no seed is not reproducible"
);
assert!(repro(&seeded, &[0.2]));
assert!(!repro(&unseeded, &[0.0, 0.2]));
assert!(repro(&seeded, &[0.0, 0.2]));
assert!(!Provenance::from_options(&unseeded, &compute, 0.0, false).is_reproducible());
assert!(
!Provenance::from_options(&seeded, &compute, 0.0, false).is_reproducible(),
"a seed cannot rescue what the constructor never observed — the truncation"
);
assert!(!Provenance::from_options(&seeded, &compute, 0.2, true).is_reproducible());
}
#[test]
fn identity_uses_the_full_option_vocabulary() {
let base = Provenance::from_options(&DecodingOptions::new(), &ComputeOptions::new(), 0.0, false);
let built = base
.clone()
.with_model_id("openai_whisper-tiny")
.with_model_revision("a1b2c3d")
.with_tokenizer_id("openai/whisper-tiny")
.with_tokenizer_revision("deadbeef");
assert_eq!(built.model_id(), Some("openai_whisper-tiny"));
assert_eq!(built.model_revision(), Some("a1b2c3d"));
assert_eq!(built.tokenizer_id(), Some("openai/whisper-tiny"));
assert_eq!(built.tokenizer_revision(), Some("deadbeef"));
let mut m = built.clone();
m.clear_model_id().clear_tokenizer_revision();
assert_eq!(m.model_id(), None);
assert_eq!(m.tokenizer_revision(), None);
m.set_model_id("other").update_tokenizer_id(None);
assert_eq!(m.model_id(), Some("other"));
assert_eq!(m.tokenizer_id(), None);
let via_maybe = base.maybe_model_revision(Some("feedface".to_string()));
assert_eq!(via_maybe.model_revision(), Some("feedface"));
}
#[test]
fn provenance_never_infers_the_vad_detector() {
struct AlwaysActiveVad;
impl VoiceActivityDetector for AlwaysActiveVad {
fn voice_activity(&self, samples: &[f32]) -> Vec<bool> {
vec![true; samples.len().div_ceil(self.frame_length_samples())]
}
fn frame_length_samples(&self) -> usize {
crate::audio::whisper::audio::vad::DEFAULT_FRAME_LENGTH_SAMPLES
}
}
let installed: Box<dyn VoiceActivityDetector + Send + Sync> = Box::new(AlwaysActiveVad);
let mut audio = vec![0.1f32; 96_000];
audio[32_000..35_200].fill(0.0);
audio[64_000..67_200].fill(0.0);
let clips = prepare_seek_clips(&[], audio.len()).unwrap();
let boundaries = |vad: &(dyn VoiceActivityDetector + Send + Sync)| {
VadChunker::new()
.chunk_all(vad, &audio, 48_000, &clips)
.iter()
.map(AudioChunk::seek_offset)
.collect::<Vec<_>>()
};
assert_ne!(
boundaries(&EnergyVad::new()),
boundaries(installed.as_ref()),
"the swap must move the chunk boundaries, or the rest proves nothing"
);
let decoding = DecodingOptions::new().with_chunking_strategy(ChunkingStrategy::Vad);
let compute = ComputeOptions::new();
let result = result_at("en", &[0.0]);
let default_run = Provenance::for_result(&decoding, &compute, &result);
let installed_run = Provenance::for_result(&decoding, &compute, &result);
assert_eq!(
default_run, installed_run,
"there is no constructor parameter the detector could arrive through"
);
assert_eq!(
installed_run.decoding().chunking_strategy(),
ChunkingStrategy::Vad
);
assert_eq!(
installed_run.vad_detector(),
None,
"never inferred: not `AlwaysActiveVad`, not the default `EnergyVad`"
);
let named =
Provenance::for_result(&decoding, &compute, &result).with_vad_detector("AlwaysActiveVad");
assert_eq!(named.vad_detector(), Some("AlwaysActiveVad"));
assert_ne!(
named, default_run,
"supplied, the two runs are finally distinguishable"
);
assert_eq!(
Provenance::for_result(&decoding, &compute, &result)
.maybe_vad_detector(Some("AlwaysActiveVad".to_string())),
named
);
let mut mutated =
Provenance::for_result(&decoding, &compute, &result).with_vad_detector("AlwaysActiveVad");
mutated.clear_vad_detector();
assert_eq!(mutated.vad_detector(), None);
mutated.set_vad_detector("EnergyVad");
assert_eq!(mutated.vad_detector(), Some("EnergyVad"));
mutated.update_vad_detector(None);
assert_eq!(mutated.vad_detector(), None);
#[cfg(feature = "serde")]
{
let unsupplied: serde_json::Value = serde_json::to_value(&installed_run).unwrap();
assert!(
!unsupplied.as_object().unwrap().contains_key("vad_detector"),
"an unsupplied detector is absent, not null"
);
assert_eq!(
serde_json::from_str::<Provenance>(&unsupplied.to_string()).unwrap(),
installed_run,
"and it reads back `None`, never a guess"
);
let supplied: serde_json::Value = serde_json::to_value(&named).unwrap();
assert_eq!(supplied["vad_detector"], "AlwaysActiveVad");
assert_eq!(
serde_json::from_str::<Provenance>(&supplied.to_string()).unwrap(),
named
);
}
}
#[cfg(feature = "serde")]
#[test]
fn serde_round_trips_every_field() {
let full = Provenance::from_options(&distinctive_decoding(), &distinctive_compute(), 0.6, false)
.with_model_id("openai_whisper-tiny")
.with_model_revision("a1b2c3d")
.with_tokenizer_id("openai/whisper-tiny")
.with_tokenizer_revision("deadbeef")
.with_vad_detector("SileroVad");
let json = serde_json::to_string(&full).unwrap();
assert_eq!(serde_json::from_str::<Provenance>(&json).unwrap(), full);
}
#[cfg(feature = "serde")]
#[test]
fn unset_identity_serializes_as_absent_not_null() {
let bare = Provenance::from_options(&DecodingOptions::new(), &ComputeOptions::new(), 0.0, false);
let value: serde_json::Value = serde_json::to_value(&bare).unwrap();
let object = value.as_object().unwrap();
for absent in [
"model_id",
"model_revision",
"tokenizer_id",
"tokenizer_revision",
"vad_detector",
] {
assert!(
!object.contains_key(absent),
"unset `{absent}` must be absent, not null"
);
}
assert!(
!object["decoding"].as_object().unwrap().contains_key("seed"),
"an unset seed must be absent, not null"
);
for present in ["decoding", "compute", "effective_temperature", "task_facts"] {
assert!(object.contains_key(present), "`{present}` must be recorded");
}
let facts = object["task_facts"].as_object().unwrap();
assert!(
facts["observed_language"].is_null(),
"an unobserved detection is an explicit null, not an omission"
);
assert!(
facts["worker_schedule"].is_null(),
"an unknown worker schedule is an explicit null, never a fabricated 0"
);
assert_eq!(
serde_json::from_str::<Provenance>(&value.to_string()).unwrap(),
bare
);
}
#[cfg(feature = "serde")]
#[test]
fn a_record_missing_a_library_known_field_is_rejected() {
let full = Provenance::for_result(
&distinctive_decoding(),
&distinctive_compute(),
&result_at("es", &[0.6]),
);
let value: serde_json::Value = serde_json::to_value(&full).unwrap();
assert_eq!(
serde_json::from_str::<Provenance>(&value.to_string()).unwrap(),
full,
"the intact record must round-trip, or the removals below prove nothing"
);
for required in ["decoding", "compute", "effective_temperature", "task_facts"] {
let mut without = value.clone();
without.as_object_mut().unwrap().remove(required).unwrap();
assert!(
serde_json::from_str::<Provenance>(&without.to_string()).is_err(),
"a missing `{required}` must fail, not default"
);
}
for required in [
"drew_from_rng",
"observed_language",
"early_stopped",
"worker_schedule",
] {
let mut without = value.clone();
without
.as_object_mut()
.unwrap()
.get_mut("task_facts")
.unwrap()
.as_object_mut()
.unwrap()
.remove(required)
.unwrap();
assert!(
serde_json::from_str::<Provenance>(&without.to_string()).is_err(),
"a missing `task_facts.{required}` must fail, not default"
);
}
let mut nulled = value;
nulled.as_object_mut().unwrap()["effective_temperature"] = serde_json::Value::Null;
{
let facts = nulled
.as_object_mut()
.unwrap()
.get_mut("task_facts")
.unwrap()
.as_object_mut()
.unwrap();
facts["observed_language"] = serde_json::Value::Null;
facts["worker_schedule"] = serde_json::Value::Null;
}
let read: Provenance = serde_json::from_str(&nulled.to_string()).unwrap();
assert_eq!(read.effective_temperature(), None);
assert_eq!(read.task_facts().observed_language(), None);
assert_eq!(
read.task_facts().worker_schedule(),
None,
"a null worker schedule reads back explicit unknown, never [0]"
);
}
#[test]
fn for_result_reads_the_carried_sampling_fact_not_the_surviving_segments() {
let compute = ComputeOptions::new();
let unseeded = DecodingOptions::new();
let filtered = TranscriptionResult::new(
"Hello world.",
vec![TranscriptionSegment::new().with_temperature(0.0)],
"en",
TranscriptionTimings::new(),
)
.with_task_facts(
TaskFacts::unknown()
.with_drew_from_rng(true)
.with_early_stopped(false)
.with_had_swallowed_error(false)
.with_worker(0),
);
let record = Provenance::for_result(&unseeded, &compute, &filtered);
assert_eq!(
record.effective_temperature(),
Some(0.0),
"the survivors really are all greedy -- the fix must not come from here"
);
assert_eq!(record.task_facts().drew_from_rng(), Some(true));
assert!(
!record.is_reproducible(),
"an unseeded sampled window happened, even though nothing survived to say so"
);
assert!(
Provenance::for_result(&DecodingOptions::new().with_seed(3), &compute, &filtered)
.is_reproducible()
);
let empty = result_at("en", &[]);
assert_eq!(empty.task_facts().drew_from_rng(), Some(false));
let greedy = Provenance::for_result(&unseeded, &compute, &empty);
assert_eq!(greedy.effective_temperature(), None);
assert!(
greedy.is_reproducible(),
"nothing ever drew from the sampler, so there is nothing to replay"
);
}
#[test]
fn from_options_and_for_segment_record_the_explicit_draw_fact_not_the_temperature() {
let compute = ComputeOptions::new();
let decoding = DecodingOptions::new();
assert_eq!(
Provenance::from_options(&decoding, &compute, 0.3, false)
.task_facts()
.drew_from_rng(),
Some(false),
"an explicit no-draw at 0.3 is recorded as not sampled, not inferred from 0.3"
);
assert_eq!(
Provenance::from_options(&decoding, &compute, 0.0, true)
.task_facts()
.drew_from_rng(),
Some(true),
"an explicit draw is recorded even at temperature 0.0"
);
let never_drew = TranscriptionSegment::new().with_temperature(0.3);
let record = Provenance::for_segment(&decoding, &compute, &never_drew, false);
assert_eq!(
record.effective_temperature(),
Some(0.3),
"the segment's rung is still recorded"
);
assert_eq!(
record.task_facts().drew_from_rng(),
Some(false),
"a 0.3 segment that never drew is not sampled -- for_segment must not infer from 0.3"
);
assert_eq!(
Provenance::for_segment(&decoding, &compute, &never_drew, true)
.task_facts()
.drew_from_rng(),
Some(true),
"and an explicit draw on that same segment IS recorded"
);
}
#[test]
fn provenance_records_the_real_draw_fact_never_the_temperature() {
let compute = ComputeOptions::new();
for &temperature in &[-0.2f32, -0.0, 0.0, 0.2] {
let mut sampler = crate::audio::whisper::decode::sampler::GreedyTokenSampler::new(
temperature,
99,
&DecodingOptions::new(),
);
let token = sampler.sample(&[-10.0f32, 10.0, 0.5]).token();
let drew = sampler.drew_from_rng();
assert_eq!(
drew,
temperature != 0.0,
"temperature {temperature}: the sampler draws iff temperature != 0.0"
);
if !drew {
assert_eq!(
token, 1,
"temperature {temperature} must decode greedily (argmax)"
);
}
let unseeded = Provenance::from_options(&DecodingOptions::new(), &compute, temperature, drew);
assert_eq!(
unseeded.task_facts().drew_from_rng(),
Some(drew),
"temperature {temperature}: the recorded fact is the explicit draw"
);
assert!(
!unseeded.is_reproducible(),
"temperature {temperature}: from_options never promises reproducibility"
);
let ran_to_completion = TranscriptionResult::new(
"x",
vec![TranscriptionSegment::new().with_temperature(temperature)],
"en",
TranscriptionTimings::new(),
)
.with_task_facts(
TaskFacts::unknown()
.with_drew_from_rng(drew)
.with_early_stopped(false)
.with_had_swallowed_error(false)
.with_worker(0),
);
assert_eq!(
Provenance::for_result(&DecodingOptions::new(), &compute, &ran_to_completion)
.is_reproducible(),
!drew,
"unseeded temperature {temperature}: reproducible iff nothing was drawn"
);
assert!(
Provenance::for_result(
&DecodingOptions::new().with_seed(7),
&compute,
&ran_to_completion
)
.is_reproducible(),
"seeded temperature {temperature} replays exactly"
);
}
}
#[test]
fn for_result_reads_the_carried_flag_not_the_segment_temperature() {
let compute = ComputeOptions::new();
let decoding = DecodingOptions::new();
let never_drew = TranscriptionResult::new(
"Hello world.",
vec![TranscriptionSegment::new().with_temperature(0.3)],
"en",
TranscriptionTimings::new(),
)
.with_task_facts(
TaskFacts::unknown()
.with_drew_from_rng(false)
.with_early_stopped(false)
.with_had_swallowed_error(false),
);
assert_eq!(
never_drew.task_facts().drew_from_rng(),
Some(false),
"the result itself carries no draw"
);
let record = Provenance::for_result(&decoding, &compute, &never_drew);
assert_eq!(
record.task_facts().drew_from_rng(),
Some(false),
"the 0.3 segment must NOT be read as a draw -- the carried flag is the only witness"
);
assert!(
record.is_reproducible(),
"nothing drew, so it is reproducible despite the 0.3 segment"
);
let drew_greedy_survivors = TranscriptionResult::new(
"Hello",
vec![TranscriptionSegment::new().with_temperature(0.0)],
"en",
TranscriptionTimings::new(),
)
.with_task_facts(
TaskFacts::unknown()
.with_drew_from_rng(true)
.with_early_stopped(false)
.with_had_swallowed_error(false),
);
let record = Provenance::for_result(&decoding, &compute, &drew_greedy_survivors);
assert_eq!(
record.task_facts().drew_from_rng(),
Some(true),
"a carried draw survives even when every remaining segment reads 0.0"
);
assert!(!record.is_reproducible());
}
#[cfg(feature = "serde")]
#[test]
fn effective_temperature_non_finite_is_rejected_by_serde() {
let compute = ComputeOptions::new();
for bad in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY] {
let record = Provenance::from_options(&DecodingOptions::new(), &compute, bad, false);
assert!(record.effective_temperature().is_some());
assert!(serde_json::to_string(&record).is_err());
}
let finite = Provenance::from_options(&DecodingOptions::new(), &compute, 0.7, false);
let json = serde_json::to_string(&finite).unwrap();
assert_eq!(serde_json::from_str::<Provenance>(&json).unwrap(), finite);
}