Skip to main content

eredu_evaluation/
lib.rs

1//! Backend-neutral model evaluation drivers.
2
3#![forbid(unsafe_code)]
4#![warn(missing_docs)]
5
6mod checkpoint;
7mod distribution;
8mod evidence;
9mod parity;
10mod realtime;
11
12pub use checkpoint::{
13    compare_checkpoint_artifacts, CheckpointParityError, CheckpointParityOptions,
14    CheckpointParityReport,
15};
16pub use distribution::{compare_distributions, DistributionError, DistributionMetrics};
17
18pub use evidence::{
19    observe_f32_tensor, observe_i32_tensor, observe_realtime_frame, summarize_latencies,
20    EvaluationEvidence, EvidenceError, LatencySummary,
21};
22pub use parity::{
23    compare_observations, LogitRowMetrics, LogitTolerance, NumericMetrics, NumericTolerance,
24    ObservationParity, ParityComparison, ParityError, ParityMetrics, ParityPolicy, ParityReport,
25    ParityRule,
26};
27pub use realtime::{encoded_audio_frames, run_realtime_trace, RealtimeTrace, RealtimeTraceError};
28
29use std::{
30    error::Error,
31    fs,
32    io::{self, Write},
33    path::{Path, PathBuf},
34    time::Instant,
35};
36
37use eredu_architectures::moshi::personaplex_prompt::{
38    wrap_system_prompt, AUDIO_TOKENS_PER_STREAM, SILENCE_TOKENS, SINE_TOKENS, TEXT_PADDING_TOKEN,
39};
40use eredu_codec::mimi::Mimi;
41use eredu_core::{
42    scheduler::{RequestId, SchedulerLimits},
43    RealtimeBackend, RealtimeInputFrame, RealtimeModel, RealtimeOutputFrame, RealtimeSampling,
44    RealtimeScheduler,
45};
46use eredu_nn::Tensor;
47use sentencepiece_rs::SentencePieceProcessor;
48use serde::Serialize;
49use serde_json::json;
50
51const SAMPLE_RATE: u32 = 24_000;
52const FRAME_RATE: f64 = 12.5;
53const FRAME_SAMPLES: usize = 1_920;
54const DEADLINE_MS: f64 = 1_000.0 / FRAME_RATE;
55const TAIL_ACTIVITY_FRAMES: usize = 3;
56const ACTIVE_AUDIO_DBFS: f64 = -40.0;
57const PROMPT_SILENCE_FRAMES: usize = 6;
58
59/// Default PersonaPlex system instruction used by the evaluator.
60pub const DEFAULT_TEXT_PROMPT: &str = "You are a wise and friendly teacher. Answer questions or provide advice in a clear and engaging way.";
61/// Default deterministic sampling seed.
62pub const DEFAULT_SAMPLING_SEED: u64 = 20_260_713;
63
64/// Paths consumed and produced by one PersonaPlex comparison.
65#[derive(Debug, Clone)]
66pub struct PersonaPlexEvaluationPaths {
67    /// Dense model artifact.
68    pub dense_model: PathBuf,
69    /// Quantized model artifact.
70    pub quantized_model: PathBuf,
71    /// SentencePiece text tokenizer.
72    pub text_tokenizer: PathBuf,
73    /// Mono 24 kHz raw `f32le` voice prompt.
74    pub voice_prompt: PathBuf,
75    /// Mono 24 kHz raw `f32le` user input.
76    pub input: PathBuf,
77    /// New output directory.
78    pub output: PathBuf,
79}
80
81/// Controls for one PersonaPlex comparison.
82#[derive(Debug, Clone)]
83pub struct PersonaPlexEvaluationOptions {
84    /// Maximum input frames, or all complete frames when omitted.
85    pub frames: Option<usize>,
86    /// Unwrapped system instruction.
87    pub text_prompt: String,
88    /// Root seed used independently by both model runs.
89    pub sampling_seed: u64,
90}
91
92impl Default for PersonaPlexEvaluationOptions {
93    fn default() -> Self {
94        Self {
95            frames: None,
96            text_prompt: DEFAULT_TEXT_PROMPT.into(),
97            sampling_seed: DEFAULT_SAMPLING_SEED,
98        }
99    }
100}
101
102/// Runs the complete PersonaPlex dense-versus-quantized evaluation.
103///
104/// The model loader is the only backend composition hook. Realtime inputs,
105/// outputs, forcing, sampling, scheduling, diagnostics, codec execution, and
106/// reporting use backend-neutral contracts.
107pub fn run_personaplex_quantization<B, T, L>(
108    paths: &PersonaPlexEvaluationPaths,
109    options: &PersonaPlexEvaluationOptions,
110    mimi: &mut Mimi<T>,
111    context: &T::Context,
112    mut load_model: L,
113) -> Result<(), Box<dyn Error>>
114where
115    B: RealtimeBackend,
116    T: Tensor,
117    L: FnMut(&Path) -> Result<RealtimeModel<B>, Box<dyn Error>>,
118{
119    if paths.output.exists() {
120        return Err(invalid(format!(
121            "output directory already exists: {}",
122            paths.output.display()
123        )));
124    }
125    let voice_pcm = read_f32le(&paths.voice_prompt)?;
126    if voice_pcm.len() < FRAME_SAMPLES {
127        return Err(invalid("voice prompt contains no complete 80 ms frame"));
128    }
129    let input_pcm = read_f32le(&paths.input)?;
130    let available_frames = input_pcm.len() / FRAME_SAMPLES;
131    let frames = options
132        .frames
133        .unwrap_or(available_frames)
134        .min(available_frames);
135    if frames < 4 {
136        return Err(invalid(format!(
137            "input must contain at least four complete frames; found {available_frames}"
138        )));
139    }
140    let input_pcm = &input_pcm[..frames * FRAME_SAMPLES];
141    let input_tail = tail_max_rms_dbfs(input_pcm);
142    let input_likely_truncated = input_tail > ACTIVE_AUDIO_DBFS;
143    let input_warning = input_likely_truncated.then_some(
144        "The final 240 ms contains active audio; the frame limit may truncate the user utterance.",
145    );
146    if let Some(warning) = input_warning {
147        eprintln!("warning: {warning} tail_max_rms_dbfs={input_tail:.1}");
148    }
149
150    let codec_start = Instant::now();
151    let voice_tokens = encode_pcm(mimi, &voice_pcm, context)?;
152    let input_tokens = encode_pcm(mimi, input_pcm, context)?;
153    if input_tokens.len() != frames {
154        return Err(invalid(format!(
155            "Mimi produced {} frames for {frames} PCM frames",
156            input_tokens.len()
157        )));
158    }
159    let encode_seconds = codec_start.elapsed().as_secs_f64();
160    let offline = offline_roundtrip(mimi, input_pcm, context)?;
161    let offline_tokens = offline.tokens;
162    let offline_roundtrip = offline.pcm;
163    let streaming_roundtrip = decode_tokens(mimi, &input_tokens, context, input_pcm.len())?;
164    let codec_agreement = token_frame_agreement(&input_tokens, &offline_tokens);
165
166    let tokenizer = SentencePieceProcessor::open(&paths.text_tokenizer)?;
167    let wrapped_text_prompt = wrap_system_prompt(&options.text_prompt);
168    let text_tokens = tokenizer
169        .encode_to_ids(&wrapped_text_prompt)?
170        .into_iter()
171        .map(i32::try_from)
172        .collect::<Result<Vec<_>, _>>()?;
173    if text_tokens.is_empty() {
174        return Err(invalid("text prompt tokenized to an empty sequence"));
175    }
176    let prompt = PromptConditioning {
177        voice_frames: voice_tokens,
178        text_tokens,
179    };
180
181    let dense_load_start = Instant::now();
182    let mut dense = load_model(&paths.dense_model)?;
183    let dense_load_seconds = dense_load_start.elapsed().as_secs_f64();
184    validate_personaplex_geometry(&dense)?;
185    let dense_reference = run_model(
186        &mut dense,
187        &prompt,
188        &input_tokens,
189        RealtimeSampling::greedy(),
190        RunMode::Diagnostics,
191    )?;
192    let sampling =
193        RealtimeSampling::new(0.7, 0.8, options.sampling_seed)?.with_top_k(Some(25), Some(250))?;
194    let dense_run = run_model(&mut dense, &prompt, &input_tokens, sampling, RunMode::Free)?;
195    drop(dense);
196
197    let quantized_load_start = Instant::now();
198    let mut quantized = load_model(&paths.quantized_model)?;
199    let quantized_load_seconds = quantized_load_start.elapsed().as_secs_f64();
200    validate_personaplex_geometry(&quantized)?;
201    let quantized_teacher = run_model(
202        &mut quantized,
203        &prompt,
204        &input_tokens,
205        RealtimeSampling::greedy(),
206        RunMode::TeacherForced(&dense_reference.frames),
207    )?;
208    let quality = quality_summary(&dense_reference.frames, &quantized_teacher.frames)?;
209    let quantized_run = run_model(
210        &mut quantized,
211        &prompt,
212        &input_tokens,
213        sampling,
214        RunMode::Free,
215    )?;
216
217    let decode_start = Instant::now();
218    let dense_pcm = decode_tokens(mimi, &dense_run.emitted_audio, context, input_pcm.len())?;
219    let quantized_pcm =
220        decode_tokens(mimi, &quantized_run.emitted_audio, context, input_pcm.len())?;
221    let decode_seconds = decode_start.elapsed().as_secs_f64();
222    let dense_tail = tail_max_rms_dbfs(&dense_pcm);
223    let quantized_tail = tail_max_rms_dbfs(&quantized_pcm);
224    let swap = options.sampling_seed & 1 == 1;
225    let (sample_a, sample_b, label_a, label_b, tail_a, tail_b) = if swap {
226        (
227            &quantized_pcm,
228            &dense_pcm,
229            "quantized",
230            "dense",
231            quantized_tail,
232            dense_tail,
233        )
234    } else {
235        (
236            &dense_pcm,
237            &quantized_pcm,
238            "dense",
239            "quantized",
240            dense_tail,
241            quantized_tail,
242        )
243    };
244    let truncated_a = tail_a > ACTIVE_AUDIO_DBFS;
245    let truncated_b = tail_b > ACTIVE_AUDIO_DBFS;
246    if truncated_a || truncated_b {
247        eprintln!(
248            "warning: generated speech is active at the output boundary; sample_a_tail_dbfs={tail_a:.1} sample_b_tail_dbfs={tail_b:.1}"
249        );
250    }
251    let dense_performance = performance_summary(&dense_run.latencies_ms);
252    let quantized_performance = performance_summary(&quantized_run.latencies_ms);
253    let divergence = free_run_agreement(&dense_run.frames, &quantized_run.frames);
254
255    fs::create_dir(&paths.output)?;
256    write_wav_pcm16(&paths.output.join("input.wav"), input_pcm, SAMPLE_RATE)?;
257    write_wav_pcm16(
258        &paths.output.join("input_codec_roundtrip.wav"),
259        &streaming_roundtrip,
260        SAMPLE_RATE,
261    )?;
262    write_wav_pcm16(
263        &paths.output.join("input_codec_roundtrip_offline.wav"),
264        &offline_roundtrip,
265        SAMPLE_RATE,
266    )?;
267    write_wav_pcm16(&paths.output.join("sample_a.wav"), sample_a, SAMPLE_RATE)?;
268    write_wav_pcm16(&paths.output.join("sample_b.wav"), sample_b, SAMPLE_RATE)?;
269
270    let metrics = json!({
271        "format_version": 1,
272        "methodology": "Both models use the public backend-neutral realtime scheduler, forcing, sampling, and observation contracts.",
273        "input": {
274            "path": paths.input,
275            "sample_rate": SAMPLE_RATE,
276            "frame_rate": FRAME_RATE,
277            "frames": frames,
278            "audio_seconds": frames as f64 / FRAME_RATE,
279            "tail_max_rms_dbfs": input_tail,
280            "likely_truncated": input_likely_truncated,
281            "warning": input_warning,
282        },
283        "conditioning": {
284            "voice_prompt_path": paths.voice_prompt,
285            "voice_prompt_frames": prompt.voice_frames.len(),
286            "text_tokenizer_path": paths.text_tokenizer,
287            "text_prompt": options.text_prompt,
288            "wrapped_text_prompt": wrapped_text_prompt,
289            "text_prompt_tokens": prompt.text_tokens.len(),
290            "silence_frames_after_voice": PROMPT_SILENCE_FRAMES,
291            "silence_frames_after_text": PROMPT_SILENCE_FRAMES,
292        },
293        "codec_diagnostic": {
294            "streaming_roundtrip": "input_codec_roundtrip.wav",
295            "offline_roundtrip": "input_codec_roundtrip_offline.wav",
296            "streaming_offline_token_agreement": codec_agreement,
297        },
298        "performance": {
299            "frame_deadline_ms": DEADLINE_MS,
300            "codec_encode_seconds": encode_seconds,
301            "codec_decode_both_outputs_seconds": decode_seconds,
302            "dense": { "load_seconds": dense_load_seconds, "model": dense_performance },
303            "quantized": { "load_seconds": quantized_load_seconds, "model": quantized_performance },
304        },
305        "teacher_forced_quality": quality,
306        "free_run_divergence_diagnostic": divergence,
307        "listening_test": {
308            "input": "input.wav",
309            "sample_a": "sample_a.wav",
310            "sample_b": "sample_b.wav",
311            "sampling": {
312                "seed": options.sampling_seed,
313                "text_temperature": sampling.text_temperature(),
314                "audio_temperature": sampling.audio_temperature(),
315                "text_top_k": sampling.text_top_k(),
316                "audio_top_k": sampling.audio_top_k(),
317            },
318            "sample_a_tail_max_rms_dbfs": tail_a,
319            "sample_b_tail_max_rms_dbfs": tail_b,
320            "sample_a_likely_truncated": truncated_a,
321            "sample_b_likely_truncated": truncated_b,
322            "input_warning": input_warning,
323        },
324    });
325    fs::write(
326        paths.output.join("metrics.json"),
327        serde_json::to_vec_pretty(&metrics)?,
328    )?;
329    fs::write(
330        paths.output.join("answer_key.json"),
331        serde_json::to_vec_pretty(&json!({ "sample_a": label_a, "sample_b": label_b }))?,
332    )?;
333    fs::write(
334        paths.output.join("listening_manifest.json"),
335        serde_json::to_vec_pretty(&json!({
336            "format_version": 1,
337            "trials": [{
338                "id": "personaplex_quantization_001",
339                "input": "input.wav",
340                "codec_roundtrip": "input_codec_roundtrip.wav",
341                "sample_a": "sample_a.wav",
342                "sample_b": "sample_b.wav",
343                "input_warning": input_warning,
344                "sample_a_likely_truncated": truncated_a,
345                "sample_b_likely_truncated": truncated_b,
346            }],
347        }))?,
348    )?;
349    fs::write(
350        paths.output.join("token_diagnostics.json"),
351        serde_json::to_vec_pretty(&json!({
352            "input": input_tokens,
353            "input_offline": offline_tokens,
354            "conditioning": {
355                "voice_prompt": prompt.voice_frames,
356                "text_prompt": prompt.text_tokens,
357                "silence_frames_after_voice": PROMPT_SILENCE_FRAMES,
358                "silence_frames_after_text": PROMPT_SILENCE_FRAMES,
359            },
360            "sampling": {
361                "seed": options.sampling_seed,
362                "text_temperature": sampling.text_temperature(),
363                "audio_temperature": sampling.audio_temperature(),
364                "text_top_k": sampling.text_top_k(),
365                "audio_top_k": sampling.audio_top_k(),
366            },
367            "dense_emitted": dense_run.emitted_audio,
368            "dense_sampled_frames": reference_tokens(&dense_run.frames),
369            "dense_greedy_emitted": dense_reference.emitted_audio,
370            "dense_greedy_frames": reference_tokens(&dense_reference.frames),
371            "quantized_emitted": quantized_run.emitted_audio,
372        }))?,
373    )?;
374    Ok(())
375}
376
377fn validate_personaplex_geometry<B: RealtimeBackend>(
378    model: &RealtimeModel<B>,
379) -> Result<(), Box<dyn Error>> {
380    let config = model.speech_config();
381    if config.input_audio_codebooks() != AUDIO_TOKENS_PER_STREAM
382        || config.generated_audio_codebooks() != AUDIO_TOKENS_PER_STREAM
383    {
384        return Err(invalid(format!(
385            "PersonaPlex evaluation requires {AUDIO_TOKENS_PER_STREAM} input and generated codebooks, got {} and {}",
386            config.input_audio_codebooks(),
387            config.generated_audio_codebooks()
388        )));
389    }
390    Ok(())
391}
392
393struct PromptConditioning {
394    voice_frames: Vec<Vec<i32>>,
395    text_tokens: Vec<i32>,
396}
397
398enum RunMode<'a> {
399    Free,
400    Diagnostics,
401    TeacherForced(&'a [ReferenceFrame]),
402}
403
404struct ReferenceFrame {
405    text_token: i32,
406    decision_audio: Vec<i32>,
407    sampled_audio: Vec<i32>,
408    diagnostics: Vec<Vec<f32>>,
409}
410
411struct ModelRun {
412    frames: Vec<ReferenceFrame>,
413    emitted_audio: Vec<Vec<i32>>,
414    latencies_ms: Vec<f64>,
415}
416
417fn run_model<B: RealtimeBackend>(
418    model: &mut RealtimeModel<B>,
419    prompt: &PromptConditioning,
420    input_tokens: &[Vec<i32>],
421    sampling: RealtimeSampling,
422    mode: RunMode<'_>,
423) -> Result<ModelRun, Box<dyn Error>> {
424    let request = RequestId::new(0);
425    let mut scheduler = RealtimeScheduler::new(model, SchedulerLimits::new(1, 1)?)?;
426    scheduler.register_request(model, request, sampling)?;
427    for frame in prompt_frames(prompt) {
428        run_frame(model, &mut scheduler, request, frame)?;
429    }
430    let mut frames = Vec::with_capacity(input_tokens.len());
431    let mut emitted_audio = Vec::new();
432    let mut latencies_ms = Vec::with_capacity(input_tokens.len());
433    for (index, tokens) in input_tokens.iter().enumerate() {
434        let mut frame = RealtimeInputFrame::new(1, tokens.clone());
435        match mode {
436            RunMode::Free => {}
437            RunMode::Diagnostics => frame = frame.with_diagnostics(),
438            RunMode::TeacherForced(reference) => {
439                let reference = reference
440                    .get(index)
441                    .ok_or_else(|| invalid("teacher-forced reference is shorter than input"))?;
442                frame = frame
443                    .with_forced_text(vec![reference.text_token])
444                    .with_forced_generated_audio(reference.sampled_audio.clone())
445                    .with_diagnostics();
446            }
447        }
448        let start = Instant::now();
449        let output = run_frame(model, &mut scheduler, request, frame)?;
450        latencies_ms.push(start.elapsed().as_secs_f64() * 1_000.0);
451        if let Some(tokens) = output.output_audio_tokens() {
452            emitted_audio.push(tokens.to_vec());
453        }
454        frames.push(ReferenceFrame {
455            text_token: *output
456                .text_tokens()
457                .first()
458                .ok_or_else(|| invalid("realtime output has no text token"))?,
459            decision_audio: output.decision_audio_tokens().to_vec(),
460            sampled_audio: output.sampled_audio_tokens().to_vec(),
461            diagnostics: output
462                .diagnostics()
463                .iter()
464                .map(|diagnostic| diagnostic.logits().to_vec())
465                .collect(),
466        });
467    }
468    scheduler.finish_request(request)?;
469    Ok(ModelRun {
470        frames,
471        emitted_audio,
472        latencies_ms,
473    })
474}
475
476fn run_frame<B: RealtimeBackend>(
477    model: &mut RealtimeModel<B>,
478    scheduler: &mut RealtimeScheduler<B>,
479    request: RequestId,
480    frame: RealtimeInputFrame,
481) -> Result<RealtimeOutputFrame, Box<dyn Error>> {
482    let input = model.backend().materialize_input(model.model(), &frame)?;
483    scheduler.enqueue(model, request, input)?;
484    loop {
485        if let Some(completed) = scheduler.run_queued(model)?.pop() {
486            return Ok(model.backend().observe_output(completed.output())?);
487        }
488        std::thread::yield_now();
489    }
490}
491
492fn prompt_frames(prompt: &PromptConditioning) -> Vec<RealtimeInputFrame> {
493    let forced = |audio: Vec<i32>, text: i32| {
494        RealtimeInputFrame::new(1, SINE_TOKENS.to_vec())
495            .with_forced_generated_audio(audio)
496            .with_forced_text(vec![text])
497    };
498    let mut frames = prompt
499        .voice_frames
500        .iter()
501        .cloned()
502        .map(|audio| forced(audio, TEXT_PADDING_TOKEN))
503        .collect::<Vec<_>>();
504    frames.extend(
505        std::iter::repeat_with(|| forced(SILENCE_TOKENS.to_vec(), TEXT_PADDING_TOKEN))
506            .take(PROMPT_SILENCE_FRAMES),
507    );
508    frames.extend(
509        prompt
510            .text_tokens
511            .iter()
512            .map(|token| forced(SILENCE_TOKENS.to_vec(), *token)),
513    );
514    frames.extend(
515        std::iter::repeat_with(|| forced(SILENCE_TOKENS.to_vec(), TEXT_PADDING_TOKEN))
516            .take(PROMPT_SILENCE_FRAMES),
517    );
518    frames
519}
520
521fn encode_pcm<T: Tensor>(
522    mimi: &mut Mimi<T>,
523    pcm: &[f32],
524    context: &T::Context,
525) -> Result<Vec<Vec<i32>>, Box<dyn Error>> {
526    mimi.reset_encode_state();
527    let mut frames = Vec::with_capacity(pcm.len() / FRAME_SAMPLES);
528    for frame in pcm.as_chunks::<FRAME_SAMPLES>().0 {
529        let frame = T::from_f32_slice(frame, &[1, 1, FRAME_SAMPLES as i32], context)?;
530        if let Some(tokens) = mimi.encode_step(&frame, context)? {
531            frames.push(tokens.to_i32_vec(context)?);
532        }
533    }
534    Ok(frames)
535}
536
537fn decode_tokens<T: Tensor>(
538    mimi: &mut Mimi<T>,
539    frames: &[Vec<i32>],
540    context: &T::Context,
541    target_samples: usize,
542) -> Result<Vec<f32>, Box<dyn Error>> {
543    mimi.reset_decode_state();
544    let mut pcm = Vec::with_capacity(target_samples);
545    for frame in frames {
546        let tokens = T::from_i32_slice(frame, &[1, frame.len() as i32], context)?;
547        pcm.extend(mimi.decode_step(&tokens, context)?.to_f32_vec(context)?);
548    }
549    pcm.truncate(target_samples);
550    pcm.resize(target_samples, 0.0);
551    Ok(pcm)
552}
553
554fn offline_roundtrip<T: Tensor>(
555    mimi: &mut Mimi<T>,
556    pcm: &[f32],
557    context: &T::Context,
558) -> Result<OfflineRoundtrip, Box<dyn Error>> {
559    let input = T::from_f32_slice(pcm, &[1, 1, pcm.len() as i32], context)?;
560    let codes = mimi.encode(&input, context)?;
561    let code_shape = codes.shape().to_vec();
562    if code_shape.len() != 3 || code_shape[0] != 1 {
563        return Err(invalid(format!(
564            "offline Mimi codes have unexpected shape {code_shape:?}"
565        )));
566    }
567    let values = codes.to_i32_vec(context)?;
568    let codebooks = code_shape[1] as usize;
569    let frame_count = code_shape[2] as usize;
570    let mut frames = vec![vec![0; codebooks]; frame_count];
571    for codebook in 0..codebooks {
572        for frame in 0..frame_count {
573            frames[frame][codebook] = values[codebook * frame_count + frame];
574        }
575    }
576    let mut roundtrip = mimi.decode(&codes, context)?.to_f32_vec(context)?;
577    roundtrip.truncate(pcm.len());
578    roundtrip.resize(pcm.len(), 0.0);
579    Ok(OfflineRoundtrip {
580        tokens: frames,
581        pcm: roundtrip,
582    })
583}
584
585struct OfflineRoundtrip {
586    tokens: Vec<Vec<i32>>,
587    pcm: Vec<f32>,
588}
589
590#[derive(Debug, Clone, Default)]
591struct DistributionAccumulator {
592    count: usize,
593    target_count: usize,
594    kl_sum: f64,
595    entropy_sum: f64,
596    target_nll_delta_sum: f64,
597    centered_rmse_sum: f64,
598    top1_matches: usize,
599    top5_overlap_sum: f64,
600}
601
602impl DistributionAccumulator {
603    fn update(
604        &mut self,
605        dense: &[f32],
606        candidate: &[f32],
607        target: usize,
608    ) -> Result<(), Box<dyn Error>> {
609        let metrics = compare_distributions(
610            dense,
611            candidate,
612            (target < dense.len()).then_some(target),
613            5,
614        )?;
615        self.count += 1;
616        self.kl_sum += metrics.kl_nats;
617        self.entropy_sum += metrics.reference_entropy_nats;
618        self.centered_rmse_sum += metrics.centered_logit_rmse;
619        self.top1_matches += usize::from(metrics.top1_agreement);
620        self.top5_overlap_sum += metrics.top_k_overlap;
621        if let Some(delta) = metrics.target_nll_delta_nats {
622            self.target_count += 1;
623            self.target_nll_delta_sum += delta;
624        }
625        Ok(())
626    }
627
628    fn merge(&mut self, other: &Self) {
629        self.count += other.count;
630        self.target_count += other.target_count;
631        self.kl_sum += other.kl_sum;
632        self.entropy_sum += other.entropy_sum;
633        self.target_nll_delta_sum += other.target_nll_delta_sum;
634        self.centered_rmse_sum += other.centered_rmse_sum;
635        self.top1_matches += other.top1_matches;
636        self.top5_overlap_sum += other.top5_overlap_sum;
637    }
638
639    fn summary(&self) -> MetricSummary {
640        let count = self.count.max(1) as f64;
641        MetricSummary {
642            distributions: self.count,
643            target_distributions: self.target_count,
644            mean_kl_nats: self.kl_sum / count,
645            mean_dense_entropy_nats: self.entropy_sum / count,
646            mean_target_nll_delta_nats: self.target_nll_delta_sum / self.target_count.max(1) as f64,
647            mean_centered_logit_rmse: self.centered_rmse_sum / count,
648            top1_agreement: self.top1_matches as f64 / count,
649            mean_top5_overlap: self.top5_overlap_sum / count,
650        }
651    }
652}
653
654#[derive(Debug, Clone, Serialize)]
655struct MetricSummary {
656    distributions: usize,
657    target_distributions: usize,
658    mean_kl_nats: f64,
659    mean_dense_entropy_nats: f64,
660    mean_target_nll_delta_nats: f64,
661    mean_centered_logit_rmse: f64,
662    top1_agreement: f64,
663    mean_top5_overlap: f64,
664}
665
666#[derive(Debug, Clone, Serialize)]
667struct QualitySummary {
668    methodology: &'static str,
669    text: MetricSummary,
670    audio_generated: MetricSummary,
671    audio_input_conditioned: MetricSummary,
672    audio_overall: MetricSummary,
673    audio_by_codebook: Vec<MetricSummary>,
674}
675
676fn quality_summary(
677    dense: &[ReferenceFrame],
678    candidate: &[ReferenceFrame],
679) -> Result<QualitySummary, Box<dyn Error>> {
680    if dense.len() != candidate.len() {
681        return Err(invalid("teacher-forced run lengths differ"));
682    }
683    let mut text = DistributionAccumulator::default();
684    let mut audio = Vec::<DistributionAccumulator>::new();
685    for (dense, candidate) in dense.iter().zip(candidate) {
686        if dense.diagnostics.len() != candidate.diagnostics.len() || dense.diagnostics.is_empty() {
687            return Err(invalid(
688                "teacher-forced diagnostic counts differ or are empty",
689            ));
690        }
691        text.update(
692            &dense.diagnostics[0],
693            &candidate.diagnostics[0],
694            dense.text_token as usize,
695        )?;
696        if audio.is_empty() {
697            audio.resize(
698                dense.diagnostics.len() - 1,
699                DistributionAccumulator::default(),
700            );
701        }
702        for (codebook, accumulator) in audio.iter_mut().enumerate() {
703            accumulator.update(
704                &dense.diagnostics[codebook + 1],
705                &candidate.diagnostics[codebook + 1],
706                *dense
707                    .decision_audio
708                    .get(codebook)
709                    .ok_or_else(|| invalid("teacher-forced decision token is missing"))?
710                    as usize,
711            )?;
712        }
713    }
714    let mut overall = DistributionAccumulator::default();
715    for value in &audio {
716        overall.merge(value);
717    }
718    let mut generated = DistributionAccumulator::default();
719    for value in audio.iter().take(AUDIO_TOKENS_PER_STREAM) {
720        generated.merge(value);
721    }
722    let mut input_conditioned = DistributionAccumulator::default();
723    for value in audio.iter().skip(AUDIO_TOKENS_PER_STREAM) {
724        input_conditioned.merge(value);
725    }
726    Ok(QualitySummary {
727        methodology: "The candidate is teacher-forced onto the dense model's exact text and generated-audio history; KL uses the dense distribution as reference.",
728        text: text.summary(),
729        audio_generated: generated.summary(),
730        audio_input_conditioned: input_conditioned.summary(),
731        audio_overall: overall.summary(),
732        audio_by_codebook: audio.iter().map(DistributionAccumulator::summary).collect(),
733    })
734}
735
736#[derive(Debug, Clone, Serialize)]
737struct PerformanceSummary {
738    frames: usize,
739    mean_ms: f64,
740    p50_ms: f64,
741    p95_ms: f64,
742    max_ms: f64,
743    deadline_misses: usize,
744}
745
746fn performance_summary(latencies: &[f64]) -> PerformanceSummary {
747    let summary = summarize_latencies(latencies, Some(DEADLINE_MS))
748        .expect("every model run records at least one finite nonnegative latency");
749    PerformanceSummary {
750        frames: summary.samples,
751        mean_ms: summary.mean_ms,
752        p50_ms: summary.p50_ms,
753        p95_ms: summary.p95_ms,
754        max_ms: summary.max_ms,
755        deadline_misses: summary.deadline_misses,
756    }
757}
758
759fn free_run_agreement(dense: &[ReferenceFrame], quantized: &[ReferenceFrame]) -> serde_json::Value {
760    let frames = dense.len().min(quantized.len());
761    let text_matches = dense
762        .iter()
763        .zip(quantized)
764        .filter(|(left, right)| left.text_token == right.text_token)
765        .count();
766    let mut audio_matches = 0usize;
767    let mut audio_total = 0usize;
768    for (left, right) in dense.iter().zip(quantized) {
769        for (left, right) in left.sampled_audio.iter().zip(&right.sampled_audio) {
770            audio_matches += usize::from(left == right);
771            audio_total += 1;
772        }
773    }
774    json!({
775        "frames": frames,
776        "text_token_agreement": text_matches as f64 / frames.max(1) as f64,
777        "audio_token_agreement": audio_matches as f64 / audio_total.max(1) as f64,
778    })
779}
780
781fn reference_tokens(frames: &[ReferenceFrame]) -> Vec<serde_json::Value> {
782    frames
783        .iter()
784        .map(|frame| json!({ "text": frame.text_token, "sampled_audio": frame.sampled_audio }))
785        .collect()
786}
787
788fn token_frame_agreement(left: &[Vec<i32>], right: &[Vec<i32>]) -> f64 {
789    let mut matches = 0usize;
790    let mut total = 0usize;
791    for (left, right) in left.iter().zip(right) {
792        for (left, right) in left.iter().zip(right) {
793            matches += usize::from(left == right);
794            total += 1;
795        }
796    }
797    matches as f64 / total.max(1) as f64
798}
799
800fn read_f32le(path: &Path) -> Result<Vec<f32>, Box<dyn Error>> {
801    let bytes = fs::read(path)?;
802    if bytes.len() % 4 != 0 {
803        return Err(invalid(format!(
804            "raw f32le input length must be divisible by four, got {} bytes",
805            bytes.len()
806        )));
807    }
808    Ok(bytes
809        .as_chunks::<4>()
810        .0
811        .iter()
812        .map(|chunk| f32::from_le_bytes(*chunk))
813        .collect())
814}
815
816fn rms_dbfs(samples: &[f32]) -> f64 {
817    let mean_square = samples
818        .iter()
819        .map(|sample| (*sample as f64) * (*sample as f64))
820        .sum::<f64>()
821        / samples.len().max(1) as f64;
822    20.0 * mean_square.sqrt().max(1e-12).log10()
823}
824
825fn tail_max_rms_dbfs(samples: &[f32]) -> f64 {
826    samples
827        .as_chunks::<FRAME_SAMPLES>()
828        .0
829        .iter()
830        .rev()
831        .take(TAIL_ACTIVITY_FRAMES)
832        .map(|frame| rms_dbfs(frame))
833        .fold(f64::NEG_INFINITY, f64::max)
834}
835
836fn write_wav_pcm16(path: &Path, samples: &[f32], sample_rate: u32) -> Result<(), Box<dyn Error>> {
837    let data_bytes = u32::try_from(
838        samples
839            .len()
840            .checked_mul(2)
841            .ok_or_else(|| invalid("WAV size overflow"))?,
842    )?;
843    let mut file = fs::File::create(path)?;
844    file.write_all(b"RIFF")?;
845    file.write_all(&(36u32 + data_bytes).to_le_bytes())?;
846    file.write_all(b"WAVEfmt ")?;
847    file.write_all(&16u32.to_le_bytes())?;
848    file.write_all(&1u16.to_le_bytes())?;
849    file.write_all(&1u16.to_le_bytes())?;
850    file.write_all(&sample_rate.to_le_bytes())?;
851    file.write_all(&(sample_rate * 2).to_le_bytes())?;
852    file.write_all(&2u16.to_le_bytes())?;
853    file.write_all(&16u16.to_le_bytes())?;
854    file.write_all(b"data")?;
855    file.write_all(&data_bytes.to_le_bytes())?;
856    for sample in samples {
857        let value = (sample.clamp(-1.0, 1.0) * i16::MAX as f32).round() as i16;
858        file.write_all(&value.to_le_bytes())?;
859    }
860    Ok(())
861}
862
863fn invalid(message: impl Into<String>) -> Box<dyn Error> {
864    Box::new(io::Error::new(io::ErrorKind::InvalidInput, message.into()))
865}
866
867#[cfg(test)]
868mod tests {
869    use super::*;
870
871    #[test]
872    fn identical_distribution_metrics_are_exact() {
873        let values = [0.0, 1.0, -1.0, 0.5, 0.25];
874        let mut metric = DistributionAccumulator::default();
875        metric.update(&values, &values, 1).unwrap();
876        let summary = metric.summary();
877        assert!(summary.mean_kl_nats.abs() < 1e-12);
878        assert_eq!(summary.top1_agreement, 1.0);
879        assert_eq!(summary.mean_top5_overlap, 1.0);
880    }
881
882    #[test]
883    fn prompt_frames_preserve_released_conditioning_order() {
884        let prompt = PromptConditioning {
885            voice_frames: vec![vec![1; AUDIO_TOKENS_PER_STREAM]],
886            text_tokens: vec![7, 8],
887        };
888        let frames = prompt_frames(&prompt);
889        assert_eq!(
890            frames.len(),
891            1 + PROMPT_SILENCE_FRAMES + 2 + PROMPT_SILENCE_FRAMES
892        );
893        assert_eq!(frames[0].forced_generated_audio_tokens(), Some(&[1; 8][..]));
894        assert_eq!(
895            frames[1 + PROMPT_SILENCE_FRAMES].forced_text_tokens(),
896            Some(&[7][..])
897        );
898    }
899}