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