1use anyhow::{Context, Result};
16use rlx_dac::DacCodec;
17use rlx_runtime::Device;
18use std::path::{Path, PathBuf};
19
20pub fn best_device() -> Device {
24 for d in [Device::Mlx, Device::Metal, Device::Gpu] {
25 if rlx_runtime::is_available(d) {
26 return d;
27 }
28 }
29 Device::Cpu
30}
31
32pub fn default_dir() -> PathBuf {
35 std::env::var("RLX_DAC_DIR")
36 .map(PathBuf::from)
37 .unwrap_or_else(|_| PathBuf::from(".cache/dac44"))
38}
39
40pub fn open(device: Device) -> Result<DacCodec> {
45 open_in(&default_dir(), device)
46}
47
48pub fn open_in(model_dir: &Path, device: Device) -> Result<DacCodec> {
50 rlx_dac::download::ensure_weights(model_dir)?;
51 DacCodec::open_on(model_dir, device)
52}
53
54pub fn weights_available(dir: &Path) -> bool {
56 dir.join("model.safetensors").is_file() && dir.join("config.json").is_file()
57}
58
59pub struct CorrectCodec {
64 dac: DacCodec,
65 num_quantizers: Option<usize>,
66}
67
68const MAGIC_V1: &[u8; 4] = b"TSRX";
70const MAGIC_V2: &[u8; 4] = b"TSR2";
75
76fn bits_per_code(codebook_size: usize) -> u32 {
78 codebook_size.next_power_of_two().trailing_zeros().max(1)
79}
80
81struct BitWriter {
83 buf: Vec<u8>,
84 acc: u64,
85 nbits: u32,
86}
87impl BitWriter {
88 fn new() -> Self {
89 Self {
90 buf: Vec::new(),
91 acc: 0,
92 nbits: 0,
93 }
94 }
95 fn put(&mut self, value: u32, bits: u32) {
96 self.acc |= (value as u64) << self.nbits;
97 self.nbits += bits;
98 while self.nbits >= 8 {
99 self.buf.push((self.acc & 0xff) as u8);
100 self.acc >>= 8;
101 self.nbits -= 8;
102 }
103 }
104 fn finish(mut self) -> Vec<u8> {
105 if self.nbits > 0 {
106 self.buf.push((self.acc & 0xff) as u8);
107 }
108 self.buf
109 }
110}
111
112struct BitReader<'a> {
113 bytes: &'a [u8],
114 pos: usize,
115 acc: u64,
116 nbits: u32,
117}
118impl<'a> BitReader<'a> {
119 fn new(bytes: &'a [u8]) -> Self {
120 Self {
121 bytes,
122 pos: 0,
123 acc: 0,
124 nbits: 0,
125 }
126 }
127 fn get(&mut self, bits: u32) -> u32 {
128 while self.nbits < bits {
129 let byte = self.bytes.get(self.pos).copied().unwrap_or(0);
130 self.pos += 1;
131 self.acc |= (byte as u64) << self.nbits;
132 self.nbits += 8;
133 }
134 let mask = if bits >= 32 {
135 u32::MAX
136 } else {
137 (1u32 << bits) - 1
138 };
139 let v = (self.acc as u32) & mask;
140 self.acc >>= bits;
141 self.nbits -= bits;
142 v
143 }
144}
145
146impl CorrectCodec {
147 pub fn open(device: Device, quality: Option<u8>) -> Result<Self> {
150 Self::open_in(&default_dir(), device, quality)
151 }
152
153 pub fn open_in(model_dir: &Path, device: Device, quality: Option<u8>) -> Result<Self> {
154 let dac = open_in(model_dir, device)?;
155 let num_quantizers = quality.map(|q| (q as usize).clamp(1, 9));
156 Ok(Self {
157 dac,
158 num_quantizers,
159 })
160 }
161
162 pub fn sample_rate(&self) -> u32 {
163 self.dac.sample_rate()
164 }
165
166 pub fn encode_file(&self, in_audio: &Path, out_tsac: &Path) -> Result<()> {
169 let codes = self.dac.encode_wav(in_audio, self.num_quantizers)?;
170 let orig = mono_len_44k(in_audio)?;
171 let bits = bits_per_code(self.dac.config().codebook_size);
172 let mut buf = Vec::with_capacity(
173 21 + (codes.num_frames() * codes.num_quantizers * bits as usize).div_ceil(8),
174 );
175 buf.extend_from_slice(MAGIC_V2);
176 buf.extend_from_slice(&self.dac.sample_rate().to_le_bytes());
177 buf.extend_from_slice(&(orig as u32).to_le_bytes());
178 buf.extend_from_slice(&(codes.num_frames() as u32).to_le_bytes());
179 buf.extend_from_slice(&(codes.num_quantizers as u32).to_le_bytes());
180 buf.push(bits as u8);
181 let mut bw = BitWriter::new();
182 for frame in &codes.frames {
183 for &c in frame {
184 bw.put(c, bits);
185 }
186 }
187 buf.extend_from_slice(&bw.finish());
188 if let Some(parent) = out_tsac.parent() {
189 if !parent.as_os_str().is_empty() {
190 std::fs::create_dir_all(parent).ok();
191 }
192 }
193 std::fs::write(out_tsac, buf).with_context(|| format!("write {}", out_tsac.display()))?;
194 Ok(())
195 }
196
197 pub fn decode_file(&self, in_tsac: &Path, out_wav: &Path) -> Result<()> {
199 let b = std::fs::read(in_tsac).with_context(|| format!("read {}", in_tsac.display()))?;
200 anyhow::ensure!(b.len() >= 20, "container too short");
201 let rd = |o: usize| u32::from_le_bytes([b[o], b[o + 1], b[o + 2], b[o + 3]]) as usize;
202 let orig = rd(8);
203 let n_frames = rd(12);
204 let n_cb = rd(16);
205 let frames = if &b[0..4] == MAGIC_V2 {
206 anyhow::ensure!(b.len() >= 21, "TSR2 truncated header");
207 let bits = b[20] as u32;
208 anyhow::ensure!((1..=16).contains(&bits), "TSR2 bad bits_per_code {bits}");
209 let need = (n_frames * n_cb * bits as usize).div_ceil(8);
210 anyhow::ensure!(b.len() >= 21 + need, "TSR2 truncated payload");
211 let mut br = BitReader::new(&b[21..]);
212 let mut frames = Vec::with_capacity(n_frames);
213 for _ in 0..n_frames {
214 let mut row = Vec::with_capacity(n_cb);
215 for _ in 0..n_cb {
216 row.push(br.get(bits));
217 }
218 frames.push(row);
219 }
220 frames
221 } else if &b[0..4] == MAGIC_V1 {
222 anyhow::ensure!(b.len() >= 20 + n_frames * n_cb * 2, "TSRX truncated");
223 let mut frames = Vec::with_capacity(n_frames);
224 let mut off = 20;
225 for _ in 0..n_frames {
226 let mut row = Vec::with_capacity(n_cb);
227 for _ in 0..n_cb {
228 row.push(u16::from_le_bytes([b[off], b[off + 1]]) as u32);
229 off += 2;
230 }
231 frames.push(row);
232 }
233 frames
234 } else {
235 anyhow::bail!("not a TSRX/TSR2 container");
236 };
237 let codes = rlx_dac::DacCodes {
238 frames,
239 num_quantizers: n_cb,
240 };
241 self.dac.decode_wav(&codes, out_wav, Some(orig))
242 }
243}
244
245fn mono_len_44k(path: &Path) -> Result<usize> {
246 let r = hound::WavReader::open(path).with_context(|| format!("open {}", path.display()))?;
247 let s = r.spec();
248 let frames = (r.len() as usize) / (s.channels as usize).max(1);
249 Ok((frames as u64 * 44_100 / s.sample_rate.max(1) as u64) as usize)
250}
251
252#[cfg(test)]
253mod tests {
254 use super::*;
255
256 #[test]
259 fn bit_pack_roundtrip_10bit() {
260 let bits = bits_per_code(1024);
261 assert_eq!(bits, 10);
262 let codes: Vec<u32> = (0..5000u32).map(|i| (i * 37 + 11) % 1024).collect();
263 let mut bw = BitWriter::new();
264 for &c in &codes {
265 bw.put(c, bits);
266 }
267 let packed = bw.finish();
268 assert_eq!(packed.len(), (codes.len() * bits as usize).div_ceil(8));
270 let mut br = BitReader::new(&packed);
271 for &c in &codes {
272 assert_eq!(br.get(bits), c);
273 }
274 }
275
276 #[test]
277 fn bits_per_code_sizes() {
278 assert_eq!(bits_per_code(1024), 10);
279 assert_eq!(bits_per_code(512), 9);
280 assert_eq!(bits_per_code(2048), 11);
281 assert_eq!(bits_per_code(1), 1);
282 }
283}