use crate::toc::Bandwidth;
use crate::Error;
pub fn interp_phase_samples(bandwidth: Bandwidth) -> Result<usize, Error> {
Ok(match bandwidth {
Bandwidth::Nb => 64,
Bandwidth::Mb => 96,
Bandwidth::Wb => 128,
_ => return Err(Error::MalformedPacket),
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct StereoWeightsQ13 {
pub w0_q13: i32,
pub w1_q13: i32,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct StereoUnmixState {
mid_hist: [f32; 2],
side_hist: f32,
prev_weights: StereoWeightsQ13,
}
impl Default for StereoUnmixState {
fn default() -> Self {
Self::new()
}
}
impl StereoUnmixState {
pub fn new() -> Self {
StereoUnmixState {
mid_hist: [0.0; 2],
side_hist: 0.0,
prev_weights: StereoWeightsQ13::default(),
}
}
pub fn reset(&mut self) {
*self = Self::new();
}
pub fn prev_weights(&self) -> StereoWeightsQ13 {
self.prev_weights
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct StereoFrame {
pub left: Vec<f32>,
pub right: Vec<f32>,
}
pub fn stereo_ms_to_lr(
bandwidth: Bandwidth,
mid: &[f32],
side: Option<&[f32]>,
weights: StereoWeightsQ13,
state: &mut StereoUnmixState,
) -> Result<StereoFrame, Error> {
let n2 = mid.len();
if n2 == 0 {
return Err(Error::MalformedPacket);
}
if let Some(s) = side {
if s.len() != n2 {
return Err(Error::MalformedPacket);
}
}
let n1 = interp_phase_samples(bandwidth)?;
let prev = state.prev_weights;
let w0_q13 = weights.w0_q13 as f32;
let w1_q13 = weights.w1_q13 as f32;
let prev_w0_q13 = prev.w0_q13 as f32;
let prev_w1_q13 = prev.w1_q13 as f32;
let n1_f = n1 as f32;
let w0_base = prev_w0_q13 / 8192.0;
let w1_base = prev_w1_q13 / 8192.0;
let w0_step = (w0_q13 - prev_w0_q13) / (8192.0 * n1_f);
let w1_step = (w1_q13 - prev_w1_q13) / (8192.0 * n1_f);
let mut left = vec![0.0f32; n2];
let mut right = vec![0.0f32; n2];
let mid_m2 = state.mid_hist[0]; let mid_m1 = state.mid_hist[1]; let side_m1 = if side.is_some() {
state.side_hist
} else {
0.0
};
for i in 0..n2 {
let ramp = (i.min(n1)) as f32;
let w0 = w0_base + ramp * w0_step;
let w1 = w1_base + ramp * w1_step;
let m_i = mid[i];
let m_i1 = if i >= 1 { mid[i - 1] } else { mid_m1 };
let m_i2 = match i {
0 => mid_m2,
1 => mid_m1,
_ => mid[i - 2],
};
let s_i1 = match side {
Some(s) if i >= 1 => s[i - 1],
Some(_) => side_m1,
None => 0.0,
};
let p0 = (m_i2 + 2.0 * m_i1 + m_i) / 4.0;
let l = (1.0 + w1) * m_i1 + s_i1 + w0 * p0;
let r = (1.0 - w1) * m_i1 - s_i1 - w0 * p0;
left[i] = l.clamp(-1.0, 1.0);
right[i] = r.clamp(-1.0, 1.0);
}
state.mid_hist = if n2 >= 2 {
[mid[n2 - 2], mid[n2 - 1]]
} else {
[mid_m1, mid[n2 - 1]]
};
state.side_hist = match side {
Some(s) => s[n2 - 1],
None => 0.0,
};
state.prev_weights = weights;
Ok(StereoFrame { left, right })
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct StereoDownmixState {
prev_mid: f32,
prev_weights: StereoWeightsQ13,
}
impl Default for StereoDownmixState {
fn default() -> Self {
Self::new()
}
}
impl StereoDownmixState {
pub fn new() -> Self {
StereoDownmixState {
prev_mid: 0.0,
prev_weights: StereoWeightsQ13::default(),
}
}
pub fn reset(&mut self) {
*self = Self::new();
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct MidSideFrame {
pub mid: Vec<f32>,
pub side: Vec<f32>,
}
pub fn stereo_lr_to_ms(
bandwidth: Bandwidth,
left: &[f32],
right: &[f32],
weights: StereoWeightsQ13,
next_lr: Option<(f32, f32)>,
state: &mut StereoDownmixState,
) -> Result<MidSideFrame, Error> {
let n2 = left.len();
if n2 == 0 || right.len() != n2 {
return Err(Error::MalformedPacket);
}
let n1 = interp_phase_samples(bandwidth)?;
let prev = state.prev_weights;
let n1_f = n1 as f32;
let w0_base = prev.w0_q13 as f32 / 8192.0;
let w1_base = prev.w1_q13 as f32 / 8192.0;
let w0_step = (weights.w0_q13 - prev.w0_q13) as f32 / (8192.0 * n1_f);
let w1_step = (weights.w1_q13 - prev.w1_q13) as f32 / (8192.0 * n1_f);
let mid: Vec<f32> = left
.iter()
.zip(right)
.map(|(&l, &r)| (l + r) / 2.0)
.collect();
let mid_next = match next_lr {
Some((l, r)) => (l + r) / 2.0,
None => mid[n2 - 1],
};
let mut side = vec![0.0f32; n2];
for k in 0..n2 {
let ramp = ((k + 1).min(n1)) as f32;
let w0 = w0_base + ramp * w0_step;
let w1 = w1_base + ramp * w1_step;
let m_km1 = if k >= 1 { mid[k - 1] } else { state.prev_mid };
let m_kp1 = if k + 1 < n2 { mid[k + 1] } else { mid_next };
let p0 = (m_km1 + 2.0 * mid[k] + m_kp1) / 4.0;
side[k] = (left[k] - right[k]) / 2.0 - w1 * mid[k] - w0 * p0;
}
state.prev_mid = mid[n2 - 1];
state.prev_weights = weights;
Ok(MidSideFrame { mid, side })
}
pub fn estimate_stereo_weights(
mid: &[f32],
side_raw: &[f32],
prev_mid: f32,
mid_next: f32,
) -> Result<StereoWeightsQ13, Error> {
let n2 = mid.len();
if n2 == 0 || side_raw.len() != n2 {
return Err(Error::MalformedPacket);
}
let mut sum_pp = 0f64;
let mut sum_pm = 0f64;
let mut sum_mm = 0f64;
let mut sum_ps = 0f64;
let mut sum_ms = 0f64;
for k in 0..n2 {
let m_km1 = if k >= 1 { mid[k - 1] } else { prev_mid } as f64;
let m_kp1 = if k + 1 < n2 { mid[k + 1] } else { mid_next } as f64;
let m = mid[k] as f64;
let p0 = (m_km1 + 2.0 * m + m_kp1) / 4.0;
let s = side_raw[k] as f64;
sum_pp += p0 * p0;
sum_pm += p0 * m;
sum_mm += m * m;
sum_ps += p0 * s;
sum_ms += m * s;
}
let det = sum_pp * sum_mm - sum_pm * sum_pm;
if det.abs() < 1e-12 {
return Ok(StereoWeightsQ13::default());
}
let w0 = (sum_ps * sum_mm - sum_ms * sum_pm) / det;
let w1 = (sum_ms * sum_pp - sum_ps * sum_pm) / det;
let to_q13 = |w: f64| -> i32 {
(w * 8192.0)
.round()
.clamp(-(1 << 30) as f64, (1 << 30) as f64) as i32
};
Ok(StereoWeightsQ13 {
w0_q13: to_q13(w0),
w1_q13: to_q13(w1),
})
}
#[cfg(test)]
mod tests {
use super::*;
fn approx(a: f32, b: f32) {
assert!(
(a - b).abs() < 1e-5,
"expected {b}, got {a} (delta {})",
(a - b).abs()
);
}
#[test]
fn interp_phase_table() {
assert_eq!(interp_phase_samples(Bandwidth::Nb).unwrap(), 64);
assert_eq!(interp_phase_samples(Bandwidth::Mb).unwrap(), 96);
assert_eq!(interp_phase_samples(Bandwidth::Wb).unwrap(), 128);
assert!(interp_phase_samples(Bandwidth::Swb).is_err());
assert!(interp_phase_samples(Bandwidth::Fb).is_err());
}
#[test]
fn state_starts_and_resets_zero() {
let mut s = StereoUnmixState::new();
assert_eq!(s.mid_hist, [0.0, 0.0]);
assert_eq!(s.side_hist, 0.0);
assert_eq!(s.prev_weights, StereoWeightsQ13::default());
s.mid_hist = [0.3, -0.2];
s.side_hist = 0.1;
s.prev_weights = StereoWeightsQ13 {
w0_q13: 5,
w1_q13: 7,
};
s.reset();
assert_eq!(s, StereoUnmixState::new());
}
#[test]
fn rejects_empty_and_mismatched() {
let mut s = StereoUnmixState::new();
assert!(stereo_ms_to_lr(
Bandwidth::Wb,
&[],
None,
StereoWeightsQ13::default(),
&mut s
)
.is_err());
let mid = vec![0.0f32; 80];
let side = vec![0.0f32; 79];
assert!(stereo_ms_to_lr(
Bandwidth::Wb,
&mid,
Some(&side),
StereoWeightsQ13::default(),
&mut s
)
.is_err());
}
#[test]
fn zero_weights_no_side_is_delayed_mono() {
let mut s = StereoUnmixState::new();
let mid: Vec<f32> = vec![0.1, 0.2, 0.3, 0.4];
let out = stereo_ms_to_lr(
Bandwidth::Wb,
&mid,
None,
StereoWeightsQ13::default(),
&mut s,
)
.unwrap();
let expect = [0.0, 0.1, 0.2, 0.3];
for (i, &e) in expect.iter().enumerate() {
approx(out.left[i], e);
approx(out.right[i], e);
}
assert_eq!(out.left, out.right);
}
#[test]
fn known_midside_reconstruction_constant_weights() {
let w = StereoWeightsQ13 {
w0_q13: 4096, w1_q13: 8192, };
let mut s = StereoUnmixState::new();
s.prev_weights = w;
let mid = vec![0.4f32, -0.2, 0.1, 0.3, -0.1];
let side = vec![0.05f32, 0.0, -0.1, 0.2, 0.1];
let out = stereo_ms_to_lr(Bandwidth::Wb, &mid, Some(&side), w, &mut s).unwrap();
let w0 = 0.5f32;
let w1 = 1.0f32;
let mut mhist = [0.0f32, 0.0]; let mut shist = 0.0f32; for i in 0..mid.len() {
let m_i = mid[i];
let m_i1 = mhist[1];
let m_i2 = mhist[0];
let s_i1 = shist;
let p0 = (m_i2 + 2.0 * m_i1 + m_i) / 4.0;
let l = ((1.0 + w1) * m_i1 + s_i1 + w0 * p0).clamp(-1.0, 1.0);
let r = ((1.0 - w1) * m_i1 - s_i1 - w0 * p0).clamp(-1.0, 1.0);
approx(out.left[i], l);
approx(out.right[i], r);
mhist = [m_i1, m_i];
shist = side[i];
}
}
#[test]
fn phase1_ramp_endpoints() {
let n1 = 64usize;
let n2 = n1 + 4;
let w_cur = StereoWeightsQ13 {
w0_q13: 0,
w1_q13: 8192,
}; let mut s = StereoUnmixState::new();
s.prev_weights = StereoWeightsQ13 {
w0_q13: 0,
w1_q13: 0,
};
let m = 0.4f32;
let mid = vec![m; n2];
let out = stereo_ms_to_lr(Bandwidth::Nb, &mid, None, w_cur, &mut s).unwrap();
approx(out.left[1], (1.0 + 1.0 / 64.0) * m);
approx(out.left[n1], 2.0 * m);
approx(out.left[n1 + 1], 2.0 * m);
approx(out.right[1], (1.0 - 1.0 / 64.0) * m);
approx(out.right[n1 + 1], 0.0);
}
#[test]
fn history_carries_across_frames() {
let w = StereoWeightsQ13 {
w0_q13: 0,
w1_q13: 0,
}; let mut s = StereoUnmixState::new();
s.prev_weights = w;
let frame1 = vec![0.1f32, 0.2, 0.3, 0.4];
let _ = stereo_ms_to_lr(Bandwidth::Wb, &frame1, None, w, &mut s).unwrap();
assert_eq!(s.mid_hist, [0.3, 0.4]);
let frame2 = vec![0.5f32, 0.6, 0.7, 0.8];
let out2 = stereo_ms_to_lr(Bandwidth::Wb, &frame2, None, w, &mut s).unwrap();
approx(out2.left[0], 0.4);
approx(out2.left[1], 0.5);
}
#[test]
fn side_history_carries_across_frames() {
let w = StereoWeightsQ13 {
w0_q13: 0,
w1_q13: 0,
};
let mut s = StereoUnmixState::new();
s.prev_weights = w;
let mid1 = vec![0.0f32; 4];
let side1 = vec![0.1f32, 0.2, 0.3, 0.4];
let _ = stereo_ms_to_lr(Bandwidth::Wb, &mid1, Some(&side1), w, &mut s).unwrap();
assert_eq!(s.side_hist, 0.4);
let mid2 = vec![0.0f32; 4];
let side2 = vec![0.5f32, 0.6, 0.7, 0.8];
let out2 = stereo_ms_to_lr(Bandwidth::Wb, &mid2, Some(&side2), w, &mut s).unwrap();
approx(out2.left[0], 0.4);
approx(out2.right[0], -0.4);
approx(out2.left[1], 0.5);
approx(out2.right[1], -0.5);
}
struct Lcg(u64);
impl Lcg {
fn next_f32(&mut self) -> f32 {
self.0 = self
.0
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((self.0 >> 40) as f32 / (1u64 << 24) as f32 - 0.5) * 0.8
}
}
#[test]
fn lr_to_ms_roundtrips_through_ms_to_lr() {
for (bandwidth, n2) in [
(Bandwidth::Nb, 160usize),
(Bandwidth::Mb, 240),
(Bandwidth::Wb, 320),
] {
let mut rng = Lcg(0xD0_1985 ^ n2 as u64);
let frames = 3usize;
let left: Vec<f32> = (0..frames * n2).map(|_| rng.next_f32()).collect();
let right: Vec<f32> = (0..frames * n2).map(|_| rng.next_f32()).collect();
let frame_weights = [
StereoWeightsQ13 {
w0_q13: -2950,
w1_q13: 820,
},
StereoWeightsQ13 {
w0_q13: 5000,
w1_q13: -6500,
},
StereoWeightsQ13 {
w0_q13: 820,
w1_q13: 10050,
},
];
let mut enc = StereoDownmixState::new();
let mut dec = StereoUnmixState::new();
let mut out_left = Vec::new();
let mut out_right = Vec::new();
for f in 0..frames {
let l = &left[f * n2..(f + 1) * n2];
let r = &right[f * n2..(f + 1) * n2];
let next_lr = if f + 1 < frames {
Some((left[(f + 1) * n2], right[(f + 1) * n2]))
} else {
None
};
let ms =
stereo_lr_to_ms(bandwidth, l, r, frame_weights[f], next_lr, &mut enc).unwrap();
for k in 0..n2 {
approx(ms.mid[k], (l[k] + r[k]) / 2.0);
}
let rec = stereo_ms_to_lr(
bandwidth,
&ms.mid,
Some(&ms.side),
frame_weights[f],
&mut dec,
)
.unwrap();
out_left.extend_from_slice(&rec.left);
out_right.extend_from_slice(&rec.right);
}
assert!(out_left[0].abs() < 1e-5, "{bandwidth:?}");
assert!(out_right[0].abs() < 1e-5, "{bandwidth:?}");
for i in 1..frames * n2 {
assert!(
(out_left[i] - left[i - 1]).abs() < 1e-4,
"{bandwidth:?} left sample {i}: {} vs {}",
out_left[i],
left[i - 1]
);
assert!(
(out_right[i] - right[i - 1]).abs() < 1e-4,
"{bandwidth:?} right sample {i}: {} vs {}",
out_right[i],
right[i - 1]
);
}
}
}
#[test]
fn lr_to_ms_single_frame_roundtrip_ignores_lookahead() {
let n2 = 160usize;
let mut rng = Lcg(0x0AC4);
let left: Vec<f32> = (0..n2).map(|_| rng.next_f32()).collect();
let right: Vec<f32> = (0..n2).map(|_| rng.next_f32()).collect();
let w = StereoWeightsQ13 {
w0_q13: 6500,
w1_q13: -820,
};
for next in [None, Some((0.35f32, -0.2f32))] {
let mut enc = StereoDownmixState::new();
let mut dec = StereoUnmixState::new();
let ms = stereo_lr_to_ms(Bandwidth::Nb, &left, &right, w, next, &mut enc).unwrap();
let rec = stereo_ms_to_lr(Bandwidth::Nb, &ms.mid, Some(&ms.side), w, &mut dec).unwrap();
for i in 1..n2 {
approx(rec.left[i], left[i - 1]);
approx(rec.right[i], right[i - 1]);
}
}
}
#[test]
fn estimate_recovers_planted_weights() {
let mut rng = Lcg(0xE571_0001);
let n2 = 320usize;
let mid: Vec<f32> = (0..n2).map(|_| rng.next_f32()).collect();
let prev_mid = rng.next_f32();
let mid_next = rng.next_f32();
let (w0, w1) = (0.37f64, -0.61f64);
let mut side_raw = vec![0.0f32; n2];
for k in 0..n2 {
let m_km1 = if k >= 1 { mid[k - 1] } else { prev_mid } as f64;
let m_kp1 = if k + 1 < n2 { mid[k + 1] } else { mid_next } as f64;
let p0 = (m_km1 + 2.0 * mid[k] as f64 + m_kp1) / 4.0;
side_raw[k] = (w0 * p0 + w1 * mid[k] as f64) as f32;
}
let est = estimate_stereo_weights(&mid, &side_raw, prev_mid, mid_next).unwrap();
assert!(
(est.w0_q13 - (w0 * 8192.0).round() as i32).abs() <= 1,
"w0 {} vs planted {}",
est.w0_q13,
(w0 * 8192.0).round()
);
assert!(
(est.w1_q13 - (w1 * 8192.0).round() as i32).abs() <= 1,
"w1 {} vs planted {}",
est.w1_q13,
(w1 * 8192.0).round()
);
}
#[test]
fn estimate_zero_mid_and_bad_input() {
let side = vec![0.25f32; 160];
let mid = vec![0.0f32; 160];
let est = estimate_stereo_weights(&mid, &side, 0.0, 0.0).unwrap();
assert_eq!(est, StereoWeightsQ13::default());
assert!(estimate_stereo_weights(&[], &[], 0.0, 0.0).is_err());
assert!(estimate_stereo_weights(&mid, &side[..159], 0.0, 0.0).is_err());
}
#[test]
fn estimate_quantize_downmix_reduces_side_energy() {
use crate::silk_frame::StereoWeightSymbols;
let mut rng = Lcg(0xE571_C0DE);
let n2 = 320usize;
let left: Vec<f32> = (0..n2).map(|_| rng.next_f32()).collect();
let right: Vec<f32> = left
.iter()
.map(|&l| 0.4 * l + 0.12 * rng.next_f32())
.collect();
let mid: Vec<f32> = left
.iter()
.zip(&right)
.map(|(&l, &r)| (l + r) / 2.0)
.collect();
let side_raw: Vec<f32> = left
.iter()
.zip(&right)
.map(|(&l, &r)| (l - r) / 2.0)
.collect();
let target = estimate_stereo_weights(&mid, &side_raw, 0.0, mid[n2 - 1]).unwrap();
let quintuple = StereoWeightSymbols::quantize(crate::silk_frame::StereoPredictionWeights {
w0_q13: target.w0_q13,
w1_q13: target.w1_q13,
});
let coded = quintuple.weights();
let coded_w = StereoWeightsQ13 {
w0_q13: coded.w0_q13,
w1_q13: coded.w1_q13,
};
let mut enc = StereoDownmixState::new();
enc.prev_weights = coded_w;
let ms = stereo_lr_to_ms(Bandwidth::Wb, &left, &right, coded_w, None, &mut enc).unwrap();
let energy = |v: &[f32]| -> f64 { v.iter().map(|&x| (x as f64) * (x as f64)).sum() };
assert!(
energy(&ms.side) < 0.7 * energy(&side_raw),
"predicted side energy {} not below unpredicted {}",
energy(&ms.side),
energy(&side_raw)
);
let mut dec = StereoUnmixState::new();
dec.prev_weights = coded_w;
let rec =
stereo_ms_to_lr(Bandwidth::Wb, &ms.mid, Some(&ms.side), coded_w, &mut dec).unwrap();
for i in 1..n2 {
assert!((rec.left[i] - left[i - 1]).abs() < 1e-4, "left {i}");
assert!((rec.right[i] - right[i - 1]).abs() < 1e-4, "right {i}");
}
}
#[test]
fn lr_to_ms_rejects_bad_input_and_resets() {
let mut s = StereoDownmixState::new();
assert!(stereo_lr_to_ms(
Bandwidth::Wb,
&[],
&[],
StereoWeightsQ13::default(),
None,
&mut s
)
.is_err());
assert!(stereo_lr_to_ms(
Bandwidth::Wb,
&[0.0; 4],
&[0.0; 3],
StereoWeightsQ13::default(),
None,
&mut s
)
.is_err());
assert!(stereo_lr_to_ms(
Bandwidth::Swb,
&[0.0; 4],
&[0.0; 4],
StereoWeightsQ13::default(),
None,
&mut s
)
.is_err());
let _ = stereo_lr_to_ms(
Bandwidth::Wb,
&[0.1; 320],
&[0.2; 320],
StereoWeightsQ13 {
w0_q13: 820,
w1_q13: 820,
},
None,
&mut s,
)
.unwrap();
assert_ne!(s, StereoDownmixState::new());
s.reset();
assert_eq!(s, StereoDownmixState::new());
}
#[test]
fn output_is_clamped() {
let w = StereoWeightsQ13 {
w0_q13: 8192 * 4, w1_q13: 8192 * 4, };
let mut s = StereoUnmixState::new();
s.prev_weights = w;
let mid = vec![1.0f32; 8];
let side = vec![1.0f32; 8];
let out = stereo_ms_to_lr(Bandwidth::Wb, &mid, Some(&side), w, &mut s).unwrap();
for i in 0..8 {
assert!((-1.0..=1.0).contains(&out.left[i]), "left {}", out.left[i]);
assert!(
(-1.0..=1.0).contains(&out.right[i]),
"right {}",
out.right[i]
);
}
}
}