1#![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
59pub const DEFAULT_TEXT_PROMPT: &str = "You are a wise and friendly teacher. Answer questions or provide advice in a clear and engaging way.";
61pub const DEFAULT_SAMPLING_SEED: u64 = 20_260_713;
63
64#[derive(Debug, Clone)]
66pub struct PersonaPlexEvaluationPaths {
67 pub dense_model: PathBuf,
69 pub quantized_model: PathBuf,
71 pub text_tokenizer: PathBuf,
73 pub voice_prompt: PathBuf,
75 pub input: PathBuf,
77 pub output: PathBuf,
79}
80
81#[derive(Debug, Clone)]
83pub struct PersonaPlexEvaluationOptions {
84 pub frames: Option<usize>,
86 pub text_prompt: String,
88 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
102pub 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}