use anyhow::{Context, Result};
use rlx_dac::DacCodec;
use rlx_runtime::Device;
use std::path::{Path, PathBuf};
pub fn best_device() -> Device {
for d in [Device::Mlx, Device::Metal, Device::Gpu] {
if rlx_runtime::is_available(d) {
return d;
}
}
Device::Cpu
}
pub fn default_dir() -> PathBuf {
std::env::var("RLX_DAC_DIR")
.map(PathBuf::from)
.unwrap_or_else(|_| PathBuf::from(".cache/dac44"))
}
pub fn open(device: Device) -> Result<DacCodec> {
open_in(&default_dir(), device)
}
pub fn open_in(model_dir: &Path, device: Device) -> Result<DacCodec> {
rlx_dac::download::ensure_weights(model_dir)?;
DacCodec::open_on(model_dir, device)
}
pub fn weights_available(dir: &Path) -> bool {
dir.join("model.safetensors").is_file() && dir.join("config.json").is_file()
}
pub struct CorrectCodec {
dac: DacCodec,
num_quantizers: Option<usize>,
}
const MAGIC_V1: &[u8; 4] = b"TSRX";
const MAGIC_V2: &[u8; 4] = b"TSR2";
fn bits_per_code(codebook_size: usize) -> u32 {
codebook_size.next_power_of_two().trailing_zeros().max(1)
}
struct BitWriter {
buf: Vec<u8>,
acc: u64,
nbits: u32,
}
impl BitWriter {
fn new() -> Self {
Self {
buf: Vec::new(),
acc: 0,
nbits: 0,
}
}
fn put(&mut self, value: u32, bits: u32) {
self.acc |= (value as u64) << self.nbits;
self.nbits += bits;
while self.nbits >= 8 {
self.buf.push((self.acc & 0xff) as u8);
self.acc >>= 8;
self.nbits -= 8;
}
}
fn finish(mut self) -> Vec<u8> {
if self.nbits > 0 {
self.buf.push((self.acc & 0xff) as u8);
}
self.buf
}
}
struct BitReader<'a> {
bytes: &'a [u8],
pos: usize,
acc: u64,
nbits: u32,
}
impl<'a> BitReader<'a> {
fn new(bytes: &'a [u8]) -> Self {
Self {
bytes,
pos: 0,
acc: 0,
nbits: 0,
}
}
fn get(&mut self, bits: u32) -> u32 {
while self.nbits < bits {
let byte = self.bytes.get(self.pos).copied().unwrap_or(0);
self.pos += 1;
self.acc |= (byte as u64) << self.nbits;
self.nbits += 8;
}
let mask = if bits >= 32 {
u32::MAX
} else {
(1u32 << bits) - 1
};
let v = (self.acc as u32) & mask;
self.acc >>= bits;
self.nbits -= bits;
v
}
}
impl CorrectCodec {
pub fn open(device: Device, quality: Option<u8>) -> Result<Self> {
Self::open_in(&default_dir(), device, quality)
}
pub fn open_in(model_dir: &Path, device: Device, quality: Option<u8>) -> Result<Self> {
let dac = open_in(model_dir, device)?;
let num_quantizers = quality.map(|q| (q as usize).clamp(1, 9));
Ok(Self {
dac,
num_quantizers,
})
}
pub fn sample_rate(&self) -> u32 {
self.dac.sample_rate()
}
pub fn encode_file(&self, in_audio: &Path, out_tsac: &Path) -> Result<()> {
let codes = self.dac.encode_wav(in_audio, self.num_quantizers)?;
let orig = mono_len_44k(in_audio)?;
let bits = bits_per_code(self.dac.config().codebook_size);
let mut buf = Vec::with_capacity(
21 + (codes.num_frames() * codes.num_quantizers * bits as usize).div_ceil(8),
);
buf.extend_from_slice(MAGIC_V2);
buf.extend_from_slice(&self.dac.sample_rate().to_le_bytes());
buf.extend_from_slice(&(orig as u32).to_le_bytes());
buf.extend_from_slice(&(codes.num_frames() as u32).to_le_bytes());
buf.extend_from_slice(&(codes.num_quantizers as u32).to_le_bytes());
buf.push(bits as u8);
let mut bw = BitWriter::new();
for frame in &codes.frames {
for &c in frame {
bw.put(c, bits);
}
}
buf.extend_from_slice(&bw.finish());
if let Some(parent) = out_tsac.parent() {
if !parent.as_os_str().is_empty() {
std::fs::create_dir_all(parent).ok();
}
}
std::fs::write(out_tsac, buf).with_context(|| format!("write {}", out_tsac.display()))?;
Ok(())
}
pub fn decode_file(&self, in_tsac: &Path, out_wav: &Path) -> Result<()> {
let b = std::fs::read(in_tsac).with_context(|| format!("read {}", in_tsac.display()))?;
anyhow::ensure!(b.len() >= 20, "container too short");
let rd = |o: usize| u32::from_le_bytes([b[o], b[o + 1], b[o + 2], b[o + 3]]) as usize;
let orig = rd(8);
let n_frames = rd(12);
let n_cb = rd(16);
let frames = if &b[0..4] == MAGIC_V2 {
anyhow::ensure!(b.len() >= 21, "TSR2 truncated header");
let bits = b[20] as u32;
anyhow::ensure!((1..=16).contains(&bits), "TSR2 bad bits_per_code {bits}");
let need = (n_frames * n_cb * bits as usize).div_ceil(8);
anyhow::ensure!(b.len() >= 21 + need, "TSR2 truncated payload");
let mut br = BitReader::new(&b[21..]);
let mut frames = Vec::with_capacity(n_frames);
for _ in 0..n_frames {
let mut row = Vec::with_capacity(n_cb);
for _ in 0..n_cb {
row.push(br.get(bits));
}
frames.push(row);
}
frames
} else if &b[0..4] == MAGIC_V1 {
anyhow::ensure!(b.len() >= 20 + n_frames * n_cb * 2, "TSRX truncated");
let mut frames = Vec::with_capacity(n_frames);
let mut off = 20;
for _ in 0..n_frames {
let mut row = Vec::with_capacity(n_cb);
for _ in 0..n_cb {
row.push(u16::from_le_bytes([b[off], b[off + 1]]) as u32);
off += 2;
}
frames.push(row);
}
frames
} else {
anyhow::bail!("not a TSRX/TSR2 container");
};
let codes = rlx_dac::DacCodes {
frames,
num_quantizers: n_cb,
};
self.dac.decode_wav(&codes, out_wav, Some(orig))
}
}
fn mono_len_44k(path: &Path) -> Result<usize> {
let r = hound::WavReader::open(path).with_context(|| format!("open {}", path.display()))?;
let s = r.spec();
let frames = (r.len() as usize) / (s.channels as usize).max(1);
Ok((frames as u64 * 44_100 / s.sample_rate.max(1) as u64) as usize)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bit_pack_roundtrip_10bit() {
let bits = bits_per_code(1024);
assert_eq!(bits, 10);
let codes: Vec<u32> = (0..5000u32).map(|i| (i * 37 + 11) % 1024).collect();
let mut bw = BitWriter::new();
for &c in &codes {
bw.put(c, bits);
}
let packed = bw.finish();
assert_eq!(packed.len(), (codes.len() * bits as usize).div_ceil(8));
let mut br = BitReader::new(&packed);
for &c in &codes {
assert_eq!(br.get(bits), c);
}
}
#[test]
fn bits_per_code_sizes() {
assert_eq!(bits_per_code(1024), 10);
assert_eq!(bits_per_code(512), 9);
assert_eq!(bits_per_code(2048), 11);
assert_eq!(bits_per_code(1), 1);
}
}