use anyhow::{Context, Result};
use ndarray::Array3;
use rlx_dac::ops::weight_norm;
use rlx_dac::weights::WeightStore;
use rlx_tsac::default_tsac_dir;
use rlx_tsac::rlx_decode::{q8_conv_std, q8_convt_std, q8_f32, q8_v_raw};
fn stats(name: &str, real: &[f32], mine: &[f32]) {
let n = real.len().min(mine.len());
let (mut dot, mut rr, mut mm, mut err, mut mx) = (0f64, 0f64, 0f64, 0f64, 0f64);
for i in 0..n {
let (a, b) = (real[i] as f64, mine[i] as f64);
dot += a * b;
rr += a * a;
mm += b * b;
err += (a - b) * (a - b);
mx = mx.max((a - b).abs());
}
let corr = dot / (rr.sqrt() * mm.sqrt()).max(1e-12);
let rel = (err / rr.max(1e-12)).sqrt();
eprintln!(
"{name:42} len={n:>8} corr={corr:.4} rel_rms_err={rel:.4} max_abs_err={mx:.4} (real_rms={:.4} mine_rms={:.4})",
(rr / n as f64).sqrt(),
(mm / n as f64).sqrt(),
);
}
fn real_conv(store: &WeightStore, prefix: &str) -> Result<Array3<f32>> {
let (g, gs) = store.get(&format!("{prefix}.weight_g"))?;
let (v, vs) = store.get(&format!("{prefix}.weight_v"))?;
let gv = ndarray::ArrayView3::from_shape((gs[0], gs[1], gs[2]), g)?;
let vv = ndarray::ArrayView3::from_shape((vs[0], vs[1], vs[2]), v)?;
Ok(weight_norm(gv, vv))
}
fn main() -> Result<()> {
let tdir = default_tsac_dir();
let ddir = std::env::var("RLX_DAC_DIR").unwrap_or_else(|_| ".cache/dac44".into());
let ddir = std::path::PathBuf::from(ddir);
std::fs::create_dir_all(&ddir).ok();
let st_path = ddir.join("model.safetensors");
if !st_path.is_file() {
rlx_dac::download::fetch_dac(&ddir, "44khz").context("download real DAC-44kHz")?;
}
let store = WeightStore::open(&st_path)?;
eprintln!("=== real DAC key sample ===");
for k in store.keys().take(6) {
eprintln!(" {k}");
}
{
let (rv, rs) = store.get("decoder.model.0.weight_v")?; let (mv, md) = q8_v_raw(&tdir, "decoder.model.0")?; eprintln!("raw v: real shape={rs:?} q8 dims={md:?}");
let perms = [
[0, 1, 2],
[0, 2, 1],
[1, 0, 2],
[1, 2, 0],
[2, 0, 1],
[2, 1, 0],
];
let (d0, d1, d2) = (md[0], md[1], md[2]);
for p in perms {
let pd = [md[p[0]], md[p[1]], md[p[2]]];
if pd != [rs[0], rs[1], rs[2]] {
continue; }
let mut perm_flat = vec![0f32; mv.len()];
let mut idx = 0;
for a in 0..pd[0] {
for b in 0..pd[1] {
for c in 0..pd[2] {
let coord = [a, b, c];
let mut q = [0usize; 3];
for (pos, &ax) in p.iter().enumerate() {
q[ax] = coord[pos];
}
perm_flat[idx] = mv[(q[0] * d1 + q[1]) * d2 + q[2]];
idx += 1;
}
}
}
let mut dot = 0f64;
let mut rr = 0f64;
let mut mm = 0f64;
for i in 0..rv.len() {
let (x, y) = (rv[i] as f64, perm_flat[i] as f64);
dot += x * y;
rr += x * x;
mm += y * y;
}
let corr = dot / (rr.sqrt() * mm.sqrt()).max(1e-12);
eprintln!(" perm {p:?} -> shape {pd:?} corr={corr:.4}");
}
}
for prefix in [
"decoder.model.0",
"decoder.model.6",
"encoder.block.0",
"encoder.block.6",
] {
match (real_conv(&store, prefix), q8_conv_std(&tdir, prefix)) {
(Ok(r), Ok(m)) => {
eprintln!("{prefix}: real{:?} mine{:?}", r.dim(), m.dim());
stats(prefix, r.as_slice().unwrap(), m.as_slice().unwrap());
}
(r, m) => eprintln!("{prefix}: real_ok={} mine_ok={}", r.is_ok(), m.is_ok()),
}
}
for prefix in ["decoder.model.1.block.1", "decoder.model.4.block.1"] {
match (real_conv(&store, prefix), q8_convt_std(&tdir, prefix)) {
(Ok(r), Ok(m)) => {
eprintln!("{prefix} (convT): real{:?} mine{:?}", r.dim(), m.dim());
stats(prefix, r.as_slice().unwrap(), m.as_slice().unwrap());
}
(r, m) => eprintln!("{prefix}: real_ok={} mine_ok={}", r.is_ok(), m.is_ok()),
}
}
let cbk = "quantizer.quantizers.0.codebook.weight";
if let (Ok((r, _)), Ok(m)) = (store.get(cbk), q8_f32(&tdir, cbk)) {
stats(cbk, r, &m);
}
Ok(())
}