Files
rustytorch/crates/models/rtx-csm/examples/silentcipher_smoke.rs
T
osobhandClaude Opus 4.7 a5cedfb46a rtx-csm: emotional_speech_guide — CREMA-D vs RAVDESS firdhokk verdict
8-gen bench (4 emotions × 2 corpora) at seed=42 against firdhokk
Whisper-LV3:

  target    RAVDESS              CREMA-D
  happy     happy (0.999) ✓      happy (0.999) ✓
  angry     neutral (0.92)       sad (0.99)
  fearful   happy (0.998)        fearful (0.984) ✓
  sad       angry (0.99)         fearful (0.99)

CREMA-D 2/4 vs RAVDESS 1/4. Larger / more naturalistic corpus
produces more class-pure fearful direction. Neither corpus solves
angry or sad — recipe shifts into 'vague expressivity' rather than
class-specific corners.

Practical: prefer CREMA-D when available; A/B both per emotion if
class precision matters.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
2026-04-30 00:01:02 -07:00

215 lines
7.9 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Phase 10.3 smoke test for SilentCipher embed. Downloads the
//! 16 kHz checkpoint from `sony/silentcipher`, loads all 3 networks,
//! runs the encode pipeline on a synthetic sine + a real speech clip,
//! verifies the output preserves shape + has signal.
//!
//! Stops short of detection (Phase 10.4 will add detect + ship the
//! Watermarker trait impl).
//!
//! Usage:
//! ```bash
//! cargo run -p rtx-csm --release --features metal --example silentcipher_smoke
//! ```
use anyhow::{Context, Result};
use candle_core::Device;
use hf_hub::api::sync::Api;
use rtx_csm::silentcipher::{SilentCipherConfig, SilentCipherWatermarker};
use std::time::Instant;
const REPO: &str = "sony/silentcipher";
const CKPT_DIR: &str = "16_khz/97561_iteration";
fn main() -> Result<()> {
tracing_subscriber::fmt()
.with_max_level(tracing::Level::WARN)
.init();
let device = if candle_core::utils::metal_is_available() {
Device::new_metal(0)?
} else {
Device::Cpu
};
eprintln!("device: {device:?}");
// Download the 3 ckpts + hparams.
let api = Api::new()?;
let repo = api.model(REPO.to_string());
let enc_c = repo
.get(&format!("{CKPT_DIR}/enc_c.ckpt"))
.context("download enc_c.ckpt")?;
let dec_c = repo
.get(&format!("{CKPT_DIR}/dec_c.ckpt"))
.context("download dec_c.ckpt")?;
let dec_m_0 = repo
.get(&format!("{CKPT_DIR}/dec_m_0.ckpt"))
.context("download dec_m_0.ckpt")?;
eprintln!("downloaded: {}", enc_c.parent().unwrap().display());
// Construct.
let cfg = SilentCipherConfig::sixteen_khz();
let load_t = Instant::now();
let wm = SilentCipherWatermarker::from_ckpts(cfg, &enc_c, &dec_c, &dec_m_0, &device)?;
eprintln!("watermarker built in {} ms", load_t.elapsed().as_millis());
// Synthetic 1 s 440 Hz sine at 16 kHz.
let n = 16_000;
let samples: Vec<f32> = (0..n)
.map(|i| (2.0 * std::f32::consts::PI * 440.0 * i as f32 / 16_000.0).sin() * 0.3)
.collect();
eprintln!();
eprintln!("=== synthetic sine ===");
eprintln!("input samples: {}", samples.len());
// 16 kHz model carries ~23.78 bits per patch (15 base-3 codes).
// Pick a payload that's well below 3^15 = 14_348_907.
let payload = 12345678u32;
let codes = wm.encode_bits(payload);
eprintln!(
"payload 0x{payload:08X} -> {} codes (each in 0..{}): {:?}",
codes.len(),
wm.cfg.message_dim,
&codes
);
let t = Instant::now();
let watermarked = wm.embed(&samples, &codes).context("embed sine")?;
let embed_ms = t.elapsed().as_millis();
eprintln!(
"embed: {embed_ms} ms ({} samples in / {} samples out)",
samples.len(),
watermarked.len()
);
// Sanity: shape preserved + samples differ from original (watermark
// present) but not catastrophically.
assert_eq!(
samples.len(),
watermarked.len(),
"embed must preserve length"
);
let diff_rms: f32 = samples
.iter()
.zip(watermarked.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt()
/ (samples.len() as f32).sqrt();
let orig_rms: f32 = (samples.iter().map(|&s| s * s).sum::<f32>() / samples.len() as f32).sqrt();
let ratio = diff_rms / orig_rms;
eprintln!("RMS diff vs original: {diff_rms:.5} (orig RMS {orig_rms:.4}, ratio {ratio:.4})");
if ratio < 1e-5 {
println!("WARN: watermark RMS is near-zero — model may not be loading correctly");
} else if ratio > 0.5 {
println!("WARN: watermark RMS is huge — corruption likely");
} else {
println!("PASS: shape preserved + watermark added (ratio {ratio:.4})");
}
// Try a real speech clip if /tmp/asr_test.flac exists.
let asr_path = std::path::Path::new("/tmp/asr_test.flac");
if asr_path.exists() {
eprintln!();
eprintln!("=== real speech (/tmp/asr_test.flac) ===");
let speech = rtx_csm::audio_io::load_mono_at_rate(asr_path, 16_000)?;
eprintln!(
"input samples: {} ({:.2}s)",
speech.len(),
speech.len() as f32 / 16_000.0
);
let t = Instant::now();
let watermarked = wm.embed(&speech, &codes).context("embed speech")?;
let speech_ms = t.elapsed().as_millis();
eprintln!(
"embed: {speech_ms} ms ({:.2}× realtime)",
speech_ms as f32 / (speech.len() as f32 / 16_000.0 * 1000.0)
);
let diff_rms: f32 = speech
.iter()
.zip(watermarked.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt()
/ (speech.len() as f32).sqrt();
let orig_rms: f32 =
(speech.iter().map(|&s| s * s).sum::<f32>() / speech.len() as f32).sqrt();
let snr_db = 20.0 * (orig_rms / diff_rms.max(1e-12)).log10();
println!("speech embed: SNR vs original = {snr_db:.1} dB (target ~47 dB per hparams)");
// Save the output WAV for ear test.
rtx_csm::audio_io::write_wav_mono(
std::path::Path::new("/tmp/silentcipher_smoke.wav"),
&watermarked,
16_000,
)?;
eprintln!("wrote /tmp/silentcipher_smoke.wav for ear test");
// Phase 10.4 round-trip: detect on the watermarked audio,
// verify we recover the embedded codes.
let detect_t = Instant::now();
let result = wm.detect(&watermarked).context("detect speech")?;
let detect_ms = detect_t.elapsed().as_millis();
eprintln!();
eprintln!("=== detect round-trip ===");
eprintln!(
"detect: {detect_ms} ms ({:.4}× realtime)",
detect_ms as f32 / (speech.len() as f32 / 16_000.0 * 1000.0)
);
eprintln!("recovered codes: {:?}", result.codes);
eprintln!("confidence: {:.3}", result.confidence);
let recovered_payload = wm.decode_bits(&result.codes);
eprintln!("recovered payload: 0x{recovered_payload:08X}");
eprintln!("expected payload: 0x{payload:08X}");
let matching_codes = result
.codes
.iter()
.zip(codes.iter())
.filter(|(a, b)| a == b)
.count();
eprintln!(
"matching codes: {} / {} ({:.1}%)",
matching_codes,
codes.len(),
100.0 * matching_codes as f32 / codes.len() as f32
);
if matching_codes == codes.len() {
println!("PASS: full round-trip recovers the embedded payload");
} else if matching_codes >= codes.len() * 3 / 4 {
println!(
"PARTIAL: most codes recovered ({matching_codes}/{}); good enough for confidence-based detection",
codes.len()
);
} else {
println!(
"WEAK: only {matching_codes}/{} codes recovered. Likely a normalization bug (Phase 10.5 to fix)",
codes.len()
);
}
// Also detect on UN-watermarked audio — should NOT recover the
// payload (or low confidence).
let result_clean = wm.detect(&speech).context("detect clean")?;
eprintln!();
eprintln!("=== detect on un-watermarked audio ===");
eprintln!("recovered codes: {:?}", result_clean.codes);
eprintln!("confidence: {:.3}", result_clean.confidence);
let clean_matching = result_clean
.codes
.iter()
.zip(codes.iter())
.filter(|(a, b)| a == b)
.count();
if clean_matching < codes.len() / 2 {
println!(
"PASS: clean audio does NOT recover the payload ({clean_matching}/{} match)",
codes.len()
);
} else {
println!(
"WARN: clean audio matches {clean_matching}/{} codes — confidence threshold needed for false-positive control",
codes.len()
);
}
}
Ok(())
}