Skip to main content

opus_rs/
celt.rs

1use crate::bands::{
2    SPREAD_NONE, SPREAD_NORMAL, compute_band_energies, denormalise_bands, haar1, log2amp,
3    normalise_bands, quant_all_bands, spreading_decision,
4};
5use crate::modes::{CeltMode, SPREAD_ICDF, TAPSET_ICDF, TF_SELECT_TABLE, TRIM_ICDF};
6use crate::quant_bands::{
7    quant_coarse_energy_advanced, quant_energy_finalise, quant_fine_energy, unquant_coarse_energy,
8    unquant_energy_finalise, unquant_fine_energy,
9};
10use crate::range_coder::RangeCoder;
11use crate::rate::{BITRES, clt_compute_allocation};
12
13#[cfg(not(feature = "std"))]
14use crate::compat::Math;
15use crate::fixedvec::FixedVec;
16
17// --- Heap-free buffer capacity constants (all sized for the worst case:
18//     2 channels, 48 kHz, max frame). Runtime construction uses smaller logical
19//     lengths when fewer channels / bands are active; the FixedVec tracks that. ---
20/// Maximum channel count supported by the API.
21const CELT_MAX_CHANNELS: usize = 2;
22/// Number of energy bands (`nb_ebands`) for the 48 kHz / 120-overlap mode.
23const CELT_NB_EBANDS: usize = 21;
24/// `nb_ebands * max_channels` = worst-case per-band-per-channel element count.
25const CELT_NB_X_CH: usize = CELT_NB_EBANDS * CELT_MAX_CHANNELS;
26/// MDCT overlap for the default mode.
27const CELT_OVERLAP: usize = 120;
28/// CELT synthesis / decode memory window per channel.
29const CELT_CHANNEL_MEM: usize = 2048 + CELT_OVERLAP;
30/// `channels * channel_mem` worst case.
31const CELT_SYN_MEM: usize = CELT_MAX_CHANNELS * CELT_CHANNEL_MEM;
32/// Prefilter comb-filter memory: `channels * COMBFILTER_MAXPERIOD`.
33const CELT_PREFILTER_MEM: usize = CELT_MAX_CHANNELS * COMBFILTER_MAXPERIOD;
34/// Encoder input buffer stride: `(MAX_FRAME_SIZE + overlap) * channels`.
35const CELT_BUFSTRIDE: usize = (MAX_FRAME_SIZE + CELT_OVERLAP) * CELT_MAX_CHANNELS;
36/// `MAX_FRAME_SIZE * channels`.
37const CELT_FRAME_X_CH: usize = MAX_FRAME_SIZE * CELT_MAX_CHANNELS;
38/// Frequency scratch (+4 padding for NEON pre-rotation over-read).
39const CELT_W_FREQ: usize = CELT_FRAME_X_CH + 4;
40/// PVQ stride-access scratch: `frame_x_ch + STRIDE_ACCESS_PAD`.
41const CELT_W_X_ENC: usize = CELT_FRAME_X_CH + STRIDE_ACCESS_PAD;
42/// Prefilter pre-scratch: `channels * (COMBFILTER_MAXPERIOD + MAX_FRAME_SIZE)`.
43const CELT_PREFILTER_PRE: usize = CELT_MAX_CHANNELS * (COMBFILTER_MAXPERIOD + MAX_FRAME_SIZE);
44/// Prefilter pitch buffer: `(COMBFILTER_MAXPERIOD + MAX_FRAME_SIZE) >> 1`.
45const CELT_PREFILTER_PITCH: usize = (COMBFILTER_MAXPERIOD + MAX_FRAME_SIZE) >> 1;
46/// Decoder `w_x`: `DECODE_BUFFER_SIZE * channels + STRIDE_ACCESS_PAD`.
47const CELT_W_X_DEC: usize = DECODE_BUFFER_SIZE * CELT_MAX_CHANNELS + STRIDE_ACCESS_PAD;
48/// Decoder `w_freq`: `DECODE_BUFFER_SIZE * channels + 4`.
49const CELT_W_FREQ_DEC: usize = DECODE_BUFFER_SIZE * CELT_MAX_CHANNELS + 4;
50/// Decoder `decode_mem`: `channels * (DECODE_BUFFER_SIZE + overlap)`.
51const CELT_DECODE_MEM: usize = CELT_MAX_CHANNELS * (DECODE_BUFFER_SIZE + CELT_OVERLAP);
52
53/// CELT internal-to-API decimation factor (port of libopus `resampling_factor`).
54/// The CELT decoder always runs at 48 kHz internally; the output is decimated
55/// by this factor to reach the API sampling rate.
56fn resampling_factor(sampling_rate: i32) -> usize {
57    match sampling_rate {
58        48000 => 1,
59        24000 => 2,
60        16000 => 3,
61        12000 => 4,
62        8000 => 6,
63        _ => 1,
64    }
65}
66
67#[cfg(target_arch = "aarch64")]
68use core::arch::aarch64::*;
69
70#[cfg(target_arch = "aarch64")]
71#[inline(always)]
72#[allow(unsafe_op_in_unsafe_fn)]
73unsafe fn sum_abs_neon(x: &[f32], n: usize) -> f32 {
74    let mut sum_vec = vdupq_n_f32(0.0);
75    let mut i = 0;
76
77    while i + 16 <= n {
78        let x0 = vld1q_f32(x.as_ptr().add(i));
79        let x1 = vld1q_f32(x.as_ptr().add(i + 4));
80        let x2 = vld1q_f32(x.as_ptr().add(i + 8));
81        let x3 = vld1q_f32(x.as_ptr().add(i + 12));
82
83        sum_vec = vfmaq_f32(sum_vec, vabsq_f32(x0), vdupq_n_f32(1.0));
84        sum_vec = vfmaq_f32(sum_vec, vabsq_f32(x1), vdupq_n_f32(1.0));
85        sum_vec = vfmaq_f32(sum_vec, vabsq_f32(x2), vdupq_n_f32(1.0));
86        sum_vec = vfmaq_f32(sum_vec, vabsq_f32(x3), vdupq_n_f32(1.0));
87
88        i += 16;
89    }
90
91    while i + 8 <= n {
92        let x0 = vld1q_f32(x.as_ptr().add(i));
93        let x1 = vld1q_f32(x.as_ptr().add(i + 4));
94        sum_vec = vfmaq_f32(sum_vec, vabsq_f32(x0), vdupq_n_f32(1.0));
95        sum_vec = vfmaq_f32(sum_vec, vabsq_f32(x1), vdupq_n_f32(1.0));
96        i += 8;
97    }
98
99    while i + 4 <= n {
100        let x0 = vld1q_f32(x.as_ptr().add(i));
101        sum_vec = vfmaq_f32(sum_vec, vabsq_f32(x0), vdupq_n_f32(1.0));
102        i += 4;
103    }
104
105    let mut sum = vaddvq_f32(sum_vec);
106
107    for j in i..n {
108        sum += x[j].abs();
109    }
110
111    sum
112}
113
114#[inline(always)]
115fn sum_abs(x: &[f32]) -> f32 {
116    #[cfg(target_arch = "x86_64")]
117    unsafe {
118        if crate::compat::x86_has_avx() {
119            return sum_abs_avx(x, x.len());
120        }
121    }
122    #[cfg(target_arch = "aarch64")]
123    unsafe {
124        sum_abs_neon(x, x.len())
125    }
126    #[cfg(not(target_arch = "aarch64"))]
127    {
128        x.iter().map(|&v| v.abs()).sum()
129    }
130}
131
132const MAX_FRAME_SIZE: usize = 2880;
133
134/// Largest CELT-internal synthesis buffer. The CELT decoder writes `n + overlap`
135/// MDCT samples per channel where `n ≤ frame_size ≤ 960` (20 ms @ 48 kHz); 2048
136/// is the libopus reference value (`DECODE_BUFFER_SIZE`) and gives ample margin.
137/// (Reduced from a previous 3072 to shrink the heap-free footprint.)
138const DECODE_BUFFER_SIZE: usize = 2048;
139
140const INV_TABLE: [u8; 128] = [
141    255, 255, 156, 110, 86, 70, 59, 51, 45, 40, 37, 33, 31, 28, 26, 25, 23, 22, 21, 20, 19, 18, 17,
142    16, 16, 15, 15, 14, 13, 13, 12, 12, 12, 12, 11, 11, 11, 10, 10, 10, 9, 9, 9, 9, 9, 9, 8, 8, 8,
143    8, 8, 7, 7, 7, 7, 7, 7, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5,
144    5, 5, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 3, 3, 3, 3,
145    3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 2,
146];
147
148const MAX_TRANSIENT_LEN: usize = 3000;
149
150#[derive(Debug, Clone, Copy)]
151pub struct AnalysisInfo {
152    pub valid: bool,
153    pub tonality: f32,
154    pub tonality_slope: f32,
155    pub noisiness: f32,
156    pub activity: f32,
157    pub music_prob: f32,
158    pub music_prob_min: f32,
159    pub music_prob_max: f32,
160    pub bandwidth: i32,
161    pub activity_probability: f32,
162    pub max_pitch_ratio: f32,
163    pub leak_boost: [u8; 19], // LEAK_BANDS = 19
164}
165
166impl Default for AnalysisInfo {
167    fn default() -> Self {
168        Self {
169            valid: false,
170            tonality: 0.0,
171            tonality_slope: 0.0,
172            noisiness: 0.0,
173            activity: 0.0,
174            music_prob: 0.0,
175            music_prob_min: 0.0,
176            music_prob_max: 0.0,
177            bandwidth: 0,
178            activity_probability: 0.0,
179            max_pitch_ratio: 1.0,
180            leak_boost: [0; 19],
181        }
182    }
183}
184
185#[allow(clippy::too_many_arguments)]
186fn transient_analysis(
187    input: &[f32],
188    len: usize,
189    channels: usize,
190    tf_estimate: &mut f32,
191    tf_chan: &mut usize,
192    allow_weak_transients: bool,
193    weak_transient: &mut bool,
194    _tone_freq: f32,
195    toneishness: f32,
196    tmp: &mut [f32],
197    tmp2: &mut [f32],
198) -> bool {
199    let mut mask_metric = 0.0f32;
200    let mut forward_decay = 0.0625f32;
201
202    *weak_transient = false;
203    if allow_weak_transients {
204        forward_decay = 0.03125f32;
205    }
206
207    let len2 = len / 2;
208    debug_assert!(len <= MAX_TRANSIENT_LEN);
209
210    for c in 0..channels {
211        let mut mem0 = 0.0f32;
212        let mut mem1 = 0.0f32;
213
214        for i in 0..len {
215            let x = input[c * len + i];
216            let y = mem0 + x;
217            let mem00 = mem0;
218            mem0 = mem0 - x + 0.5 * mem1;
219            mem1 = x - mem00;
220            tmp[i] = y;
221        }
222
223        tmp[..12].fill(0.0);
224
225        let mut mean = 0.0f32;
226        mem0 = 0.0f32;
227        for i in 0..len2 {
228            let x2 = (tmp[2 * i] * tmp[2 * i] + tmp[2 * i + 1] * tmp[2 * i + 1]) / 16.0;
229            mean += x2 / 4096.0;
230            mem0 = x2 + (1.0 - forward_decay) * mem0;
231            tmp2[i] = forward_decay * mem0;
232        }
233
234        mem0 = 0.0f32;
235        let mut max_e = 0.0f32;
236        for i in (0..len2).rev() {
237            mem0 = tmp2[i] + 0.875 * mem0;
238            tmp2[i] = 0.125 * mem0;
239            if tmp2[i] > max_e {
240                max_e = tmp2[i];
241            }
242        }
243
244        mean = (mean * max_e * 0.5 * (len2 as f32)).sqrt();
245        let norm = (len2 as f32) / (1e-10 + mean);
246
247        let mut unmask = 0.0f32;
248        for i in (12..(len2 - 5)).step_by(4) {
249            let id = (64.0 * norm * (tmp2[i] + 1e-10)).floor() as i32;
250            let id = id.clamp(0, 127) as usize;
251            unmask += INV_TABLE[id] as f32;
252        }
253
254        unmask = 64.0 * unmask * 4.0 / (6.0 * (len2 as f32 - 17.0));
255        if unmask > mask_metric {
256            *tf_chan = c;
257            mask_metric = unmask;
258        }
259    }
260
261    let mut is_transient = mask_metric > 200.0;
262
263    if toneishness > 0.98 && _tone_freq < 0.026 {
264        is_transient = false;
265        mask_metric = 0.0;
266    }
267
268    *tf_estimate = (mask_metric - 150.0).clamp(0.0, 1.0);
269
270    is_transient
271}
272
273fn l1_metric(tmp: &[f32], n: usize, lm: i32, bias: f32) -> f32 {
274    #[cfg(target_arch = "x86_64")]
275    unsafe {
276        if n >= 16 && crate::compat::x86_has_avx() {
277            return l1_metric_avx(tmp, n, lm, bias);
278        }
279    }
280    #[cfg(target_arch = "aarch64")]
281    {
282        if n >= 16 {
283            return unsafe { l1_metric_neon(tmp, n, lm, bias) };
284        }
285    }
286
287    let mut l1 = 0.0f32;
288    for &tv in tmp[..n].iter() {
289        l1 += tv.abs();
290    }
291    l1 + (lm as f32) * bias * l1
292}
293
294#[cfg(target_arch = "x86_64")]
295#[target_feature(enable = "avx")]
296unsafe fn sum_abs_avx(x: &[f32], n: usize) -> f32 {
297    use core::arch::x86_64::*;
298
299    let mut sum0 = _mm256_setzero_ps();
300    let mut sum1 = _mm256_setzero_ps();
301    let mut i = 0usize;
302    let sign_mask = _mm256_set1_ps(-0.0);
303
304    while i + 16 <= n {
305        let v0 = _mm256_loadu_ps(x.as_ptr().add(i));
306        let v1 = _mm256_loadu_ps(x.as_ptr().add(i + 8));
307        sum0 = _mm256_add_ps(sum0, _mm256_andnot_ps(sign_mask, v0));
308        sum1 = _mm256_add_ps(sum1, _mm256_andnot_ps(sign_mask, v1));
309        i += 16;
310    }
311
312    while i + 8 <= n {
313        let v = _mm256_loadu_ps(x.as_ptr().add(i));
314        sum0 = _mm256_add_ps(sum0, _mm256_andnot_ps(sign_mask, v));
315        i += 8;
316    }
317
318    let sum = _mm256_add_ps(sum0, sum1);
319    let hi = _mm256_extractf128_ps(sum, 1);
320    let lo = _mm256_castps256_ps128(sum);
321    let s4 = _mm_add_ps(lo, hi);
322    let t1 = _mm_movehl_ps(s4, s4);
323    let s2 = _mm_add_ps(s4, t1);
324    let t2 = _mm_shuffle_ps(s2, s2, 0x55);
325    let mut out = _mm_cvtss_f32(_mm_add_ss(s2, t2));
326
327    for j in i..n {
328        out += x[j].abs();
329    }
330
331    out
332}
333
334#[cfg(target_arch = "x86_64")]
335#[target_feature(enable = "avx")]
336unsafe fn l1_metric_avx(tmp: &[f32], n: usize, lm: i32, bias: f32) -> f32 {
337    let l1 = sum_abs_avx(tmp, n);
338    l1 + (lm as f32) * bias * l1
339}
340
341#[cfg(target_arch = "aarch64")]
342#[target_feature(enable = "neon")]
343unsafe fn l1_metric_neon(tmp: &[f32], n: usize, lm: i32, bias: f32) -> f32 {
344    unsafe {
345        let mut sum4 = vdupq_n_f32(0.0);
346        let mut i = 0;
347
348        while i + 15 < n {
349            let v0 = vld1q_f32(tmp.as_ptr().add(i));
350            let v1 = vld1q_f32(tmp.as_ptr().add(i + 4));
351            let v2 = vld1q_f32(tmp.as_ptr().add(i + 8));
352            let v3 = vld1q_f32(tmp.as_ptr().add(i + 12));
353
354            sum4 = vaddq_f32(sum4, vabsq_f32(v0));
355            sum4 = vaddq_f32(sum4, vabsq_f32(v1));
356            sum4 = vaddq_f32(sum4, vabsq_f32(v2));
357            sum4 = vaddq_f32(sum4, vabsq_f32(v3));
358
359            i += 16;
360        }
361
362        while i + 3 < n {
363            let v = vld1q_f32(tmp.as_ptr().add(i));
364            sum4 = vaddq_f32(sum4, vabsq_f32(v));
365            i += 4;
366        }
367
368        let sum2 = vpaddq_f32(sum4, sum4);
369        let sum1 = vpaddq_f32(sum2, sum2);
370        let mut l1 = vgetq_lane_f32(sum1, 0);
371
372        while i < n {
373            l1 += tmp[i].abs();
374            i += 1;
375        }
376
377        l1 + (lm as f32) * bias * l1
378    }
379}
380
381const MAX_NB_EBANDS: usize = 21;
382
383const MAX_TF_TMP: usize = 176;
384
385#[allow(clippy::too_many_arguments)]
386fn tf_analysis(
387    mode: &CeltMode,
388    len: usize,
389    is_transient: bool,
390    tf_res: &mut [i32],
391    lambda: i32,
392    x: &[f32],
393    n0: usize,
394    lm: i32,
395    tf_estimate: f32,
396    tf_chan: usize,
397) -> i32 {
398    debug_assert!(len <= MAX_NB_EBANDS);
399    let mut metric = [0i32; MAX_NB_EBANDS];
400    let mut tmp = [0.0f32; MAX_TF_TMP];
401    let mut tmp_1 = [0.0f32; MAX_TF_TMP];
402
403    let bias = 0.04 * (-0.25f32).max(0.5 - tf_estimate);
404
405    for (i, metric_i) in metric[..len].iter_mut().enumerate() {
406        let n = ((mode.e_bands[i + 1] - mode.e_bands[i]) as usize) << lm;
407        let narrow = (mode.e_bands[i + 1] - mode.e_bands[i]) == 1;
408        let offset = tf_chan * n0 + ((mode.e_bands[i] as usize) << lm);
409        tmp[..n].copy_from_slice(&x[offset..offset + n]);
410
411        let mut l1 = l1_metric(&tmp[..n], n, if is_transient { lm } else { 0 }, bias);
412        let mut best_l1 = l1;
413        let mut best_level = 0;
414
415        if is_transient && !narrow {
416            tmp_1[..n].copy_from_slice(&tmp[..n]);
417            haar1(&mut tmp_1[..n], n >> lm, 1 << lm);
418            l1 = l1_metric(&tmp_1[..n], n, lm + 1, bias);
419            if l1 < best_l1 {
420                best_l1 = l1;
421                best_level = -1;
422            }
423        }
424
425        for k in 0..(lm + if is_transient || narrow { 0 } else { 1 }) {
426            let b = if is_transient { lm - k - 1 } else { k + 1 };
427
428            haar1(&mut tmp[..n], n >> k, 1 << k);
429            l1 = l1_metric(&tmp[..n], n, b, bias);
430
431            if l1 < best_l1 {
432                best_l1 = l1;
433                best_level = k + 1;
434            }
435        }
436
437        if is_transient {
438            *metric_i = 2 * best_level;
439        } else {
440            *metric_i = -2 * best_level;
441        }
442
443        if narrow && (*metric_i == 0 || *metric_i == -2 * lm) {
444            *metric_i -= 1;
445        }
446    }
447
448    let mut tf_select = 0;
449    let importance = [1.0f32; MAX_NB_EBANDS];
450    let mut selcost = [0.0f32; 2];
451
452    for sel in 0..2 {
453        let mut cost0 = importance[0]
454            * ((metric[0]
455                - 2 * TF_SELECT_TABLE[lm as usize][4 * (is_transient as usize) + 2 * sel] as i32)
456                as f32)
457                .abs();
458        let mut cost1 = importance[0]
459            * ((metric[0]
460                - 2 * TF_SELECT_TABLE[lm as usize][4 * (is_transient as usize) + 2 * sel + 1]
461                    as i32) as f32)
462                .abs()
463            + (if is_transient { 0.0 } else { lambda as f32 });
464
465        for i in 1..len {
466            let curr0 = cost0.min(cost1 + lambda as f32);
467            let curr1 = (cost0 + lambda as f32).min(cost1);
468            cost0 = curr0
469                + importance[i]
470                    * ((metric[i]
471                        - 2 * TF_SELECT_TABLE[lm as usize][4 * (is_transient as usize) + 2 * sel]
472                            as i32) as f32)
473                        .abs();
474            cost1 = curr1
475                + importance[i]
476                    * ((metric[i]
477                        - 2 * TF_SELECT_TABLE[lm as usize]
478                            [4 * (is_transient as usize) + 2 * sel + 1]
479                            as i32) as f32)
480                        .abs();
481        }
482        selcost[sel] = cost0.min(cost1);
483    }
484
485    if selcost[1] < selcost[0] {
486        tf_select = 1;
487    }
488
489    let mut cost0 = importance[0]
490        * ((metric[0]
491            - 2 * TF_SELECT_TABLE[lm as usize][4 * (is_transient as usize) + 2 * tf_select] as i32)
492            as f32)
493            .abs();
494    let mut cost1 = importance[0]
495        * ((metric[0]
496            - 2 * TF_SELECT_TABLE[lm as usize][4 * (is_transient as usize) + 2 * tf_select + 1]
497                as i32) as f32)
498            .abs()
499        + (if is_transient { 0.0 } else { lambda as f32 });
500
501    tf_res[0] = if cost0 < cost1 { 0 } else { 1 };
502
503    for i in 1..len {
504        let curr0 = cost0.min(cost1 + lambda as f32);
505        let curr1 = (cost0 + lambda as f32).min(cost1);
506        cost0 = curr0
507            + importance[i]
508                * ((metric[i]
509                    - 2 * TF_SELECT_TABLE[lm as usize][4 * (is_transient as usize) + 2 * tf_select]
510                        as i32) as f32)
511                    .abs();
512        cost1 = curr1
513            + importance[i]
514                * ((metric[i]
515                    - 2 * TF_SELECT_TABLE[lm as usize]
516                        [4 * (is_transient as usize) + 2 * tf_select + 1]
517                        as i32) as f32)
518                    .abs();
519        tf_res[i] = if cost0 < cost1 { 0 } else { 1 };
520    }
521
522    tf_select as i32
523}
524
525fn tf_encode(
526    start: usize,
527    end: usize,
528    is_transient: bool,
529    tf_res: &mut [i32],
530    lm: i32,
531    mut tf_select: i32,
532    rc: &mut RangeCoder,
533) -> i32 {
534    let mut curr = 0;
535    let mut tf_changed = 0;
536    let mut logp = if is_transient { 2 } else { 4 };
537    let mut budget = rc.storage as i32 * 8;
538    let mut tell = rc.tell();
539
540    let tf_select_rsv = if lm > 0 && tell + logp < budget { 1 } else { 0 };
541    budget -= tf_select_rsv;
542
543    for tf_res_i in tf_res[start..end].iter_mut() {
544        if tell + logp <= budget {
545            rc.encode_bit_logp(*tf_res_i ^ curr != 0, logp as u32);
546            tell = rc.tell();
547            curr = *tf_res_i;
548            tf_changed |= curr;
549        } else {
550            *tf_res_i = curr;
551        }
552        logp = if is_transient { 4 } else { 5 };
553    }
554
555    if tf_select_rsv != 0
556        && TF_SELECT_TABLE[lm as usize][4 * (is_transient as usize) + (tf_changed as usize)]
557            != TF_SELECT_TABLE[lm as usize][4 * (is_transient as usize) + 2 + (tf_changed as usize)]
558    {
559        rc.encode_bit_logp(tf_select != 0, 1);
560    } else {
561        tf_select = 0;
562    }
563
564    for tf_res_i in tf_res[start..end].iter_mut() {
565        *tf_res_i = TF_SELECT_TABLE[lm as usize]
566            [4 * (is_transient as usize) + 2 * (tf_select as usize) + (*tf_res_i as usize)]
567            as i32;
568    }
569
570    tf_changed
571}
572
573fn tf_decode(
574    start: usize,
575    end: usize,
576    is_transient: bool,
577    tf_res: &mut [i32],
578    lm: i32,
579    rc: &mut RangeCoder,
580) {
581    let mut curr = 0;
582    let mut tf_changed = 0;
583    let mut logp = if is_transient { 2 } else { 4 };
584    let budget = rc.storage as i32 * 8;
585    let mut tell = rc.tell();
586
587    let tf_select_rsv = if lm > 0 && tell + logp < budget { 1 } else { 0 };
588    let budget = budget - tf_select_rsv;
589
590    for tf_res_i in tf_res[start..end].iter_mut() {
591        if tell + logp <= budget {
592            curr ^= if rc.decode_bit_logp(logp as u32) {
593                1
594            } else {
595                0
596            };
597            tell = rc.tell();
598            tf_changed |= curr;
599        }
600        *tf_res_i = curr;
601        logp = if is_transient { 4 } else { 5 };
602    }
603
604    let mut tf_select = 0;
605    let _budget = budget + tf_select_rsv;
606    if tf_select_rsv > 0
607        && TF_SELECT_TABLE[lm as usize][4 * (is_transient as usize) + (tf_changed as usize)]
608            != TF_SELECT_TABLE[lm as usize][4 * (is_transient as usize) + 2 + (tf_changed as usize)]
609    {
610        tf_select = if rc.decode_bit_logp(1) { 1 } else { 0 };
611    }
612
613    for tf_res_i in tf_res[start..end].iter_mut() {
614        *tf_res_i = TF_SELECT_TABLE[lm as usize]
615            [4 * (is_transient as usize) + 2 * (tf_select as usize) + (*tf_res_i as usize)]
616            as i32;
617    }
618}
619
620fn stereo_analysis(m: &CeltMode, x: &[f32], lm: i32, n0: usize) -> bool {
621    let mut sum_lr = 1e-9f32;
622    let mut sum_ms = 1e-9f32;
623
624    for i in 0..13 {
625        let start = (m.e_bands[i] as usize) << lm;
626        let end = (m.e_bands[i + 1] as usize) << lm;
627        for j in start..end {
628            let l = x[j];
629            let r = x[n0 + j];
630            let m_val = l + r;
631            let s_val = l - r;
632            sum_lr += l.abs() + r.abs();
633            sum_ms += m_val.abs() + s_val.abs();
634        }
635    }
636
637    sum_ms *= core::f32::consts::FRAC_1_SQRT_2;
638    let mut thetas = 13;
639    if lm <= 1 {
640        thetas -= 8;
641    }
642
643    let left = (((m.e_bands[13] as usize) << (lm + 1)) + thetas) as f32 * sum_ms;
644    let right = ((m.e_bands[13] as usize) << (lm + 1)) as f32 * sum_lr;
645
646    left > right
647}
648
649const COMBFILTER_MINPERIOD: usize = 15;
650const COMBFILTER_MAXPERIOD: usize = 1024;
651
652const PREFILTER_GAINS: [[f32; 3]; 3] = [
653    [0.306_640_6, 0.217_041, 0.129_638_7],
654    [0.463_867_2, 0.268_066_4, 0.0],
655    [0.799_804_7, 0.100_097_7, 0.0],
656];
657
658#[allow(clippy::too_many_arguments)]
659fn comb_filter_const(
660    y: &mut [f32],
661    x: &[f32],
662    y_idx: usize,
663    x_idx: usize,
664    t: usize,
665    n: usize,
666    g10: f32,
667    g11: f32,
668    g12: f32,
669) {
670    #[cfg(target_arch = "aarch64")]
671    {
672        comb_filter_const_neon(y, x, y_idx, x_idx, t, n, g10, g11, g12);
673    }
674    #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
675    unsafe {
676        if crate::compat::x86_has_avx() {
677            comb_filter_const_avx(y, x, y_idx, x_idx, t, n, g10, g11, g12);
678            return;
679        }
680    }
681    #[cfg(all(target_arch = "x86_64", target_feature = "sse"))]
682    unsafe {
683        comb_filter_const_sse(y, x, y_idx, x_idx, t, n, g10, g11, g12);
684        #[allow(clippy::needless_return)]
685        return;
686    }
687    #[cfg(not(any(
688        target_arch = "aarch64",
689        all(target_arch = "x86_64", target_feature = "sse")
690    )))]
691    {
692        comb_filter_const_scalar(y, x, y_idx, x_idx, t, n, g10, g11, g12);
693    }
694}
695
696#[inline]
697#[allow(dead_code)]
698fn comb_filter_const_scalar(
699    y: &mut [f32],
700    x: &[f32],
701    y_idx: usize,
702    x_idx: usize,
703    t: usize,
704    n: usize,
705    g10: f32,
706    g11: f32,
707    g12: f32,
708) {
709    let mut x1;
710    let mut x2;
711    let mut x3;
712    let mut x4;
713    let mut x0;
714
715    x4 = x[x_idx - t - 2];
716    x3 = x[x_idx - t - 1];
717    x2 = x[x_idx - t];
718    x1 = x[x_idx - t + 1];
719
720    for i in 0..n {
721        x0 = x[x_idx + i - t + 2];
722        y[y_idx + i] = x[x_idx + i] + g10 * x2 + g11 * (x1 + x3) + g12 * (x0 + x4);
723        x4 = x3;
724        x3 = x2;
725        x2 = x1;
726        x1 = x0;
727    }
728}
729
730#[cfg(target_arch = "aarch64")]
731fn comb_filter_const_neon(
732    y: &mut [f32],
733    x: &[f32],
734    y_idx: usize,
735    x_idx: usize,
736    t: usize,
737    n: usize,
738    g10: f32,
739    g11: f32,
740    g12: f32,
741) {
742    unsafe { comb_filter_const_neon_impl(y, x, y_idx, x_idx, t, n, g10, g11, g12) }
743}
744
745#[cfg(target_arch = "aarch64")]
746#[inline(always)]
747#[allow(unsafe_op_in_unsafe_fn)]
748unsafe fn comb_filter_const_neon_impl(
749    y: &mut [f32],
750    x: &[f32],
751    y_idx: usize,
752    x_idx: usize,
753    t: usize,
754    n: usize,
755    g10: f32,
756    g11: f32,
757    g12: f32,
758) {
759    use core::arch::aarch64::*;
760
761    let g10v = vdupq_n_f32(g10);
762    let g11v = vdupq_n_f32(g11);
763    let g12v = vdupq_n_f32(g12);
764
765    let xbase = x.as_ptr().add(x_idx);
766    let ybase = y.as_mut_ptr().add(y_idx);
767
768    let mut x0v = vld1q_f32(xbase.sub(t + 2));
769
770    let mut i = 0;
771    while i + 4 <= n {
772        let x4v = vld1q_f32(xbase.add(i).sub(t - 2));
773
774        let x2v = vextq_f32(x0v, x4v, 2);
775
776        let x1v = vextq_f32(x0v, x4v, 1);
777
778        let x3v = vextq_f32(x0v, x4v, 3);
779
780        let xi = vld1q_f32(xbase.add(i));
781
782        let mut yi = xi;
783        yi = vfmaq_f32(yi, g10v, x2v);
784        yi = vfmaq_f32(yi, g11v, vaddq_f32(x1v, x3v));
785        yi = vfmaq_f32(yi, g12v, vaddq_f32(x4v, x0v));
786        vst1q_f32(ybase.add(i), yi);
787
788        x0v = x4v;
789        i += 4;
790    }
791
792    let x0v_arr: [f32; 4] = core::mem::transmute(x0v);
793    let mut sx4 = x0v_arr[0];
794    let mut sx3 = x0v_arr[1];
795    let mut sx2 = x0v_arr[2];
796    let mut sx1 = x0v_arr[3];
797
798    while i < n {
799        let sx0 = x[x_idx + i - t + 2];
800        y[y_idx + i] = x[x_idx + i] + g10 * sx2 + g11 * (sx1 + sx3) + g12 * (sx0 + sx4);
801        sx4 = sx3;
802        sx3 = sx2;
803        sx2 = sx1;
804        sx1 = sx0;
805        i += 1;
806    }
807}
808
809#[cfg(all(target_arch = "x86_64", target_feature = "sse"))]
810#[inline(always)]
811#[allow(unsafe_op_in_unsafe_fn)]
812unsafe fn comb_filter_const_sse(
813    y: &mut [f32],
814    x: &[f32],
815    y_idx: usize,
816    x_idx: usize,
817    t: usize,
818    n: usize,
819    g10: f32,
820    g11: f32,
821    g12: f32,
822) {
823    use core::arch::x86_64::*;
824
825    let g10v = _mm_set1_ps(g10);
826    let g11v = _mm_set1_ps(g11);
827    let g12v = _mm_set1_ps(g12);
828
829    let xbase = x.as_ptr().add(x_idx);
830    let ybase = y.as_mut_ptr().add(y_idx);
831    let mut x0v = _mm_loadu_ps(xbase.sub(t + 2));
832
833    let mut i = 0;
834    while i + 4 <= n {
835        let x4v = _mm_loadu_ps(xbase.add(i).sub(t - 2));
836
837        let x2v = _mm_shuffle_ps(x0v, x4v, 0x4e);
838
839        let x1v = _mm_shuffle_ps(x0v, x2v, 0x99);
840
841        let x3v = _mm_shuffle_ps(x2v, x4v, 0x99);
842
843        let xi = _mm_loadu_ps(xbase.add(i));
844
845        let mut yi = xi;
846        yi = _mm_add_ps(yi, _mm_mul_ps(g10v, x2v));
847        let yi2 = _mm_add_ps(
848            _mm_mul_ps(g11v, _mm_add_ps(x3v, x1v)),
849            _mm_mul_ps(g12v, _mm_add_ps(x4v, x0v)),
850        );
851        yi = _mm_add_ps(yi, yi2);
852        _mm_storeu_ps(ybase.add(i), yi);
853
854        x0v = x4v;
855        i += 4;
856    }
857
858    let x0v_arr: [f32; 4] = core::mem::transmute(x0v);
859    let mut sx4 = x0v_arr[0];
860    let mut sx3 = x0v_arr[1];
861    let mut sx2 = x0v_arr[2];
862    let mut sx1 = x0v_arr[3];
863
864    while i < n {
865        let sx0 = x[x_idx + i - t + 2];
866        y[y_idx + i] = x[x_idx + i] + g10 * sx2 + g11 * (sx1 + sx3) + g12 * (sx0 + sx4);
867        sx4 = sx3;
868        sx3 = sx2;
869        sx2 = sx1;
870        sx1 = sx0;
871        i += 1;
872    }
873}
874
875#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
876#[target_feature(enable = "avx,fma")]
877#[allow(unsafe_op_in_unsafe_fn)]
878unsafe fn comb_filter_const_avx(
879    y: &mut [f32],
880    x: &[f32],
881    y_idx: usize,
882    x_idx: usize,
883    t: usize,
884    n: usize,
885    g10: f32,
886    g11: f32,
887    g12: f32,
888) {
889    use core::arch::x86_64::*;
890
891    let g10v = _mm256_set1_ps(g10);
892    let g11v = _mm256_set1_ps(g11);
893    let g12v = _mm256_set1_ps(g12);
894
895    let xbase = x.as_ptr().add(x_idx);
896    let ybase = y.as_mut_ptr().add(y_idx);
897
898    let mut i = 0;
899
900    while i + 16 <= n {
901        let xi_a = _mm256_loadu_ps(xbase.add(i));
902        let x0_a = _mm256_loadu_ps(xbase.add(i).sub(t + 2));
903        let x4_a = _mm256_loadu_ps(xbase.add(i).sub(t - 2));
904
905        let x2_a = _mm256_loadu_ps(xbase.add(i).sub(t));
906        let x1x3_a = _mm256_add_ps(
907            _mm256_loadu_ps(xbase.add(i).sub(t + 1)),
908            _mm256_loadu_ps(xbase.add(i).sub(t - 1)),
909        );
910        let x0x4_a = _mm256_add_ps(x0_a, x4_a);
911
912        let mut yi_a = xi_a;
913        yi_a = _mm256_fmadd_ps(g10v, x2_a, yi_a);
914        yi_a = _mm256_fmadd_ps(g11v, x1x3_a, yi_a);
915        yi_a = _mm256_fmadd_ps(g12v, x0x4_a, yi_a);
916        _mm256_storeu_ps(ybase.add(i), yi_a);
917
918        let j = i + 8;
919        let xi_b = _mm256_loadu_ps(xbase.add(j));
920        let x0_b = _mm256_loadu_ps(xbase.add(j).sub(t + 2));
921        let x4_b = _mm256_loadu_ps(xbase.add(j).sub(t - 2));
922        let x2_b = _mm256_loadu_ps(xbase.add(j).sub(t));
923        let x1x3_b = _mm256_add_ps(
924            _mm256_loadu_ps(xbase.add(j).sub(t + 1)),
925            _mm256_loadu_ps(xbase.add(j).sub(t - 1)),
926        );
927        let x0x4_b = _mm256_add_ps(x0_b, x4_b);
928
929        let mut yi_b = xi_b;
930        yi_b = _mm256_fmadd_ps(g10v, x2_b, yi_b);
931        yi_b = _mm256_fmadd_ps(g11v, x1x3_b, yi_b);
932        yi_b = _mm256_fmadd_ps(g12v, x0x4_b, yi_b);
933        _mm256_storeu_ps(ybase.add(j), yi_b);
934
935        i += 16;
936    }
937
938    while i + 8 <= n {
939        let xi = _mm256_loadu_ps(xbase.add(i));
940        let x0 = _mm256_loadu_ps(xbase.add(i).sub(t + 2));
941        let x4 = _mm256_loadu_ps(xbase.add(i).sub(t - 2));
942        let x2 = _mm256_loadu_ps(xbase.add(i).sub(t));
943        let x1x3 = _mm256_add_ps(
944            _mm256_loadu_ps(xbase.add(i).sub(t + 1)),
945            _mm256_loadu_ps(xbase.add(i).sub(t - 1)),
946        );
947        let x0x4 = _mm256_add_ps(x0, x4);
948
949        let mut yi = xi;
950        yi = _mm256_fmadd_ps(g10v, x2, yi);
951        yi = _mm256_fmadd_ps(g11v, x1x3, yi);
952        yi = _mm256_fmadd_ps(g12v, x0x4, yi);
953        _mm256_storeu_ps(ybase.add(i), yi);
954
955        i += 8;
956    }
957
958    if i + 4 <= n {
959        comb_filter_const_sse_fma(y, x, y_idx + i, x_idx + i, t, n - i, g10, g11, g12);
960        return;
961    }
962
963    let mut sx4 = x[x_idx + i - t - 2];
964    let mut sx3 = x[x_idx + i - t - 1];
965    let mut sx2 = x[x_idx + i - t];
966    let mut sx1 = x[x_idx + i - t + 1];
967    while i < n {
968        let sx0 = x[x_idx + i - t + 2];
969        y[y_idx + i] = x[x_idx + i] + g10 * sx2 + g11 * (sx1 + sx3) + g12 * (sx0 + sx4);
970        sx4 = sx3;
971        sx3 = sx2;
972        sx2 = sx1;
973        sx1 = sx0;
974        i += 1;
975    }
976}
977
978#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
979#[target_feature(enable = "avx,fma")]
980#[allow(unsafe_op_in_unsafe_fn)]
981unsafe fn comb_filter_const_sse_fma(
982    y: &mut [f32],
983    x: &[f32],
984    y_idx: usize,
985    x_idx: usize,
986    t: usize,
987    n: usize,
988    g10: f32,
989    g11: f32,
990    g12: f32,
991) {
992    use core::arch::x86_64::*;
993
994    let g10v = _mm_set1_ps(g10);
995    let g11v = _mm_set1_ps(g11);
996    let g12v = _mm_set1_ps(g12);
997
998    let xbase = x.as_ptr().add(x_idx);
999    let ybase = y.as_mut_ptr().add(y_idx);
1000    let mut x0v = _mm_loadu_ps(xbase.sub(t + 2));
1001
1002    let mut i = 0;
1003    while i + 4 <= n {
1004        let x4v = _mm_loadu_ps(xbase.add(i).sub(t - 2));
1005        let x2v = _mm_shuffle_ps(x0v, x4v, 0x4e);
1006        let x1v = _mm_shuffle_ps(x0v, x2v, 0x99);
1007        let x3v = _mm_shuffle_ps(x2v, x4v, 0x99);
1008        let xi = _mm_loadu_ps(xbase.add(i));
1009
1010        let mut yi = xi;
1011        yi = _mm_fmadd_ps(g10v, x2v, yi);
1012        yi = _mm_fmadd_ps(g11v, _mm_add_ps(x1v, x3v), yi);
1013        yi = _mm_fmadd_ps(g12v, _mm_add_ps(x0v, x4v), yi);
1014        _mm_storeu_ps(ybase.add(i), yi);
1015
1016        x0v = x4v;
1017        i += 4;
1018    }
1019
1020    let x0v_arr: [f32; 4] = core::mem::transmute(x0v);
1021    let mut sx4 = x0v_arr[0];
1022    let mut sx3 = x0v_arr[1];
1023    let mut sx2 = x0v_arr[2];
1024    let mut sx1 = x0v_arr[3];
1025    while i < n {
1026        let sx0 = x[x_idx + i - t + 2];
1027        y[y_idx + i] = x[x_idx + i] + g10 * sx2 + g11 * (sx1 + sx3) + g12 * (sx0 + sx4);
1028        sx4 = sx3;
1029        sx3 = sx2;
1030        sx2 = sx1;
1031        sx1 = sx0;
1032        i += 1;
1033    }
1034}
1035
1036#[allow(clippy::too_many_arguments)]
1037fn comb_filter(
1038    y: &mut [f32],
1039    x: &[f32],
1040    y_idx: usize,
1041    x_idx: usize,
1042    t0: usize,
1043    t1: usize,
1044    n: usize,
1045    g0: f32,
1046    g1: f32,
1047    tapset0: i32,
1048    tapset1: i32,
1049    window: &[f32],
1050    overlap: usize,
1051) {
1052    if g0 == 0.0 && g1 == 0.0 {
1053        if x_idx != y_idx || !core::ptr::eq(x.as_ptr(), y.as_ptr()) {
1054            y[y_idx..y_idx + n].copy_from_slice(&x[x_idx..x_idx + n]);
1055        }
1056        return;
1057    }
1058
1059    let t0 = t0.clamp(
1060        COMBFILTER_MINPERIOD,
1061        x_idx.saturating_sub(2).max(COMBFILTER_MINPERIOD),
1062    );
1063    let t1 = t1.clamp(
1064        COMBFILTER_MINPERIOD,
1065        x_idx.saturating_sub(2).max(COMBFILTER_MINPERIOD),
1066    );
1067
1068    let g00 = g0 * PREFILTER_GAINS[tapset0 as usize][0];
1069    let g01 = g0 * PREFILTER_GAINS[tapset0 as usize][1];
1070    let g02 = g0 * PREFILTER_GAINS[tapset0 as usize][2];
1071
1072    let g10 = g1 * PREFILTER_GAINS[tapset1 as usize][0];
1073    let g11 = g1 * PREFILTER_GAINS[tapset1 as usize][1];
1074    let g12 = g1 * PREFILTER_GAINS[tapset1 as usize][2];
1075
1076    let mut x1 = x[x_idx - t1 + 1];
1077    let mut x2 = x[x_idx - t1];
1078    let mut x3 = x[x_idx - t1 - 1];
1079    let mut x4 = x[x_idx - t1 - 2];
1080
1081    let mut inner_overlap = overlap;
1082    if g0 == g1 && t0 == t1 && tapset0 == tapset1 {
1083        inner_overlap = 0;
1084    }
1085
1086    let mut i = 0;
1087    while i < inner_overlap && i < n {
1088        let x0 = x[x_idx + i - t1 + 2];
1089        let f = window[i] * window[i];
1090        y[y_idx + i] = x[x_idx + i]
1091            + (1.0 - f)
1092                * (g00 * x[x_idx + i - t0]
1093                    + g01 * (x[x_idx + i - t0 + 1] + x[x_idx + i - t0 - 1])
1094                    + g02 * (x[x_idx + i - t0 + 2] + x[x_idx + i - t0 - 2]))
1095            + f * (g10 * x2 + g11 * (x1 + x3) + g12 * (x0 + x4));
1096
1097        x4 = x3;
1098        x3 = x2;
1099        x2 = x1;
1100        x1 = x0;
1101        i += 1;
1102    }
1103
1104    if i < n {
1105        if g1 == 0.0 {
1106            y[y_idx + i..y_idx + n].copy_from_slice(&x[x_idx + i..x_idx + n]);
1107        } else {
1108            comb_filter_const(y, x, y_idx + i, x_idx + i, t1, n - i, g10, g11, g12);
1109        }
1110    }
1111}
1112
1113/// In-place comb filter: buf[y_idx..y_idx+n] is both input and output.
1114/// Reference samples at buf[y_idx + i - T + offset] may already be filtered
1115/// if T < i, matching C libopus's in-place comb_filter(out, out, ...) behavior.
1116fn comb_filter_inplace(
1117    buf: &mut [f32],
1118    y_idx: usize,
1119    t0: usize,
1120    t1: usize,
1121    n: usize,
1122    g0: f32,
1123    g1: f32,
1124    tapset0: i32,
1125    tapset1: i32,
1126    window: &[f32],
1127    overlap: usize,
1128) {
1129    if g0 == 0.0 && g1 == 0.0 {
1130        // nothing to do; buf[y_idx..] already holds the input
1131        return;
1132    }
1133
1134    let t0 = t0.clamp(COMBFILTER_MINPERIOD, y_idx - 2);
1135    let t1 = t1.clamp(COMBFILTER_MINPERIOD, y_idx - 2);
1136
1137    let g00 = g0 * PREFILTER_GAINS[tapset0 as usize][0];
1138    let g01 = g0 * PREFILTER_GAINS[tapset0 as usize][1];
1139    let g02 = g0 * PREFILTER_GAINS[tapset0 as usize][2];
1140
1141    let g10 = g1 * PREFILTER_GAINS[tapset1 as usize][0];
1142    let g11 = g1 * PREFILTER_GAINS[tapset1 as usize][1];
1143    let g12 = g1 * PREFILTER_GAINS[tapset1 as usize][2];
1144
1145    let mut inner_overlap = overlap;
1146    if g0 == g1 && t0 == t1 && tapset0 == tapset1 {
1147        inner_overlap = 0;
1148    }
1149
1150    let mut i = 0;
1151    while i < inner_overlap && i < n {
1152        let idx = y_idx + i;
1153        let f = window[i] * window[i];
1154        let s = buf[idx]; // original input (not yet overwritten at idx)
1155        let r0 = buf[idx - t0];
1156        let r0p1 = buf[idx - t0 + 1];
1157        let r0m1 = buf[idx - t0 - 1];
1158        let r0p2 = buf[idx - t0 + 2];
1159        let r0m2 = buf[idx - t0 - 2];
1160        let r1 = buf[idx - t1];
1161        let r1p1 = buf[idx - t1 + 1];
1162        let r1m1 = buf[idx - t1 - 1];
1163        let r1p2 = buf[idx - t1 + 2];
1164        let r1m2 = buf[idx - t1 - 2];
1165        buf[idx] = s
1166            + (1.0 - f) * (g00 * r0 + g01 * (r0p1 + r0m1) + g02 * (r0p2 + r0m2))
1167            + f * (g10 * r1 + g11 * (r1p1 + r1m1) + g12 * (r1p2 + r1m2));
1168        i += 1;
1169    }
1170
1171    // Constant region: only new filter (t1, g1)
1172    while i < n {
1173        let idx = y_idx + i;
1174        let s = buf[idx];
1175        let r1 = buf[idx - t1];
1176        let r1p1 = buf[idx - t1 + 1];
1177        let r1m1 = buf[idx - t1 - 1];
1178        let r1p2 = buf[idx - t1 + 2];
1179        let r1m2 = buf[idx - t1 - 2];
1180        buf[idx] = s + g10 * r1 + g11 * (r1p1 + r1m1) + g12 * (r1p2 + r1m2);
1181        i += 1;
1182    }
1183}
1184
1185fn run_prefilter(
1186    in_buf: &mut [f32],
1187    prefilter_mem: &mut [f32],
1188    prefilter_period: usize,
1189    prefilter_gain: f32,
1190    prefilter_tapset: i32,
1191    tapset_decision: i32,
1192    window: &[f32],
1193    channels: usize,
1194    frame_size: usize,
1195    overlap: usize,
1196
1197    pre: &mut [f32],
1198    pitch_buf: &mut [f32],
1199    before: &mut [f32],
1200    after: &mut [f32],
1201
1202    analysis: &AnalysisInfo,
1203    loss_rate: i32,
1204) -> (bool, f32, usize) {
1205    let max_period = COMBFILTER_MAXPERIOD;
1206    let min_period = COMBFILTER_MINPERIOD;
1207    let buf_stride = frame_size + overlap;
1208    let pre_size = max_period + frame_size;
1209
1210    for c in 0..channels {
1211        pre[c * pre_size..c * pre_size + max_period]
1212            .copy_from_slice(&prefilter_mem[c * max_period..(c + 1) * max_period]);
1213        pre[c * pre_size + max_period..c * pre_size + pre_size].copy_from_slice(
1214            &in_buf[c * buf_stride + overlap..c * buf_stride + overlap + frame_size],
1215        );
1216    }
1217
1218    let pitch_buf_len = (max_period + frame_size) >> 1;
1219    {
1220        let mut pre_slices: FixedVec<&[f32], 2> = FixedVec::new();
1221        for c in 0..channels {
1222            pre_slices.push(&pre[c * pre_size..c * pre_size + pre_size]);
1223        }
1224        crate::pitch::pitch_downsample(&pre_slices, pitch_buf, pitch_buf_len, channels, 2);
1225    }
1226
1227    let search_max = max_period - 3 * min_period;
1228    let pitch_result = crate::pitch::pitch_search(
1229        &pitch_buf[max_period >> 1..],
1230        pitch_buf,
1231        frame_size,
1232        search_max,
1233    );
1234    let mut pitch_index = (max_period - pitch_result).min(max_period - 2);
1235
1236    let gain1_raw = crate::pitch::remove_doubling(
1237        pitch_buf,
1238        max_period,
1239        min_period,
1240        frame_size,
1241        &mut pitch_index,
1242        prefilter_period,
1243        prefilter_gain,
1244    );
1245    let mut gain1 = gain1_raw * 0.7;
1246
1247    // Apply max_pitch_ratio from analysis if available
1248    if analysis.valid {
1249        gain1 *= analysis.max_pitch_ratio;
1250    }
1251
1252    // Apply loss_rate scaling: halve at 2%, quarter at 4%, zero at 8%
1253    if loss_rate >= 8 {
1254        gain1 = 0.0;
1255    } else if loss_rate > 0 {
1256        gain1 *= 1.0 - (loss_rate as f32) / 8.0;
1257    }
1258
1259    let mut pf_threshold = 0.2f32;
1260    if (pitch_index as i32 - prefilter_period as i32).unsigned_abs() as usize * 10 > pitch_index {
1261        pf_threshold += 0.2;
1262    }
1263    if prefilter_gain > 0.4 {
1264        pf_threshold -= 0.1;
1265    }
1266    if prefilter_gain > 0.55 {
1267        pf_threshold -= 0.1;
1268    }
1269    pf_threshold = pf_threshold.max(0.2);
1270
1271    let pf_on;
1272    if gain1 < pf_threshold {
1273        gain1 = 0.0;
1274        pf_on = false;
1275    } else {
1276        if (gain1 - prefilter_gain).abs() < 0.1 {
1277            gain1 = prefilter_gain;
1278        }
1279        let qg = ((gain1 * 32.0 / 3.0 + 0.5).floor() as i32 - 1).clamp(0, 7);
1280        gain1 = 0.09375 * (qg + 1) as f32;
1281        pf_on = true;
1282    }
1283
1284    let before = &mut before[..channels];
1285    for c in 0..channels {
1286        let start = c * buf_stride + overlap;
1287        before[c] = sum_abs(&in_buf[start..start + frame_size]);
1288    }
1289
1290    let offset = 0usize;
1291    let prev_period = prefilter_period.clamp(COMBFILTER_MINPERIOD, max_period - 2);
1292
1293    for c in 0..channels {
1294        if offset > 0 {
1295            let pre_c = &pre[c * pre_size..];
1296            comb_filter(
1297                in_buf,
1298                pre_c,
1299                c * buf_stride + overlap,
1300                max_period,
1301                prev_period,
1302                prev_period,
1303                offset,
1304                -prefilter_gain,
1305                -prefilter_gain,
1306                prefilter_tapset,
1307                prefilter_tapset,
1308                window,
1309                0,
1310            );
1311        }
1312
1313        {
1314            let pre_c = &pre[c * pre_size..];
1315            comb_filter(
1316                in_buf,
1317                pre_c,
1318                c * buf_stride + overlap + offset,
1319                max_period + offset,
1320                prev_period,
1321                pitch_index,
1322                frame_size - offset,
1323                -prefilter_gain,
1324                -gain1,
1325                prefilter_tapset,
1326                tapset_decision,
1327                window,
1328                overlap,
1329            );
1330        }
1331    }
1332
1333    let after = &mut after[..channels];
1334    for c in 0..channels {
1335        let start = c * buf_stride + overlap;
1336        after[c] = sum_abs(&in_buf[start..start + frame_size]);
1337    }
1338
1339    let cancel_pitch = (0..channels).any(|c| after[c] > before[c]);
1340
1341    if cancel_pitch {
1342        for c in 0..channels {
1343            in_buf[c * buf_stride + overlap..c * buf_stride + overlap + frame_size]
1344                .copy_from_slice(
1345                    &pre[c * pre_size + max_period..c * pre_size + max_period + frame_size],
1346                );
1347        }
1348
1349        for c in 0..channels {
1350            if frame_size >= max_period {
1351                prefilter_mem[c * max_period..(c + 1) * max_period].copy_from_slice(
1352                    &pre[c * pre_size + frame_size..c * pre_size + frame_size + max_period],
1353                );
1354            } else {
1355                let shift = max_period - frame_size;
1356                prefilter_mem.copy_within(
1357                    c * max_period + frame_size..(c + 1) * max_period,
1358                    c * max_period,
1359                );
1360                prefilter_mem[c * max_period + shift..(c + 1) * max_period].copy_from_slice(
1361                    &pre[c * pre_size + max_period..c * pre_size + max_period + frame_size],
1362                );
1363            }
1364        }
1365        return (false, 0.0, pitch_index);
1366    }
1367
1368    for c in 0..channels {
1369        if frame_size >= max_period {
1370            prefilter_mem[c * max_period..(c + 1) * max_period].copy_from_slice(
1371                &pre[c * pre_size + frame_size..c * pre_size + frame_size + max_period],
1372            );
1373        } else {
1374            let shift = max_period - frame_size;
1375            prefilter_mem.copy_within(
1376                c * max_period + frame_size..(c + 1) * max_period,
1377                c * max_period,
1378            );
1379            prefilter_mem[c * max_period + shift..(c + 1) * max_period].copy_from_slice(
1380                &pre[c * pre_size + max_period..c * pre_size + max_period + frame_size],
1381            );
1382        }
1383    }
1384
1385    (pf_on, gain1, pitch_index)
1386}
1387
1388/// Padding appended to `w_x` to absorb any SIMD over-shoot past the last band.
1389/// Sized to `MAX_PVQ_N` (352) — far larger than the widest SIMD access (16),
1390/// which is always loop-guarded (`while i + width <= n`). (Previously
1391/// `MAX_PVQ_N * 8`, which had no access pattern justifying it.)
1392const STRIDE_ACCESS_PAD: usize = crate::pvq::MAX_PVQ_N;
1393
1394pub struct CeltEncoder {
1395    mode: &'static CeltMode,
1396    channels: usize,
1397    pub complexity: i32,
1398    syn_mem: FixedVec<f32, CELT_SYN_MEM>,
1399    enc_decode_mem: FixedVec<f32, CELT_SYN_MEM>,
1400    old_band_e: FixedVec<f32, CELT_NB_X_CH>,
1401    preemph_mem: FixedVec<f32, CELT_MAX_CHANNELS>,
1402    tonal_average: i32,
1403    hf_average: i32,
1404    tapset_decision: i32,
1405    spread_decision: i32,
1406    intensity: i32,
1407    last_coded_bands: i32,
1408    prefilter_mem: FixedVec<f32, CELT_PREFILTER_MEM>,
1409    prefilter_period: usize,
1410    prefilter_gain: f32,
1411    prefilter_tapset: i32,
1412    old_band_e2: FixedVec<f32, CELT_NB_X_CH>,
1413    old_band_e3: FixedVec<f32, CELT_NB_X_CH>,
1414    last_band_log_e: FixedVec<f32, CELT_NB_X_CH>,
1415    delayed_intra: f32,
1416
1417    w_in_buf: FixedVec<f32, CELT_BUFSTRIDE>,
1418    w_freq: FixedVec<f32, CELT_W_FREQ>,
1419    w_band_e: FixedVec<f32, CELT_NB_X_CH>,
1420    w_x: FixedVec<f32, CELT_W_X_ENC>,
1421    w_band_log_e: FixedVec<f32, CELT_NB_X_CH>,
1422    w_error: FixedVec<f32, CELT_NB_X_CH>,
1423    w_tf_res: FixedVec<i32, CELT_NB_EBANDS>,
1424    w_cap: FixedVec<i32, CELT_NB_EBANDS>,
1425    w_offsets: FixedVec<i32, CELT_NB_EBANDS>,
1426    w_pulses: FixedVec<i32, CELT_NB_EBANDS>,
1427    w_ebits: FixedVec<i32, CELT_NB_X_CH>,
1428    w_fine_priority: FixedVec<i32, CELT_NB_X_CH>,
1429    w_collapse_masks: FixedVec<u32, CELT_NB_X_CH>,
1430    w_band_amp_synth: FixedVec<f32, CELT_NB_X_CH>,
1431    consec_transient: i32,
1432
1433    w_prefilter_pre: FixedVec<f32, CELT_PREFILTER_PRE>,
1434    w_prefilter_pitch_buf: FixedVec<f32, CELT_PREFILTER_PITCH>,
1435    w_prefilter_before: FixedVec<f32, CELT_MAX_CHANNELS>,
1436    w_prefilter_after: FixedVec<f32, CELT_MAX_CHANNELS>,
1437
1438    w_transient_tmp: FixedVec<f32, MAX_TRANSIENT_LEN>,
1439    w_transient_tmp2: FixedVec<f32, { MAX_TRANSIENT_LEN / 2 }>,
1440
1441    analysis: AnalysisInfo,
1442    loss_rate: i32,
1443}
1444
1445const INTEN_THRESHOLDS: [i32; 21] = [
1446    1, 2, 3, 4, 5, 6, 7, 8, 16, 24, 36, 44, 50, 56, 62, 67, 72, 79, 88, 106, 134,
1447];
1448const INTEN_HYSTERESIS: [i32; 21] = [
1449    1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 3, 3, 4, 5, 6, 8, 8,
1450];
1451
1452fn hysteresis_decision(val: i32, thresholds: &[i32], hysteresis: &[i32], prev: i32) -> i32 {
1453    let mut i = 0;
1454    while i < thresholds.len() {
1455        if val < thresholds[i] {
1456            break;
1457        }
1458        i += 1;
1459    }
1460    let mut res = i as i32;
1461    if res > prev && val < thresholds[prev as usize] + hysteresis[prev as usize] {
1462        res = prev;
1463    }
1464    if res < prev && res > 0 && val > thresholds[prev as usize - 1] - hysteresis[prev as usize - 1]
1465    {
1466        res = prev;
1467    }
1468    res
1469}
1470
1471#[allow(clippy::too_many_arguments)]
1472fn alloc_trim_analysis(
1473    mode: &CeltMode,
1474    x: &[f32],
1475    band_log_e: &[f32],
1476    end: usize,
1477    lm: i32,
1478    channels: usize,
1479    n0: usize,
1480    stereo_saving: &mut f32,
1481    tf_estimate: f32,
1482    intensity: i32,
1483    surround_trim: f32,
1484    equiv_rate: i32,
1485) -> i32 {
1486    let mut trim = 5.0f32;
1487    if equiv_rate < 64000 {
1488        trim = 4.0;
1489    } else if equiv_rate < 80000 {
1490        let frac = (equiv_rate - 64000) as f32 / 1024.0;
1491        trim = 4.0 + (1.0 / 16.0) * frac;
1492    }
1493
1494    if channels == 2 {
1495        let mut sum = 0.0f32;
1496        for i in 0..8 {
1497            let offset = (mode.e_bands[i] as usize) << lm;
1498            let n = ((mode.e_bands[i + 1] - mode.e_bands[i]) as usize) << lm;
1499            let mut partial = 0.0f32;
1500            for j in 0..n {
1501                partial += x[offset + j] * x[n0 + offset + j];
1502            }
1503            sum += partial;
1504        }
1505        sum = (sum / 8.0).abs().min(1.0);
1506        let mut min_xc = sum;
1507        for i in 8..intensity as usize {
1508            let offset = (mode.e_bands[i] as usize) << lm;
1509            let n = ((mode.e_bands[i + 1] - mode.e_bands[i]) as usize) << lm;
1510            let mut partial = 0.0f32;
1511            for j in 0..n {
1512                partial += x[offset + j] * x[n0 + offset + j];
1513            }
1514            min_xc = min_xc.min(partial.abs());
1515        }
1516        min_xc = min_xc.min(1.0);
1517
1518        let log_xc = (1.001 - sum * sum).log2();
1519        let log_xc2 = (log_xc * 0.5).max((1.001 - min_xc * min_xc).log2());
1520
1521        trim += (-4.0f32).max(0.75 * log_xc);
1522        *stereo_saving = (*stereo_saving + 0.25).min(-0.5 * log_xc2);
1523    }
1524
1525    let mut diff = 0.0f32;
1526    for c in 0..channels {
1527        for i in 0..end - 1 {
1528            diff += band_log_e[c * mode.nb_ebands + i] * (2 + 2 * i as i32 - end as i32) as f32;
1529        }
1530    }
1531    diff /= (channels * (end - 1)) as f32;
1532    trim -= (-2.0f32).max(2.0f32.min((diff + 1.0) / 6.0));
1533    trim -= surround_trim;
1534    trim -= 2.0 * tf_estimate;
1535
1536    let trim_index = (trim + 0.5).floor() as i32;
1537    trim_index.clamp(0, 10)
1538}
1539
1540#[inline(always)]
1541fn median3(a: f32, b: f32, c: f32) -> f32 {
1542    let mut v = [a, b, c];
1543    crate::compat::sort_by(&mut v[..], |x, y| {
1544        x.partial_cmp(y).unwrap_or(core::cmp::Ordering::Equal)
1545    });
1546    v[1]
1547}
1548
1549#[inline(always)]
1550fn median5(v: &[f32]) -> f32 {
1551    let mut x = [v[0], v[1], v[2], v[3], v[4]];
1552    crate::compat::sort_by(&mut x[..], |a, b| {
1553        a.partial_cmp(b).unwrap_or(core::cmp::Ordering::Equal)
1554    });
1555    x[2]
1556}
1557
1558#[allow(clippy::too_many_arguments)]
1559fn dynalloc_analysis_simple(
1560    mode: &CeltMode,
1561    band_log_e: &[f32],
1562    old_band_e: &[f32],
1563    start: usize,
1564    end: usize,
1565    channels: usize,
1566    lm: usize,
1567    effective_bytes: usize,
1568    is_transient: bool,
1569    offsets: &mut [i32],
1570    cap: &[i32],
1571) {
1572    offsets.fill(0);
1573    if effective_bytes < (30 + 5 * lm) {
1574        return;
1575    }
1576
1577    let nb = mode.nb_ebands;
1578    let mut follower: FixedVec<f32, CELT_NB_X_CH> = FixedVec::from_value(0.0f32, nb * channels);
1579
1580    for c in 0..channels {
1581        let base = c * nb;
1582        let mut band_log_e3: FixedVec<f32, CELT_NB_EBANDS> = FixedVec::from_value(0.0f32, end);
1583        for i in 0..end {
1584            let mut e = band_log_e[base + i];
1585            if lm == 0 && i < 8 {
1586                e = e.max(old_band_e[base + i]);
1587            }
1588            band_log_e3[i] = e;
1589        }
1590
1591        let mut last = 0usize;
1592        follower[base] = band_log_e3[0];
1593        for i in 1..end {
1594            if band_log_e3[i] > band_log_e3[i - 1] + 0.5 {
1595                last = i;
1596            }
1597            follower[base + i] = (follower[base + i - 1] + 1.5).min(band_log_e3[i]);
1598        }
1599        for i in (0..last).rev() {
1600            follower[base + i] =
1601                follower[base + i].min((follower[base + i + 1] + 2.0).min(band_log_e3[i]));
1602        }
1603
1604        let offset = 1.0f32;
1605        if end >= 5 {
1606            for i in 2..end - 2 {
1607                follower[base + i] =
1608                    follower[base + i].max(median5(&band_log_e3[i - 2..i + 3]) - offset);
1609            }
1610        }
1611        if end >= 3 {
1612            let l = median3(band_log_e3[0], band_log_e3[1], band_log_e3[2]) - offset;
1613            follower[base] = follower[base].max(l);
1614            follower[base + 1] = follower[base + 1].max(l);
1615
1616            let r = median3(
1617                band_log_e3[end - 3],
1618                band_log_e3[end - 2],
1619                band_log_e3[end - 1],
1620            ) - offset;
1621            follower[base + end - 2] = follower[base + end - 2].max(r);
1622            follower[base + end - 1] = follower[base + end - 1].max(r);
1623        }
1624    }
1625
1626    if channels == 2 {
1627        for i in start..end {
1628            let l = follower[i];
1629            let r = follower[nb + i];
1630            let r2 = r.max(l - 4.0);
1631            let l2 = l.max(r - 4.0);
1632            follower[i] =
1633                ((band_log_e[i] - l2).max(0.0) + (band_log_e[nb + i] - r2).max(0.0)) * 0.5;
1634        }
1635    } else {
1636        for i in start..end {
1637            follower[i] = (band_log_e[i] - follower[i]).max(0.0);
1638        }
1639    }
1640
1641    if !is_transient {
1642        for i in start..end {
1643            follower[i] *= 0.5;
1644        }
1645    }
1646
1647    let mut tot_boost = 0i32;
1648    for i in start..end {
1649        let mut f = follower[i].min(4.0);
1650        if i < 8 {
1651            f *= 2.0;
1652        }
1653        if i >= 12 {
1654            f *= 0.5;
1655        }
1656
1657        let width = channels as i32 * (mode.e_bands[i + 1] - mode.e_bands[i]) as i32 * (1 << lm);
1658        let (boost, boost_bits) = if width < 6 {
1659            let b = f.floor().max(0.0) as i32;
1660            (b, (b * width) << BITRES)
1661        } else if width > 48 {
1662            let b = (f * 8.0).floor().max(0.0) as i32;
1663            (b, ((b * width) << BITRES) / 8)
1664        } else {
1665            let b = (f * width as f32 / 6.0).floor().max(0.0) as i32;
1666            (b, (b * 6) << BITRES)
1667        };
1668
1669        // Keep dynalloc bounded so allocator still has base bits in CBR usage.
1670        let cap_bits = ((2 * effective_bytes as i32) / 3) << (BITRES + 3);
1671        if tot_boost + boost_bits > cap_bits {
1672            offsets[i] = ((cap_bits - tot_boost) >> BITRES).max(0);
1673            break;
1674        }
1675
1676        let quanta = (width << BITRES).min((6 << BITRES).max(width));
1677        let mut boost_count = boost;
1678        let mut as_bits = boost_count * quanta;
1679        if as_bits > cap[i] {
1680            as_bits = cap[i];
1681            boost_count = as_bits / quanta;
1682        }
1683
1684        offsets[i] = boost_count.max(0);
1685        tot_boost += boost_bits.max(0);
1686    }
1687}
1688
1689impl CeltEncoder {
1690    pub fn new(mode: &'static CeltMode, channels: usize) -> Self {
1691        let overlap = mode.overlap;
1692        let channel_mem_size = 2048 + overlap;
1693        let syn_mem_size = channels * channel_mem_size;
1694        let nb_ebands = mode.nb_ebands;
1695        let nb_x_ch = nb_ebands * channels;
1696        let frame_x_ch = MAX_FRAME_SIZE * channels;
1697        let bufstride_x_ch = (MAX_FRAME_SIZE + overlap) * channels;
1698        Self {
1699            mode,
1700            channels,
1701            complexity: 9,
1702            syn_mem: FixedVec::from_value(0.0, syn_mem_size),
1703            enc_decode_mem: FixedVec::from_value(0.0, syn_mem_size),
1704            old_band_e: FixedVec::from_value(0.0, nb_x_ch),
1705            preemph_mem: FixedVec::from_value(0.0, channels),
1706            tonal_average: 256,
1707            hf_average: 0,
1708            tapset_decision: 0,
1709            spread_decision: SPREAD_NORMAL,
1710            intensity: 0,
1711            last_coded_bands: 0,
1712            prefilter_mem: FixedVec::from_value(0.0, channels * COMBFILTER_MAXPERIOD),
1713            prefilter_period: COMBFILTER_MINPERIOD,
1714            prefilter_gain: 0.0,
1715            prefilter_tapset: 0,
1716            old_band_e2: FixedVec::from_value(0.0, nb_x_ch),
1717            old_band_e3: FixedVec::from_value(0.0, nb_x_ch),
1718            last_band_log_e: FixedVec::from_value(0.0, nb_x_ch),
1719            delayed_intra: 0.0,
1720
1721            w_in_buf: FixedVec::from_value(0.0, bufstride_x_ch),
1722            w_freq: FixedVec::from_value(0.0, frame_x_ch + 4),
1723            w_band_e: FixedVec::from_value(0.0, nb_x_ch),
1724
1725            w_x: FixedVec::from_value(0.0, frame_x_ch + STRIDE_ACCESS_PAD),
1726            w_band_log_e: FixedVec::from_value(0.0, nb_x_ch),
1727            w_error: FixedVec::from_value(0.0, nb_x_ch),
1728            w_tf_res: FixedVec::from_value(0, nb_ebands),
1729            w_cap: FixedVec::from_value(0, nb_ebands),
1730            w_offsets: FixedVec::from_value(0, nb_ebands),
1731            w_pulses: FixedVec::from_value(0, nb_ebands),
1732            w_ebits: FixedVec::from_value(0, nb_x_ch),
1733            w_fine_priority: FixedVec::from_value(0, nb_x_ch),
1734            w_collapse_masks: FixedVec::from_value(0, nb_x_ch),
1735            w_band_amp_synth: FixedVec::from_value(0.0, nb_x_ch),
1736
1737            w_prefilter_pre: FixedVec::from_value(0.0, channels * (COMBFILTER_MAXPERIOD + MAX_FRAME_SIZE)),
1738            w_prefilter_pitch_buf: FixedVec::from_value(0.0, (COMBFILTER_MAXPERIOD + MAX_FRAME_SIZE) >> 1),
1739            w_prefilter_before: FixedVec::from_value(0.0, channels),
1740            w_prefilter_after: FixedVec::from_value(0.0, channels),
1741            w_transient_tmp: FixedVec::from_value(0.0, MAX_TRANSIENT_LEN),
1742            w_transient_tmp2: FixedVec::from_value(0.0, MAX_TRANSIENT_LEN / 2),
1743            consec_transient: 0,
1744
1745            analysis: AnalysisInfo::default(),
1746            loss_rate: 0,
1747        }
1748    }
1749
1750    pub fn encode(&mut self, pcm: &[f32], frame_size: usize, rc: &mut RangeCoder) {
1751        self.encode_impl(pcm, frame_size, rc, 0, None)
1752    }
1753
1754    pub fn encode_with_start_band(
1755        &mut self,
1756        pcm: &[f32],
1757        frame_size: usize,
1758        rc: &mut RangeCoder,
1759        start_band: usize,
1760    ) {
1761        self.encode_impl(pcm, frame_size, rc, start_band, None)
1762    }
1763
1764    pub fn encode_with_budget(
1765        &mut self,
1766        pcm: &[f32],
1767        frame_size: usize,
1768        rc: &mut RangeCoder,
1769        start_band: usize,
1770        total_bits: i32,
1771    ) {
1772        self.encode_impl(pcm, frame_size, rc, start_band, Some(total_bits))
1773    }
1774
1775    fn encode_impl(
1776        &mut self,
1777        pcm: &[f32],
1778        frame_size: usize,
1779        rc: &mut RangeCoder,
1780        start_band: usize,
1781        explicit_total_bits: Option<i32>,
1782    ) {
1783        let mode = self.mode;
1784        let channels = self.channels;
1785        let nb_ebands = mode.nb_ebands;
1786        let overlap = mode.overlap;
1787
1788        let mut lm = 0;
1789        while (mode.short_mdct_size << lm) != frame_size {
1790            lm += 1;
1791            if lm > mode.max_lm {
1792                break;
1793            }
1794        }
1795        if (mode.short_mdct_size << lm) != frame_size {
1796            lm = 0;
1797        }
1798
1799        let syn_mem_size = 2048 + overlap;
1800        for c in 0..channels {
1801            let channel_offset = c * syn_mem_size;
1802
1803            self.syn_mem.copy_within(
1804                channel_offset + frame_size..channel_offset + syn_mem_size,
1805                channel_offset,
1806            );
1807
1808            let mut m = self.preemph_mem[c];
1809            let coef = mode.preemph[0];
1810            for i in 0..frame_size {
1811                let x = pcm[c * frame_size + i] * 32768.0;
1812                let val = x - m;
1813                self.syn_mem[channel_offset + syn_mem_size - frame_size + i] = val;
1814                m = x * coef;
1815            }
1816            self.preemph_mem[c] = m;
1817        }
1818
1819        let buf_stride = frame_size + overlap;
1820        let in_buf = &mut self.w_in_buf[..buf_stride * channels];
1821        for c in 0..channels {
1822            let channel_offset = c * syn_mem_size;
1823            let in_buf_offset = c * buf_stride;
1824
1825            let src_start = syn_mem_size - frame_size - overlap;
1826            in_buf[in_buf_offset..in_buf_offset + buf_stride].copy_from_slice(
1827                &self.syn_mem[channel_offset + src_start..channel_offset + syn_mem_size],
1828            );
1829        }
1830
1831        let mut tf_estimate = 0.0f32;
1832        let mut tf_chan = 0;
1833        let mut weak_transient = false;
1834
1835        let is_transient = if self.complexity >= 1 {
1836            transient_analysis(
1837                in_buf,
1838                buf_stride,
1839                channels,
1840                &mut tf_estimate,
1841                &mut tf_chan,
1842                false,
1843                &mut weak_transient,
1844                0.0,
1845                0.0,
1846                &mut self.w_transient_tmp,
1847                &mut self.w_transient_tmp2,
1848            )
1849        } else {
1850            false
1851        };
1852
1853        // Check for pure tone: if tonality is very high, bypass pitch search
1854        let toneishness = if self.analysis.valid {
1855            self.analysis.tonality
1856        } else {
1857            0.0
1858        };
1859        let _tone_freq = 0.0f32; // Would be set from analysis if available
1860
1861        let pf_enabled =
1862            start_band == 0 && self.complexity >= 5 && toneishness < 0.99 && channels == 1;
1863        let (pf_on, gain1, pitch_index) = if pf_enabled {
1864            run_prefilter(
1865                in_buf,
1866                &mut self.prefilter_mem,
1867                self.prefilter_period,
1868                self.prefilter_gain,
1869                self.prefilter_tapset,
1870                self.tapset_decision,
1871                mode.window,
1872                channels,
1873                frame_size,
1874                overlap,
1875                &mut self.w_prefilter_pre,
1876                &mut self.w_prefilter_pitch_buf,
1877                &mut self.w_prefilter_before,
1878                &mut self.w_prefilter_after,
1879                &self.analysis,
1880                self.loss_rate,
1881            )
1882        } else {
1883            (false, 0.0f32, COMBFILTER_MINPERIOD)
1884        };
1885
1886        // Save the prefiltered overlap for the next frame.
1887        // In libopus, st->in_mem stores the overlap separately and run_prefilter
1888        // copies it to/from in[]. Here we emulate that by updating syn_mem with
1889        // the last overlap samples of in_buf (which were prefiltered in place).
1890        let syn_mem_size = 2048 + overlap;
1891        for c in 0..channels {
1892            let channel_offset = c * syn_mem_size;
1893            let in_buf_offset = c * buf_stride;
1894            self.syn_mem[channel_offset + syn_mem_size - overlap..channel_offset + syn_mem_size]
1895                .copy_from_slice(&in_buf[in_buf_offset + frame_size..in_buf_offset + buf_stride]);
1896        }
1897
1898        let freq = &mut self.w_freq[..frame_size * channels];
1899        let (shift, b) = if is_transient {
1900            (mode.max_lm, 1 << lm)
1901        } else {
1902            (mode.max_lm - lm, 1)
1903        };
1904        let n = frame_size / b;
1905
1906        for c in 0..channels {
1907            let c_buf_offset = c * buf_stride;
1908
1909            if c == 0 && b == 1 && channels == 1 {
1910                let mut max_val = 0.0f32;
1911                let check_len = (frame_size + overlap).min(buf_stride);
1912                for j in 0..check_len {
1913                    max_val = max_val.max(in_buf[c_buf_offset + j].abs());
1914                }
1915            }
1916
1917            for i in 0..b {
1918                mode.mdct.forward(
1919                    &in_buf[c_buf_offset + i * n..],
1920                    &mut freq[c * frame_size + i..],
1921                    mode.window,
1922                    overlap,
1923                    shift,
1924                    b,
1925                );
1926            }
1927        }
1928
1929        let band_e = &mut self.w_band_e[..nb_ebands * channels];
1930        compute_band_energies(mode, freq, band_e, nb_ebands, channels, lm);
1931
1932        let x_pad_end = (frame_size * channels + STRIDE_ACCESS_PAD).min(self.w_x.len());
1933        let x = &mut self.w_x[..x_pad_end];
1934        normalise_bands(
1935            mode,
1936            freq,
1937            x,
1938            band_e,
1939            nb_ebands,
1940            channels,
1941            (1 << lm) as usize,
1942        );
1943
1944        if channels == 1 {
1945            let _ = freq[0];
1946        }
1947
1948        let band_log_e = &mut self.w_band_log_e[..nb_ebands * channels];
1949        crate::bands::amp2log2(mode, start_band, nb_ebands, band_e, band_log_e, channels);
1950
1951        let total_bits = explicit_total_bits.unwrap_or_else(|| (rc.buf.len() * 8) as i32);
1952        self.w_error[..nb_ebands * channels].fill(0.0);
1953        let error = &mut self.w_error[..nb_ebands * channels];
1954
1955        let tell = rc.tell();
1956        let silence = false;
1957        if tell == 1 {
1958            rc.encode_bit_logp(silence, 15);
1959        }
1960
1961        if start_band == 0 && !silence && rc.tell() + 16 <= total_bits {
1962            rc.encode_bit_logp(pf_on, 1);
1963            if pf_on {
1964                let qg = (gain1 / 0.09375 - 1.0 + 0.5).floor() as i32;
1965                let qg = qg.clamp(0, 7);
1966                let pi = (pitch_index + 1) as u32;
1967                let octave = 31 - pi.leading_zeros();
1968                let octave = (octave as i32 - 5).max(0) as u32;
1969                rc.enc_uint(octave, 6);
1970                rc.enc_bits(pi - (16 << octave), 4 + octave);
1971                rc.enc_bits(qg as u32, 3);
1972                rc.encode_icdf(self.tapset_decision, &TAPSET_ICDF, 2);
1973            }
1974        }
1975
1976        let mut short_blocks = false;
1977        if lm > 0 && rc.tell() + 3 <= total_bits {
1978            rc.encode_bit_logp(is_transient, 3);
1979            if is_transient {
1980                short_blocks = true;
1981            }
1982        }
1983
1984        if short_blocks {
1985            let b = 1 << lm;
1986            let n = frame_size / b;
1987            for c in 0..channels {
1988                let c_offset = c * buf_stride;
1989                for i in 0..b {
1990                    mode.mdct.forward(
1991                        &in_buf[c_offset + i * n..c_offset + buf_stride],
1992                        &mut freq[c * frame_size + i..],
1993                        mode.window,
1994                        overlap,
1995                        mode.max_lm,
1996                        b,
1997                    );
1998                }
1999            }
2000
2001            compute_band_energies(mode, freq, band_e, nb_ebands, channels, lm);
2002            normalise_bands(
2003                mode,
2004                freq,
2005                x,
2006                band_e,
2007                nb_ebands,
2008                channels,
2009                (1 << lm) as usize,
2010            );
2011        }
2012
2013        let intra_ener = if self.complexity >= 4 {
2014            false
2015        } else {
2016            self.old_band_e[..nb_ebands * channels]
2017                .iter()
2018                .all(|&e| e <= -27.0)
2019        };
2020        quant_coarse_energy_advanced(
2021            mode,
2022            start_band,
2023            nb_ebands,
2024            nb_ebands,
2025            band_log_e,
2026            &mut self.old_band_e,
2027            total_bits as u32,
2028            error,
2029            rc,
2030            channels,
2031            lm,
2032            (total_bits / 8) as usize,
2033            is_transient || intra_ener,
2034            &mut self.delayed_intra,
2035            self.complexity >= 4,
2036            0,
2037            false,
2038        );
2039        self.w_tf_res[..nb_ebands].fill(0);
2040        let tf_res = &mut self.w_tf_res[..nb_ebands];
2041        let effective_bytes = ((total_bits / 8) as usize).max(1);
2042        let lambda = 80.max(20480 / effective_bytes + 2) as i32;
2043
2044        let tf_select = if self.complexity >= 2 && effective_bytes >= 15 * channels {
2045            tf_analysis(
2046                mode,
2047                nb_ebands,
2048                is_transient,
2049                tf_res,
2050                lambda,
2051                x,
2052                frame_size,
2053                lm as i32,
2054                tf_estimate,
2055                tf_chan,
2056            )
2057        } else {
2058            0
2059        };
2060        tf_encode(
2061            start_band,
2062            nb_ebands,
2063            is_transient,
2064            tf_res,
2065            lm as i32,
2066            tf_select,
2067            rc,
2068        );
2069
2070        let mut dual_stereo_val = if channels == 2 {
2071            stereo_analysis(mode, x, lm as i32, frame_size) as i32
2072        } else {
2073            0
2074        };
2075
2076        let mut stereo_saving = 0.0f32;
2077        let equiv_rate = (total_bits * 48000) / frame_size as i32;
2078        if channels == 2 {
2079            self.intensity = hysteresis_decision(
2080                equiv_rate / 1000,
2081                &INTEN_THRESHOLDS,
2082                &INTEN_HYSTERESIS,
2083                self.intensity,
2084            );
2085            self.intensity = self.intensity.clamp(0, nb_ebands as i32);
2086        }
2087
2088        if self.complexity == 0 {
2089            self.spread_decision = SPREAD_NONE;
2090            if rc.tell() + 4 <= total_bits {
2091                rc.encode_icdf(self.spread_decision, &SPREAD_ICDF, 5);
2092            }
2093        } else if rc.tell() + 4 <= total_bits {
2094            if is_transient || self.complexity < 3 || effective_bytes < 10 * channels {
2095                self.spread_decision = SPREAD_NORMAL;
2096            } else {
2097                let update_hf = lm == mode.max_lm;
2098                let spread_weights = [32i32; 21];
2099                self.spread_decision = spreading_decision(
2100                    mode,
2101                    x,
2102                    &mut self.tonal_average,
2103                    self.spread_decision,
2104                    &mut self.hf_average,
2105                    &mut self.tapset_decision,
2106                    update_hf,
2107                    nb_ebands,
2108                    channels,
2109                    (1 << lm) as usize,
2110                    &spread_weights,
2111                );
2112            }
2113            rc.encode_icdf(self.spread_decision, &SPREAD_ICDF, 5);
2114        } else {
2115            self.spread_decision = SPREAD_NORMAL;
2116        }
2117
2118        self.w_cap[..nb_ebands].fill(0);
2119        let cap = &mut self.w_cap[..nb_ebands];
2120        for (i, cap_i) in cap.iter_mut().enumerate() {
2121            let n = (mode.e_bands[i + 1] - mode.e_bands[i]) << lm;
2122            *cap_i = ((mode.cache.caps[nb_ebands * (2 * lm + channels - 1) + i] as i32 + 64)
2123                * channels as i32
2124                * n as i32)
2125                >> 2;
2126        }
2127
2128        self.w_offsets[..nb_ebands].fill(0);
2129        let offsets = &mut self.w_offsets[..nb_ebands];
2130
2131        dynalloc_analysis_simple(
2132            mode,
2133            band_log_e,
2134            &self.old_band_e,
2135            start_band,
2136            nb_ebands,
2137            channels,
2138            lm,
2139            effective_bytes,
2140            is_transient,
2141            offsets,
2142            cap,
2143        );
2144
2145        let mut dynalloc_logp = 6i32;
2146        let total_bits_bitres = total_bits << BITRES;
2147        let mut total_boost = 0i32;
2148        let mut tell_frac = rc.tell_frac();
2149
2150        for i in start_band..nb_ebands {
2151            let width =
2152                channels as i32 * (mode.e_bands[i + 1] - mode.e_bands[i]) as i32 * (1 << lm);
2153            let quanta = (width << BITRES).min((6 << BITRES).max(width));
2154            let mut dynalloc_loop_logp = dynalloc_logp;
2155            let mut boost = 0i32;
2156            let mut j = 0i32;
2157
2158            while tell_frac + (dynalloc_loop_logp << BITRES) < total_bits_bitres - total_boost
2159                && boost < cap[i]
2160            {
2161                let flag = j < offsets[i];
2162                rc.encode_bit_logp(flag, dynalloc_loop_logp as u32);
2163                tell_frac = rc.tell_frac();
2164                if !flag {
2165                    break;
2166                }
2167                boost += quanta;
2168                total_boost += quanta;
2169                dynalloc_loop_logp = 1;
2170                j += 1;
2171            }
2172
2173            if j > 0 {
2174                dynalloc_logp = 2.max(dynalloc_logp - 1);
2175            }
2176            offsets[i] = boost;
2177        }
2178
2179        let alloc_trim = alloc_trim_analysis(
2180            mode,
2181            x,
2182            band_log_e,
2183            nb_ebands,
2184            lm as i32,
2185            channels,
2186            frame_size,
2187            &mut stereo_saving,
2188            tf_estimate,
2189            self.intensity,
2190            0.0,
2191            equiv_rate,
2192        );
2193        if rc.tell_frac() + (6 << BITRES) <= total_bits_bitres - total_boost {
2194            rc.encode_icdf(alloc_trim, &TRIM_ICDF, 7);
2195        }
2196
2197        let mut intensity = self.intensity;
2198        self.w_pulses[..nb_ebands].fill(0);
2199        let pulses = &mut self.w_pulses[..nb_ebands];
2200
2201        let stereo = channels > 1;
2202        let ebands_stereo = if stereo {
2203            nb_ebands * channels
2204        } else {
2205            nb_ebands
2206        };
2207        self.w_fine_priority[..ebands_stereo].fill(0);
2208        let fine_priority = &mut self.w_fine_priority[..ebands_stereo];
2209        self.w_ebits[..ebands_stereo].fill(0);
2210        let ebits = &mut self.w_ebits[..ebands_stereo];
2211        let mut balance = 0;
2212
2213        self.last_coded_bands = clt_compute_allocation(
2214            mode,
2215            start_band,
2216            nb_ebands,
2217            offsets,
2218            cap,
2219            alloc_trim,
2220            &mut intensity,
2221            &mut dual_stereo_val,
2222            (total_bits << BITRES) - rc.tell_frac() - 1,
2223            &mut balance,
2224            pulses,
2225            ebits,
2226            fine_priority,
2227            channels as i32,
2228            lm as i32,
2229            rc,
2230            true,
2231            0,
2232            nb_ebands as i32 - 1,
2233        );
2234
2235        quant_fine_energy(
2236            mode,
2237            start_band,
2238            nb_ebands,
2239            &mut self.old_band_e,
2240            error,
2241            ebits,
2242            rc,
2243            channels,
2244        );
2245
2246        self.w_collapse_masks[..nb_ebands * channels].fill(0);
2247        let collapse_masks = &mut self.w_collapse_masks[..nb_ebands * channels];
2248        let (x_split, y_split) = x.split_at_mut(frame_size);
2249        let y_opt = if channels == 2 { Some(y_split) } else { None };
2250
2251        let anti_collapse_rsv = if is_transient && lm >= 2 {
2252            let remaining = (total_bits << BITRES) - rc.tell_frac() - 1;
2253            if remaining >= ((lm as i32 + 2) << BITRES) {
2254                1i32 << BITRES
2255            } else {
2256                0
2257            }
2258        } else {
2259            0
2260        };
2261
2262        let mut dual_stereo = dual_stereo_val != 0;
2263
2264        let theta_rdo = channels == 2 && !dual_stereo && self.complexity >= 8;
2265        let resynth = theta_rdo;
2266
2267        quant_all_bands(
2268            true,
2269            mode,
2270            start_band,
2271            nb_ebands,
2272            x_split,
2273            y_opt,
2274            collapse_masks,
2275            band_e,
2276            pulses,
2277            short_blocks,
2278            self.spread_decision,
2279            &mut dual_stereo,
2280            intensity as usize,
2281            tf_res,
2282            (total_bits << BITRES) - anti_collapse_rsv,
2283            &mut balance,
2284            rc,
2285            lm as i32,
2286            self.last_coded_bands,
2287            resynth,
2288            false,
2289            &mut 0u32,
2290        );
2291
2292        if anti_collapse_rsv > 0 {
2293            let anti_collapse_on = if self.consec_transient < 2 {
2294                1u32
2295            } else {
2296                0u32
2297            };
2298            rc.enc_bits(anti_collapse_on, 1);
2299        }
2300
2301        quant_energy_finalise(
2302            mode,
2303            start_band,
2304            nb_ebands,
2305            &mut self.old_band_e,
2306            error,
2307            ebits,
2308            fine_priority,
2309            total_bits - rc.tell(),
2310            rc,
2311            channels,
2312        );
2313
2314        if resynth {
2315            let band_amp_synth = &mut self.w_band_amp_synth[..nb_ebands * channels];
2316            log2amp(mode, nb_ebands, band_amp_synth, &self.old_band_e, channels);
2317            // `w_freq` is no longer needed after analysis/quant; reuse it as the
2318            // synthesis scratch (avoids a separate `w_freq_synth` buffer).
2319            self.w_freq[..frame_size * channels].fill(0.0);
2320            let freq_synth = &mut self.w_freq[..frame_size * channels];
2321            denormalise_bands(
2322                mode,
2323                x,
2324                freq_synth,
2325                band_amp_synth,
2326                start_band,
2327                nb_ebands,
2328                channels,
2329                (1 << lm) as usize,
2330            );
2331            let (syn_shift, syn_b) = if is_transient {
2332                (mode.max_lm, 1 << lm)
2333            } else {
2334                (mode.max_lm - lm, 1)
2335            };
2336            let syn_n = frame_size / syn_b;
2337            let decode_buf_size = 2048;
2338
2339            for c in 0..channels {
2340                let co = c * syn_mem_size;
2341                self.enc_decode_mem
2342                    .copy_within(co + frame_size..co + decode_buf_size + overlap, co);
2343            }
2344
2345            for c in 0..channels {
2346                let co = c * syn_mem_size;
2347                let out_syn_idx = decode_buf_size - frame_size;
2348                for bi in 0..syn_b {
2349                    let syn_stride = if is_transient {
2350                        mode.short_mdct_size
2351                    } else {
2352                        syn_n
2353                    };
2354                    mode.mdct.backward(
2355                        &freq_synth[c * frame_size + bi..],
2356                        &mut self.enc_decode_mem[co + out_syn_idx + bi * syn_stride..],
2357                        mode.window,
2358                        overlap,
2359                        syn_shift,
2360                        syn_b,
2361                    );
2362                }
2363            }
2364        }
2365
2366        self.last_band_log_e.copy_from_slice(&self.old_band_e);
2367
2368        if !is_transient {
2369            self.old_band_e3.copy_from_slice(&self.old_band_e2);
2370            self.old_band_e2.copy_from_slice(&self.old_band_e);
2371        } else {
2372            for i in 0..channels * nb_ebands {
2373                self.old_band_e2[i] = self.old_band_e2[i].min(self.old_band_e[i]);
2374            }
2375        }
2376
2377        rc.pad_to_bits(total_bits);
2378
2379        if pf_on {
2380            self.prefilter_period = pitch_index;
2381            self.prefilter_gain = gain1;
2382            self.prefilter_tapset = self.tapset_decision;
2383        } else {
2384            self.prefilter_period = COMBFILTER_MINPERIOD;
2385            self.prefilter_gain = 0.0;
2386            self.prefilter_tapset = self.tapset_decision;
2387        }
2388
2389        if is_transient {
2390            self.consec_transient += 1;
2391        } else {
2392            self.consec_transient = 0;
2393        }
2394    }
2395}
2396
2397pub struct CeltDecoder {
2398    mode: &'static CeltMode,
2399    channels: usize,
2400    /// Decimation factor for non-48kHz API rates (1=48k, 2=24k, 3=16k, 4=12k, 6=8k).
2401    /// Mirrors libopus `CELTDecoder.downsample`.
2402    downsample: usize,
2403    decode_mem: FixedVec<f32, CELT_DECODE_MEM>,
2404    old_band_e: FixedVec<f32, CELT_NB_X_CH>,
2405    preemph_mem: FixedVec<f32, CELT_MAX_CHANNELS>,
2406    prefilter_mem: FixedVec<f32, CELT_PREFILTER_MEM>,
2407    prefilter_period: usize,
2408    prefilter_period_old: usize,
2409    prefilter_gain: f32,
2410    prefilter_gain_old: f32,
2411    prefilter_tapset: i32,
2412    prefilter_tapset_old: i32,
2413    old_band_e2: FixedVec<f32, CELT_NB_X_CH>,
2414    old_band_e3: FixedVec<f32, CELT_NB_X_CH>,
2415    rng: u32,
2416
2417    w_tf_res: FixedVec<i32, CELT_NB_EBANDS>,
2418    w_cap: FixedVec<i32, CELT_NB_EBANDS>,
2419    w_offsets: FixedVec<i32, CELT_NB_EBANDS>,
2420    w_pulses: FixedVec<i32, CELT_NB_EBANDS>,
2421    w_ebits: FixedVec<i32, CELT_NB_X_CH>,
2422    w_fine_priority: FixedVec<i32, CELT_NB_X_CH>,
2423    w_x: FixedVec<f32, CELT_W_X_DEC>,
2424    w_collapse_masks: FixedVec<u32, CELT_NB_X_CH>,
2425    w_freq: FixedVec<f32, CELT_W_FREQ_DEC>,
2426    w_band_amp: FixedVec<f32, CELT_NB_X_CH>,
2427    w_pcm_frame: FixedVec<f32, DECODE_BUFFER_SIZE>,
2428    w_post: FixedVec<f32, { DECODE_BUFFER_SIZE + COMBFILTER_MAXPERIOD }>,
2429}
2430
2431impl CeltDecoder {
2432    /// Create a CELT decoder. `sampling_rate` is the API sampling rate
2433    /// (8000–48000); the decoder always operates at 48 kHz internally and
2434    /// decimates by `resampling_factor(sampling_rate)` on output.
2435    pub fn new(mode: &'static CeltMode, channels: usize, sampling_rate: i32) -> Self {
2436        let overlap = mode.overlap;
2437        let nb_ebands = mode.nb_ebands;
2438        let nb_x_ch = nb_ebands * channels;
2439        let dec_frame_x_ch = DECODE_BUFFER_SIZE * channels;
2440        Self {
2441            mode,
2442            channels,
2443            downsample: resampling_factor(sampling_rate),
2444            decode_mem: FixedVec::from_value(0.0, channels * (DECODE_BUFFER_SIZE + overlap)),
2445            old_band_e: FixedVec::from_value(0.0, nb_x_ch),
2446            preemph_mem: FixedVec::from_value(0.0, channels),
2447            prefilter_mem: FixedVec::from_value(0.0, channels * COMBFILTER_MAXPERIOD),
2448            prefilter_period: COMBFILTER_MINPERIOD,
2449            prefilter_period_old: COMBFILTER_MINPERIOD,
2450            prefilter_gain: 0.0,
2451            prefilter_gain_old: 0.0,
2452            prefilter_tapset: 0,
2453            prefilter_tapset_old: 0,
2454            old_band_e2: FixedVec::from_value(0.0, nb_x_ch),
2455            old_band_e3: FixedVec::from_value(0.0, nb_x_ch),
2456            rng: 0,
2457
2458            w_tf_res: FixedVec::from_value(0, nb_ebands),
2459            w_cap: FixedVec::from_value(0, nb_ebands),
2460            w_offsets: FixedVec::from_value(0, nb_ebands),
2461            w_pulses: FixedVec::from_value(0, nb_ebands),
2462            w_ebits: FixedVec::from_value(0, nb_x_ch),
2463            w_fine_priority: FixedVec::from_value(0, nb_x_ch),
2464
2465            w_x: FixedVec::from_value(0.0, dec_frame_x_ch + STRIDE_ACCESS_PAD),
2466            w_collapse_masks: FixedVec::from_value(0, nb_x_ch),
2467            w_freq: FixedVec::from_value(0.0, dec_frame_x_ch + 4), // +4: NEON backward pre-rotation reads up to 3 elements past n2
2468            w_band_amp: FixedVec::from_value(0.0, nb_x_ch),
2469            w_pcm_frame: FixedVec::from_value(0.0, DECODE_BUFFER_SIZE),
2470            w_post: FixedVec::from_value(0.0, DECODE_BUFFER_SIZE + COMBFILTER_MAXPERIOD),
2471        }
2472    }
2473
2474    pub fn decode(&mut self, compressed: &[u8], frame_size: usize, pcm: &mut [f32]) -> usize {
2475        self.decode_impl(compressed, frame_size, pcm, 0, self.mode.nb_ebands)
2476    }
2477
2478    /// Reset all decoder state (equivalent to libopus `OPUS_RESET_STATE`).
2479    /// Used at SILK↔CELT mode transitions to avoid cross-mode artifacts.
2480    pub fn reset_state(&mut self) {
2481        self.decode_mem.fill(0.0);
2482        self.old_band_e.fill(0.0);
2483        self.preemph_mem.fill(0.0);
2484        self.prefilter_mem.fill(0.0);
2485        self.prefilter_period = COMBFILTER_MINPERIOD;
2486        self.prefilter_period_old = COMBFILTER_MINPERIOD;
2487        self.prefilter_gain = 0.0;
2488        self.prefilter_gain_old = 0.0;
2489        self.prefilter_tapset = 0;
2490        self.prefilter_tapset_old = 0;
2491        self.old_band_e2.fill(0.0);
2492        self.old_band_e3.fill(0.0);
2493        self.rng = 0;
2494    }
2495
2496    pub fn decode_with_start_band(
2497        &mut self,
2498        compressed: &[u8],
2499        frame_size: usize,
2500        pcm: &mut [f32],
2501        start_band: usize,
2502    ) -> usize {
2503        self.decode_impl(compressed, frame_size, pcm, start_band, self.mode.nb_ebands)
2504    }
2505
2506    pub fn decode_from_range_coder(
2507        &mut self,
2508        rc: &mut RangeCoder,
2509        total_bits: i32,
2510        frame_size: usize,
2511        pcm: &mut [f32],
2512        start_band: usize,
2513    ) -> usize {
2514        self.decode_impl_from_rc(
2515            rc,
2516            total_bits,
2517            frame_size,
2518            pcm,
2519            start_band,
2520            self.mode.nb_ebands,
2521        )
2522    }
2523
2524    pub fn decode_from_range_coder_with_band_range(
2525        &mut self,
2526        rc: &mut RangeCoder,
2527        total_bits: i32,
2528        frame_size: usize,
2529        pcm: &mut [f32],
2530        start_band: usize,
2531        end_band: usize,
2532    ) -> usize {
2533        self.decode_impl_from_rc(rc, total_bits, frame_size, pcm, start_band, end_band)
2534    }
2535
2536    fn decode_impl(
2537        &mut self,
2538        compressed: &[u8],
2539        frame_size: usize,
2540        pcm: &mut [f32],
2541        start_band: usize,
2542        end_band: usize,
2543    ) -> usize {
2544        let total_bits = (compressed.len() * 8) as i32;
2545        let mut rc = RangeCoder::new_decoder(compressed);
2546        self.decode_impl_from_rc(&mut rc, total_bits, frame_size, pcm, start_band, end_band)
2547    }
2548
2549    fn decode_impl_from_rc(
2550        &mut self,
2551        rc: &mut RangeCoder,
2552        total_bits: i32,
2553        frame_size: usize,
2554        pcm: &mut [f32],
2555        start_band: usize,
2556        end_band: usize,
2557    ) -> usize {
2558        let mode = self.mode;
2559        let channels = self.channels;
2560        let nb_ebands = mode.nb_ebands;
2561        let end_band = end_band.min(nb_ebands).max(start_band);
2562        let overlap = mode.overlap;
2563
2564        // The API frame_size is in output samples. Internally CELT always
2565        // decodes at 48 kHz, so upscale by the downsample factor (libopus
2566        // celt_decoder.c:1196: `frame_size *= st->downsample`).
2567        let api_frame_size = frame_size;
2568        let frame_size = frame_size * self.downsample;
2569
2570        let mut lm = 0;
2571        while (mode.short_mdct_size << lm) != frame_size {
2572            lm += 1;
2573            if lm > mode.max_lm {
2574                break;
2575            }
2576        }
2577        if (mode.short_mdct_size << lm) != frame_size {
2578            lm = 0;
2579        }
2580
2581        let tell = rc.tell();
2582        let mut silence = false;
2583        if tell >= total_bits {
2584            silence = true;
2585        } else if tell == 1 {
2586            silence = rc.decode_bit_logp(15);
2587        }
2588
2589        if silence {
2590            pcm[..api_frame_size * channels].fill(0.0);
2591            return api_frame_size;
2592        }
2593
2594        let mut pf_on = false;
2595        let mut pitch_index = COMBFILTER_MINPERIOD;
2596        let mut gain1 = 0.0f32;
2597        let mut prefilter_tapset = 0;
2598
2599        if start_band == 0 && !silence && rc.tell() + 16 <= total_bits {
2600            pf_on = rc.decode_bit_logp(1);
2601            if pf_on {
2602                let octave = rc.dec_uint(6);
2603                pitch_index = ((16 << octave) + rc.dec_bits(4 + octave)) as usize - 1;
2604                let qg = rc.dec_bits(3);
2605                if rc.tell() + 2 <= total_bits {
2606                    prefilter_tapset = rc.decode_icdf(&TAPSET_ICDF, 2) as usize;
2607                }
2608                gain1 = 0.09375 * (qg as f32 + 1.0);
2609            }
2610        }
2611        if start_band != 0 {
2612            self.prefilter_gain = 0.0;
2613        }
2614
2615        let mut is_transient = false;
2616        if lm > 0 && rc.tell() + 3 <= total_bits {
2617            is_transient = rc.decode_bit_logp(3);
2618        }
2619        let short_blocks = is_transient;
2620
2621        let intra_ener = if rc.tell() + 3 <= total_bits {
2622            rc.decode_bit_logp(3)
2623        } else {
2624            false
2625        };
2626
2627        unquant_coarse_energy(
2628            mode,
2629            start_band,
2630            end_band,
2631            &mut self.old_band_e,
2632            intra_ener,
2633            rc,
2634            channels,
2635            lm,
2636        );
2637        self.w_tf_res[..nb_ebands].fill(0);
2638        let tf_res = &mut self.w_tf_res[..nb_ebands];
2639        tf_decode(start_band, end_band, is_transient, tf_res, lm as i32, rc);
2640
2641        let spread_decision = if rc.tell() + 4 <= total_bits {
2642            rc.decode_icdf(&SPREAD_ICDF, 5)
2643        } else {
2644            SPREAD_NORMAL
2645        };
2646
2647        self.w_cap[..nb_ebands].fill(0);
2648        let cap = &mut self.w_cap[..nb_ebands];
2649        for (i, cap_i) in cap.iter_mut().enumerate() {
2650            let n = (mode.e_bands[i + 1] - mode.e_bands[i]) << lm;
2651            *cap_i = ((mode.cache.caps[nb_ebands * (2 * lm + channels - 1) + i] as i32 + 64)
2652                * channels as i32
2653                * n as i32)
2654                >> 2;
2655        }
2656
2657        self.w_offsets[..nb_ebands].fill(0);
2658        let offsets = &mut self.w_offsets[..nb_ebands];
2659        let mut dynalloc_logp = 6i32;
2660        let mut total_bits_bitres = total_bits << BITRES;
2661        let mut tell_frac = rc.tell_frac();
2662        for i in start_band..end_band {
2663            let width =
2664                channels as i32 * (mode.e_bands[i + 1] - mode.e_bands[i]) as i32 * (1 << lm);
2665            let quanta = (width << BITRES).min((6i32 << BITRES).max(width));
2666            let mut dynalloc_loop_logp = dynalloc_logp;
2667            let mut boost = 0i32;
2668            while tell_frac + (dynalloc_loop_logp << BITRES) < total_bits_bitres && boost < cap[i] {
2669                let flag = rc.decode_bit_logp(dynalloc_loop_logp as u32);
2670                tell_frac = rc.tell_frac();
2671                if !flag {
2672                    break;
2673                }
2674                boost += quanta;
2675                total_bits_bitres -= quanta;
2676                dynalloc_loop_logp = 1;
2677            }
2678            offsets[i] = boost;
2679            if boost > 0 {
2680                dynalloc_logp = dynalloc_logp.max(2) - 1;
2681                dynalloc_logp = dynalloc_logp.max(2);
2682            }
2683        }
2684
2685        let alloc_trim = if rc.tell_frac() + (6 << BITRES) <= total_bits_bitres {
2686            rc.decode_icdf(&TRIM_ICDF, 7)
2687        } else {
2688            5
2689        };
2690        let anti_collapse_rsv = if is_transient && lm >= 2 {
2691            let remaining = (total_bits << BITRES) - rc.tell_frac() - 1;
2692            if remaining >= ((lm as i32 + 2) << BITRES) {
2693                1i32 << BITRES
2694            } else {
2695                0
2696            }
2697        } else {
2698            0
2699        };
2700
2701        let mut intensity = 0;
2702        let mut dual_stereo_val = if channels == 2 { 1 } else { 0 };
2703        let mut balance = 0;
2704        self.w_pulses[..nb_ebands].fill(0);
2705        let pulses = &mut self.w_pulses[..nb_ebands];
2706
2707        let ebands_stereo = if channels > 1 {
2708            nb_ebands * channels
2709        } else {
2710            nb_ebands
2711        };
2712        self.w_fine_priority[..ebands_stereo].fill(0);
2713        let fine_priority = &mut self.w_fine_priority[..ebands_stereo];
2714        self.w_ebits[..ebands_stereo].fill(0);
2715        let ebits = &mut self.w_ebits[..ebands_stereo];
2716
2717        let alloc_bits = (total_bits << BITRES) - rc.tell_frac() - 1 - anti_collapse_rsv;
2718        let coded_bands = clt_compute_allocation(
2719            mode,
2720            start_band,
2721            end_band,
2722            offsets,
2723            cap,
2724            alloc_trim,
2725            &mut intensity,
2726            &mut dual_stereo_val,
2727            alloc_bits,
2728            &mut balance,
2729            pulses,
2730            ebits,
2731            fine_priority,
2732            channels as i32,
2733            lm as i32,
2734            rc,
2735            false,
2736            0,
2737            end_band as i32 - 1,
2738        );
2739
2740        unquant_fine_energy(
2741            mode,
2742            start_band,
2743            end_band,
2744            &mut self.old_band_e,
2745            ebits,
2746            rc,
2747            channels,
2748        );
2749
2750        if frame_size > DECODE_BUFFER_SIZE + overlap {
2751            return 0;
2752        }
2753
2754        self.w_x[..frame_size * channels].fill(0.0);
2755
2756        let x_pad_end = (frame_size * channels + STRIDE_ACCESS_PAD).min(self.w_x.len());
2757        let x = &mut self.w_x[..x_pad_end];
2758        self.w_collapse_masks[..nb_ebands * channels].fill(0);
2759        let collapse_masks = &mut self.w_collapse_masks[..nb_ebands * channels];
2760
2761        let (x_split, y_split) = x.split_at_mut(frame_size);
2762        let y_opt = if channels == 2 { Some(y_split) } else { None };
2763
2764        let mut dual_stereo = dual_stereo_val != 0;
2765        self.w_band_amp[..nb_ebands * channels].fill(0.0);
2766        let band_amp = &mut self.w_band_amp[..nb_ebands * channels];
2767        log2amp(mode, nb_ebands, band_amp, &self.old_band_e, channels);
2768        quant_all_bands(
2769            false,
2770            mode,
2771            start_band,
2772            end_band,
2773            x_split,
2774            y_opt,
2775            collapse_masks,
2776            band_amp,
2777            pulses,
2778            short_blocks,
2779            spread_decision,
2780            &mut dual_stereo,
2781            intensity as usize,
2782            tf_res,
2783            (total_bits << BITRES) - anti_collapse_rsv,
2784            &mut balance,
2785            rc,
2786            lm as i32,
2787            coded_bands,
2788            true,
2789            false,
2790            &mut self.rng,
2791        );
2792        // Trace X values for comparison with C decoder
2793        let mut anti_collapse_on = false;
2794        if anti_collapse_rsv > 0 {
2795            anti_collapse_on = rc.dec_bits(1) != 0;
2796        }
2797
2798        unquant_energy_finalise(
2799            mode,
2800            start_band,
2801            end_band,
2802            &mut self.old_band_e,
2803            ebits,
2804            fine_priority,
2805            total_bits - rc.tell(),
2806            rc,
2807            channels,
2808        );
2809        if anti_collapse_on {
2810            self.rng = crate::bands::anti_collapse(
2811                mode,
2812                x,
2813                collapse_masks,
2814                lm as i32,
2815                channels,
2816                frame_size,
2817                start_band,
2818                nb_ebands,
2819                &self.old_band_e,
2820                &self.old_band_e2,
2821                &self.old_band_e3,
2822                pulses,
2823                self.rng,
2824            );
2825        }
2826
2827        // Recompute band_amp after unquant_energy_finalise, which adjusts old_band_e.
2828        // (Mirrors the encoder's resynth path: log2amp is called after quant_energy_finalise.)
2829        log2amp(mode, nb_ebands, band_amp, &self.old_band_e, channels);
2830        self.w_freq[..frame_size * channels].fill(0.0);
2831        let freq = &mut self.w_freq[..frame_size * channels];
2832        denormalise_bands(
2833            mode,
2834            x,
2835            freq,
2836            band_amp,
2837            start_band,
2838            end_band,
2839            channels,
2840            (1 << lm) as usize,
2841        );
2842        // Anti-aliasing: zero MDCT bins above the output Nyquist when
2843        // downsampling (libopus denormalise_bands `if(downsample!=1)
2844        // bound=IMIN(bound,N/downsample)`).
2845        if self.downsample > 1 {
2846            let bound = frame_size / self.downsample;
2847            for c in 0..channels {
2848                for i in bound..frame_size {
2849                    freq[c * frame_size + i] = 0.0;
2850                }
2851            }
2852        }
2853        // Always trace freq and band_amp for comparison
2854
2855        let (shift, b) = if short_blocks {
2856            (mode.max_lm, 1 << lm)
2857        } else {
2858            (mode.max_lm - lm, 1)
2859        };
2860        let n = frame_size / b;
2861
2862        for c in 0..channels {
2863            let channel_mem_offset = c * (DECODE_BUFFER_SIZE + overlap);
2864
2865            let mem_size = DECODE_BUFFER_SIZE + overlap;
2866            self.decode_mem.copy_within(
2867                channel_mem_offset + frame_size..channel_mem_offset + mem_size,
2868                channel_mem_offset,
2869            );
2870
2871            let out_syn_idx = DECODE_BUFFER_SIZE - frame_size;
2872
2873            for i in 0..b {
2874                let block_freq_idx = c * frame_size + i;
2875                // Stride between short-block MDCT outputs is short_mdct_size (not n).
2876                // In libopus: out_syn[c] + NB*b, where NB = mode->shortMdctSize.
2877                // For non-transient b=1, i*n == 0 either way.
2878                let block_stride = if short_blocks {
2879                    mode.short_mdct_size
2880                } else {
2881                    n
2882                };
2883                let block_out_idx = channel_mem_offset + out_syn_idx + i * block_stride;
2884                let available_len = self.decode_mem.len() - block_out_idx;
2885                if available_len < n + overlap {
2886                    panic!(
2887                        "MDCT backward buffer too small: need {}, have {} (out_syn_idx={}, n={}, overlap={})",
2888                        n + overlap,
2889                        available_len,
2890                        out_syn_idx,
2891                        n,
2892                        overlap
2893                    );
2894                }
2895                self.mode.mdct.backward(
2896                    &freq[block_freq_idx..],
2897                    &mut self.decode_mem[block_out_idx..],
2898                    mode.window,
2899                    overlap,
2900                    shift,
2901                    b,
2902                );
2903            }
2904
2905            const SIG_SAT: f32 = 536870911.0;
2906            for i in 0..frame_size {
2907                let v = &mut self.decode_mem[channel_mem_offset + out_syn_idx + i];
2908                *v = v.clamp(-SIG_SAT, SIG_SAT);
2909            }
2910
2911            self.w_pcm_frame[..frame_size].fill(0.0);
2912            let pcm_frame = &mut self.w_pcm_frame[..frame_size];
2913
2914            pcm_frame.copy_from_slice(
2915                &self.decode_mem[channel_mem_offset + out_syn_idx
2916                    ..channel_mem_offset + out_syn_idx + frame_size],
2917            );
2918            if pf_on || self.prefilter_gain > 0.0 || self.prefilter_gain_old > 0.0 {
2919                // Set up w_post = [prefilter_mem | pcm_frame] for history access.
2920                // We apply combfilter in-place on w_post[COMBFILTER_MAXPERIOD..] so that
2921                // later samples can reference already-filtered earlier samples, matching C's
2922                // in-place comb_filter behavior.
2923                self.w_post[..COMBFILTER_MAXPERIOD].copy_from_slice(
2924                    &self.prefilter_mem[c * COMBFILTER_MAXPERIOD..(c + 1) * COMBFILTER_MAXPERIOD],
2925                );
2926                self.w_post[COMBFILTER_MAXPERIOD..COMBFILTER_MAXPERIOD + frame_size]
2927                    .copy_from_slice(pcm_frame);
2928
2929                let short_n = mode.short_mdct_size;
2930                // Call 1: first short_n samples, transition old→current params
2931                // Apply in-place on w_post[COMBFILTER_MAXPERIOD..], output overwrites input
2932                comb_filter_inplace(
2933                    &mut self.w_post,
2934                    COMBFILTER_MAXPERIOD,
2935                    self.prefilter_period_old,
2936                    self.prefilter_period,
2937                    short_n,
2938                    self.prefilter_gain_old,
2939                    self.prefilter_gain,
2940                    self.prefilter_tapset_old,
2941                    self.prefilter_tapset,
2942                    mode.window,
2943                    overlap,
2944                );
2945                if lm != 0 {
2946                    // Call 2: remaining N-short_n samples, transition current→new params
2947                    comb_filter_inplace(
2948                        &mut self.w_post,
2949                        COMBFILTER_MAXPERIOD + short_n,
2950                        self.prefilter_period,
2951                        pitch_index,
2952                        frame_size - short_n,
2953                        self.prefilter_gain,
2954                        gain1,
2955                        self.prefilter_tapset,
2956                        prefilter_tapset as i32,
2957                        mode.window,
2958                        overlap,
2959                    );
2960                }
2961
2962                pcm_frame.copy_from_slice(
2963                    &self.w_post[COMBFILTER_MAXPERIOD..COMBFILTER_MAXPERIOD + frame_size],
2964                );
2965
2966                self.decode_mem[channel_mem_offset + out_syn_idx
2967                    ..channel_mem_offset + out_syn_idx + frame_size]
2968                    .copy_from_slice(pcm_frame);
2969            }
2970            let mut new_mem = [0.0f32; COMBFILTER_MAXPERIOD];
2971            if frame_size >= COMBFILTER_MAXPERIOD {
2972                new_mem.copy_from_slice(&pcm_frame[frame_size - COMBFILTER_MAXPERIOD..frame_size]);
2973            } else {
2974                new_mem[..COMBFILTER_MAXPERIOD - frame_size].copy_from_slice(
2975                    &self.prefilter_mem
2976                        [c * COMBFILTER_MAXPERIOD + frame_size..(c + 1) * COMBFILTER_MAXPERIOD],
2977                );
2978                new_mem[COMBFILTER_MAXPERIOD - frame_size..].copy_from_slice(pcm_frame);
2979            }
2980            self.prefilter_mem[c * COMBFILTER_MAXPERIOD..(c + 1) * COMBFILTER_MAXPERIOD]
2981                .copy_from_slice(&new_mem);
2982
2983            let coef = mode.preemph[0];
2984            let mut m = self.preemph_mem[c];
2985            const VERY_SMALL: f32 = 1e-30f32;
2986            let ds = self.downsample;
2987            if ds == 1 {
2988                for i in 0..frame_size {
2989                    let x = pcm_frame[i];
2990                    let val = (x + VERY_SMALL + m).clamp(-SIG_SAT, SIG_SAT);
2991                    pcm[c * api_frame_size + i] = val * (1.0 / 32768.0);
2992                    m = val * coef;
2993                }
2994            } else {
2995                // Run deemphasis IIR over all internal samples, but only write
2996                // every downsample-th sample (libopus deemphasis() stride).
2997                for i in 0..frame_size {
2998                    let x = pcm_frame[i];
2999                    let val = (x + VERY_SMALL + m).clamp(-SIG_SAT, SIG_SAT);
3000                    if i % ds == 0 {
3001                        pcm[c * api_frame_size + i / ds] = val * (1.0 / 32768.0);
3002                    }
3003                    m = val * coef;
3004                }
3005            }
3006            self.preemph_mem[c] = m;
3007        }
3008
3009        self.prefilter_period_old = self.prefilter_period;
3010        self.prefilter_gain_old = self.prefilter_gain;
3011        self.prefilter_tapset_old = self.prefilter_tapset;
3012
3013        if pf_on {
3014            self.prefilter_period = pitch_index;
3015            self.prefilter_gain = gain1;
3016            self.prefilter_tapset = prefilter_tapset as i32;
3017        } else {
3018            self.prefilter_period = COMBFILTER_MINPERIOD;
3019            self.prefilter_gain = 0.0;
3020            self.prefilter_tapset = 0;
3021        }
3022
3023        if lm > 0 {
3024            self.prefilter_period_old = self.prefilter_period;
3025            self.prefilter_gain_old = self.prefilter_gain;
3026            self.prefilter_tapset_old = self.prefilter_tapset;
3027        }
3028
3029        if !is_transient {
3030            self.old_band_e3.copy_from_slice(&self.old_band_e2);
3031            self.old_band_e2.copy_from_slice(&self.old_band_e);
3032        } else {
3033            let nb_ebands = mode.nb_ebands;
3034            for i in 0..channels * nb_ebands {
3035                self.old_band_e2[i] = self.old_band_e2[i].min(self.old_band_e[i]);
3036            }
3037        }
3038
3039        self.rng = rc.rng;
3040
3041        api_frame_size
3042    }
3043}
3044
3045#[cfg(all(test, feature = "std"))]
3046mod tests {
3047    use super::*;
3048    use crate::{modes, range_coder::RangeCoder};
3049
3050    // Regression test: directly drive CeltEncoder with an invalid frame_size=48,
3051    // bypassing the OpusEncoder::encode() validation layer.
3052    //
3053    // This reproduces the crash that was reported against opus-rs 0.1.19 when
3054    // G.729-decoded PCM (8 kHz) reached the 48 kHz Opus encoder without correct
3055    // resampling, producing a 48-sample frame instead of 480.
3056    //
3057    // Root cause: the lm-search in encode_impl finds no valid match for frame_size=48
3058    // (valid sizes are 120, 240, 480, 960) and silently falls back to lm=0.
3059    // With lm=0 and shift=max_lm=3: n=1920>>3=240, n2=120, overlap2=60.
3060    // The in_buf slice has only frame_size+overlap=168 elements, but forward()
3061    // requires input.len() >= n2+overlap2 = 180, so it panics immediately.
3062    // In opus-rs 0.1.19 this assertion was absent and the crash reached the MDCT
3063    // output write: "index out of bounds: the len is 48 but the index is 119".
3064    //
3065    // Either way: the call panics, confirming the crash path is real.
3066    // The fix in OpusEncoder::encode() returns Err before reaching CeltEncoder.
3067    #[test]
3068    #[should_panic]
3069    fn test_celt_frame_size_48_panics_confirms_crash_path() {
3070        let mode = modes::default_mode();
3071        let mut enc = CeltEncoder::new(mode, 1);
3072        // frame_size=48: lm-search fails, falls back to lm=0.
3073        // forward() will panic — either on the input-size assertion (0.1.21+) or
3074        // on the output write (0.1.19): "len is 48 but the index is 119".
3075        let pcm = vec![0.0f32; 48 + mode.overlap]; // supply ≥ frame_size samples
3076        let mut rc = RangeCoder::new_encoder(100);
3077        enc.encode_with_budget(&pcm, 48, &mut rc, 0, 800);
3078    }
3079}