use crate::celt_band_layout::{celt_first_coded_band, celt_total_bins_per_channel, CeltFrameSize};
use crate::celt_deemphasis::DeemphasisFilter;
use crate::celt_denormalise::{denormalise_bands, DenormaliseError};
use crate::celt_imdct::{imdct_into, ImdctError};
use crate::celt_mdct_window::CELT_OVERLAP_48K;
use crate::celt_overlap_add::{OverlapAddError, WeightedOverlapAdd};
#[derive(Debug, Clone, PartialEq)]
pub enum CeltSynthError {
BandCountMismatch {
expected: usize,
got_shapes: usize,
got_energies: usize,
},
ChannelCountMismatch {
expected: usize,
got: usize,
},
Denormalise(DenormaliseError),
Imdct(ImdctError),
OverlapAdd(OverlapAddError),
}
impl core::fmt::Display for CeltSynthError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
CeltSynthError::BandCountMismatch {
expected,
got_shapes,
got_energies,
} => write!(
f,
"oxideav-opus: CELT synthesis expected {expected} coded bands, \
got {got_shapes} shapes / {got_energies} energies"
),
CeltSynthError::ChannelCountMismatch { expected, got } => write!(
f,
"oxideav-opus: CELT synthesis state carries {expected} channel(s), \
call supplied {got}"
),
CeltSynthError::Denormalise(e) => write!(f, "oxideav-opus: CELT §4.3.6 {e}"),
CeltSynthError::Imdct(e) => write!(f, "oxideav-opus: CELT §4.3.7 {e}"),
CeltSynthError::OverlapAdd(e) => write!(f, "oxideav-opus: CELT §4.3.7 {e}"),
}
}
}
impl std::error::Error for CeltSynthError {}
#[derive(Debug, Clone, PartialEq)]
struct CeltChannelState {
ola: WeightedOverlapAdd,
deemph: DeemphasisFilter,
}
#[derive(Debug, Clone, PartialEq)]
pub struct CeltSynthState {
n: usize,
coded_bins: usize,
coded_bands: usize,
first_coded_band: usize,
frame_size: CeltFrameSize,
is_hybrid: bool,
channels: Vec<CeltChannelState>,
}
impl CeltSynthState {
pub fn new(
frame_size: CeltFrameSize,
is_hybrid: bool,
channels: usize,
) -> Result<Self, CeltSynthError> {
if channels != 1 && channels != 2 {
return Err(CeltSynthError::ChannelCountMismatch {
expected: 1,
got: channels,
});
}
let n = (frame_size.to_frame_tenths_ms() as usize * 48) / 10;
let coded_bins = celt_total_bins_per_channel(frame_size, is_hybrid) as usize;
let overlap = CELT_OVERLAP_48K.min(n);
let first = celt_first_coded_band(is_hybrid);
let coded = crate::celt_band_layout::celt_end_coded_band() - first;
let mut chans = Vec::with_capacity(channels);
for _ in 0..channels {
let ola = WeightedOverlapAdd::new(n, overlap).map_err(CeltSynthError::OverlapAdd)?;
chans.push(CeltChannelState {
ola,
deemph: DeemphasisFilter::new(),
});
}
Ok(Self {
n,
coded_bins,
coded_bands: coded,
first_coded_band: first,
frame_size,
is_hybrid,
channels: chans,
})
}
#[must_use]
pub fn transform_half_len(&self) -> usize {
self.n
}
#[must_use]
pub fn coded_bands(&self) -> usize {
self.coded_bands
}
#[must_use]
pub fn coded_bins(&self) -> usize {
self.coded_bins
}
#[must_use]
pub fn first_coded_band(&self) -> usize {
self.first_coded_band
}
#[must_use]
pub fn channels(&self) -> usize {
self.channels.len()
}
pub fn reset(&mut self) {
for ch in self.channels.iter_mut() {
ch.ola.reset();
ch.deemph.reset();
}
}
pub fn synthesize_channel_into(
&mut self,
channel: usize,
shapes: &[&[f64]],
log2_energy: &[f64],
out: &mut [f64],
) -> Result<(), CeltSynthError> {
if channel >= self.channels.len() {
return Err(CeltSynthError::ChannelCountMismatch {
expected: self.channels.len(),
got: channel + 1,
});
}
if shapes.len() != self.coded_bands || log2_energy.len() != self.coded_bands {
return Err(CeltSynthError::BandCountMismatch {
expected: self.coded_bands,
got_shapes: shapes.len(),
got_energies: log2_energy.len(),
});
}
if out.len() != self.n {
return Err(CeltSynthError::ChannelCountMismatch {
expected: self.n,
got: out.len(),
});
}
let mut freq = vec![0.0_f64; self.n];
let written = denormalise_bands(
shapes,
log2_energy,
self.frame_size,
self.is_hybrid,
&mut freq,
)
.map_err(CeltSynthError::Denormalise)?;
debug_assert_eq!(written, self.coded_bins);
let mut block = vec![0.0_f64; 2 * self.n];
imdct_into(&freq, &mut block).map_err(CeltSynthError::Imdct)?;
let ch = &mut self.channels[channel];
ch.ola
.process_into(&block, out)
.map_err(CeltSynthError::OverlapAdd)?;
ch.deemph.process_in_place(out);
Ok(())
}
pub fn synthesize_frame_interleaved_i16(
&mut self,
per_channel: &[(&[&[f64]], &[f64])],
) -> Result<Vec<i16>, CeltSynthError> {
if per_channel.len() != self.channels.len() {
return Err(CeltSynthError::ChannelCountMismatch {
expected: self.channels.len(),
got: per_channel.len(),
});
}
let n = self.n;
let ch_count = self.channels.len();
let mut planar: Vec<Vec<f64>> = Vec::with_capacity(ch_count);
for (c, (shapes, energies)) in per_channel.iter().enumerate() {
let mut buf = vec![0.0_f64; n];
self.synthesize_channel_into(c, shapes, energies, &mut buf)?;
planar.push(buf);
}
let mut out = vec![0_i16; ch_count * n];
for (s, slot) in out.chunks_exact_mut(ch_count).enumerate() {
for (c, dst) in slot.iter_mut().enumerate() {
*dst = celt_sample_to_i16(planar[c][s]);
}
}
Ok(out)
}
}
#[inline]
#[must_use]
fn celt_sample_to_i16(v: f64) -> i16 {
let scaled = (v.clamp(-1.0, 1.0) * 32767.0).round();
scaled as i16
}
#[cfg(test)]
mod tests {
use super::*;
use crate::celt_band_layout::celt_band_bins_per_channel;
fn zero_frame(frame_size: CeltFrameSize, is_hybrid: bool) -> (Vec<Vec<f64>>, Vec<f64>) {
let first = celt_first_coded_band(is_hybrid);
let end = crate::celt_band_layout::celt_end_coded_band();
let mut shapes = Vec::new();
let mut energies = Vec::new();
for band in first..end {
let bins = celt_band_bins_per_channel(band, frame_size).unwrap() as usize;
shapes.push(vec![0.0_f64; bins]);
energies.push(0.0_f64);
}
(shapes, energies)
}
fn shape_refs(shapes: &[Vec<f64>]) -> Vec<&[f64]> {
shapes.iter().map(|s| s.as_slice()).collect()
}
#[test]
fn new_rejects_bad_channel_count() {
assert!(matches!(
CeltSynthState::new(CeltFrameSize::Ms20, false, 0),
Err(CeltSynthError::ChannelCountMismatch { .. })
));
assert!(matches!(
CeltSynthState::new(CeltFrameSize::Ms20, false, 3),
Err(CeltSynthError::ChannelCountMismatch { .. })
));
assert!(CeltSynthState::new(CeltFrameSize::Ms20, false, 1).is_ok());
assert!(CeltSynthState::new(CeltFrameSize::Ms20, false, 2).is_ok());
}
#[test]
fn transform_half_len_matches_table55_total() {
let st = CeltSynthState::new(CeltFrameSize::Ms20, false, 1).unwrap();
assert_eq!(st.transform_half_len(), 960);
assert_eq!(st.coded_bins(), 800);
assert!(st.coded_bins() < st.transform_half_len());
assert_eq!(st.coded_bands(), 21);
assert_eq!(st.first_coded_band(), 0);
let st2 = CeltSynthState::new(CeltFrameSize::Ms2_5, false, 1).unwrap();
assert_eq!(st2.transform_half_len(), 120);
assert_eq!(st2.coded_bins(), 100);
let sth = CeltSynthState::new(CeltFrameSize::Ms20, true, 1).unwrap();
assert_eq!(sth.first_coded_band(), 17);
assert_eq!(sth.coded_bands(), 4);
}
#[test]
fn silent_frame_decodes_to_silence() {
let mut st = CeltSynthState::new(CeltFrameSize::Ms20, false, 1).unwrap();
let (shapes, energies) = zero_frame(CeltFrameSize::Ms20, false);
let refs = shape_refs(&shapes);
let mut out = vec![1.0_f64; st.transform_half_len()];
st.synthesize_channel_into(0, &refs, &energies, &mut out)
.unwrap();
assert!(out.iter().all(|&s| s == 0.0), "silent frame must be silent");
}
#[test]
fn band_count_mismatch_is_rejected() {
let mut st = CeltSynthState::new(CeltFrameSize::Ms20, false, 1).unwrap();
let (shapes, energies) = zero_frame(CeltFrameSize::Ms20, false);
let mut refs = shape_refs(&shapes);
refs.pop(); let mut out = vec![0.0_f64; st.transform_half_len()];
let err = st
.synthesize_channel_into(0, &refs, &energies, &mut out)
.unwrap_err();
assert!(matches!(err, CeltSynthError::BandCountMismatch { .. }));
}
#[test]
fn out_len_mismatch_is_rejected() {
let mut st = CeltSynthState::new(CeltFrameSize::Ms20, false, 1).unwrap();
let (shapes, energies) = zero_frame(CeltFrameSize::Ms20, false);
let refs = shape_refs(&shapes);
let mut out = vec![0.0_f64; st.transform_half_len() - 1];
let err = st
.synthesize_channel_into(0, &refs, &energies, &mut out)
.unwrap_err();
assert!(matches!(err, CeltSynthError::ChannelCountMismatch { .. }));
}
#[test]
fn channel_out_of_range_is_rejected() {
let mut st = CeltSynthState::new(CeltFrameSize::Ms20, false, 1).unwrap();
let (shapes, energies) = zero_frame(CeltFrameSize::Ms20, false);
let refs = shape_refs(&shapes);
let mut out = vec![0.0_f64; st.transform_half_len()];
let err = st
.synthesize_channel_into(1, &refs, &energies, &mut out)
.unwrap_err();
assert!(matches!(err, CeltSynthError::ChannelCountMismatch { .. }));
}
#[test]
fn reset_clears_overlap_and_deemphasis() {
let mut st = CeltSynthState::new(CeltFrameSize::Ms10, false, 1).unwrap();
let (mut shapes, mut energies) = zero_frame(CeltFrameSize::Ms10, false);
shapes[5][0] = 1.0;
energies[5] = 4.0;
let refs = shape_refs(&shapes);
let mut out = vec![0.0_f64; st.transform_half_len()];
st.synthesize_channel_into(0, &refs, &energies, &mut out)
.unwrap();
assert!(
out.iter().any(|&s| s != 0.0),
"nonzero frame should be audible"
);
st.reset();
let (zshapes, zenergies) = zero_frame(CeltFrameSize::Ms10, false);
let zrefs = shape_refs(&zshapes);
let mut out2 = vec![0.0_f64; st.transform_half_len()];
st.synthesize_channel_into(0, &zrefs, &zenergies, &mut out2)
.unwrap();
assert!(
out2.iter().all(|&s| s == 0.0),
"after reset, a silent frame must be exactly silent"
);
}
#[test]
fn two_silent_frames_stay_silent_across_state() {
let mut st = CeltSynthState::new(CeltFrameSize::Ms20, false, 1).unwrap();
let (shapes, energies) = zero_frame(CeltFrameSize::Ms20, false);
let refs = shape_refs(&shapes);
for _ in 0..2 {
let mut out = vec![0.0_f64; st.transform_half_len()];
st.synthesize_channel_into(0, &refs, &energies, &mut out)
.unwrap();
assert!(out.iter().all(|&s| s == 0.0));
}
}
#[test]
fn stereo_channels_are_independent() {
let mut st = CeltSynthState::new(CeltFrameSize::Ms10, false, 2).unwrap();
assert_eq!(st.channels(), 2);
let (mut shapes, mut energies) = zero_frame(CeltFrameSize::Ms10, false);
shapes[3][0] = 1.0;
energies[3] = 2.0;
let refs = shape_refs(&shapes);
let mut out0 = vec![0.0_f64; st.transform_half_len()];
st.synthesize_channel_into(0, &refs, &energies, &mut out0)
.unwrap();
let (zshapes, zenergies) = zero_frame(CeltFrameSize::Ms10, false);
let zrefs = shape_refs(&zshapes);
let mut out1 = vec![0.0_f64; st.transform_half_len()];
st.synthesize_channel_into(1, &zrefs, &zenergies, &mut out1)
.unwrap();
assert!(out1.iter().all(|&s| s == 0.0));
}
#[test]
fn energy_increases_with_log2_energy() {
let mk = |l: f64| -> f64 {
let mut st = CeltSynthState::new(CeltFrameSize::Ms10, false, 1).unwrap();
let (mut shapes, mut energies) = zero_frame(CeltFrameSize::Ms10, false);
shapes[2][0] = 1.0;
energies[2] = l;
let refs = shape_refs(&shapes);
let mut out = vec![0.0_f64; st.transform_half_len()];
st.synthesize_channel_into(0, &refs, &energies, &mut out)
.unwrap();
out.iter().map(|&s| s * s).sum::<f64>()
};
let low = mk(0.0);
let high = mk(2.0);
assert!(
high > low,
"higher band energy must yield more output power"
);
assert!(low > 0.0, "a nonzero shape at L=0 must be audible");
}
#[test]
fn sample_to_i16_matches_decoder_convention() {
assert_eq!(celt_sample_to_i16(0.0), 0);
assert_eq!(celt_sample_to_i16(1.0), 32767);
assert_eq!(celt_sample_to_i16(-1.0), -32767);
assert_eq!(celt_sample_to_i16(2.5), 32767);
assert_eq!(celt_sample_to_i16(-3.0), -32767);
assert_eq!(celt_sample_to_i16(0.5 / 32767.0 + 0.5), 16384);
}
#[test]
fn interleaved_silent_frame_is_all_zero_i16() {
let mut st = CeltSynthState::new(CeltFrameSize::Ms20, false, 2).unwrap();
let (shapes, energies) = zero_frame(CeltFrameSize::Ms20, false);
let refs = shape_refs(&shapes);
let per_channel: Vec<(&[&[f64]], &[f64])> = vec![
(refs.as_slice(), energies.as_slice()),
(refs.as_slice(), energies.as_slice()),
];
let pcm = st.synthesize_frame_interleaved_i16(&per_channel).unwrap();
assert_eq!(pcm.len(), 2 * st.transform_half_len());
assert!(pcm.iter().all(|&s| s == 0));
}
#[test]
fn interleaved_wrong_channel_count_rejected() {
let mut st = CeltSynthState::new(CeltFrameSize::Ms20, false, 2).unwrap();
let (shapes, energies) = zero_frame(CeltFrameSize::Ms20, false);
let refs = shape_refs(&shapes);
let per_channel: Vec<(&[&[f64]], &[f64])> = vec![(refs.as_slice(), energies.as_slice())];
let err = st
.synthesize_frame_interleaved_i16(&per_channel)
.unwrap_err();
assert!(matches!(err, CeltSynthError::ChannelCountMismatch { .. }));
}
#[test]
fn interleaved_layout_places_channels_correctly() {
let mut st = CeltSynthState::new(CeltFrameSize::Ms10, false, 2).unwrap();
let (mut a_shapes, mut a_energies) = zero_frame(CeltFrameSize::Ms10, false);
a_shapes[4][0] = 1.0;
a_energies[4] = 6.0;
let a_refs = shape_refs(&a_shapes);
let (z_shapes, z_energies) = zero_frame(CeltFrameSize::Ms10, false);
let z_refs = shape_refs(&z_shapes);
let per_channel: Vec<(&[&[f64]], &[f64])> = vec![
(a_refs.as_slice(), a_energies.as_slice()),
(z_refs.as_slice(), z_energies.as_slice()),
];
let pcm = st.synthesize_frame_interleaved_i16(&per_channel).unwrap();
let n = st.transform_half_len();
assert_eq!(pcm.len(), 2 * n);
assert!(pcm.iter().skip(1).step_by(2).all(|&s| s == 0));
assert!(pcm.iter().step_by(2).any(|&s| s != 0));
}
}