Skip to main content

rusty_opus/
pvq.rs

1use crate::range_coder::RangeCoder;
2use std::mem::MaybeUninit;
3
4pub const CELT_PVQ_U_DATA: [u32; 1272] = [
5    1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
6    0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
7    0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
8    0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
9    0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
10    0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
11    1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
12    1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
13    1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
14    1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
15    1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
16    1, 3, 5, 7, 9, 11, 13, 15, 17, 19, 21, 23, 25, 27, 29, 31, 33, 35, 37, 39, 41, 43, 45, 47, 49,
17    51, 53, 55, 57, 59, 61, 63, 65, 67, 69, 71, 73, 75, 77, 79, 81, 83, 85, 87, 89, 91, 93, 95, 97,
18    99, 101, 103, 105, 107, 109, 111, 113, 115, 117, 119, 121, 123, 125, 127, 129, 131, 133, 135,
19    137, 139, 141, 143, 145, 147, 149, 151, 153, 155, 157, 159, 161, 163, 165, 167, 169, 171, 173,
20    175, 177, 179, 181, 183, 185, 187, 189, 191, 193, 195, 197, 199, 201, 203, 205, 207, 209, 211,
21    213, 215, 217, 219, 221, 223, 225, 227, 229, 231, 233, 235, 237, 239, 241, 243, 245, 247, 249,
22    251, 253, 255, 257, 259, 261, 263, 265, 267, 269, 271, 273, 275, 277, 279, 281, 283, 285, 287,
23    289, 291, 293, 295, 297, 299, 301, 303, 305, 307, 309, 311, 313, 315, 317, 319, 321, 323, 325,
24    327, 329, 331, 333, 335, 337, 339, 341, 343, 345, 347, 349, 351, 13, 25, 41, 61, 85, 113, 145,
25    181, 221, 265, 313, 365, 421, 481, 545, 613, 685, 761, 841, 925, 1013, 1105, 1201, 1301, 1405,
26    1513, 1625, 1741, 1861, 1985, 2113, 2245, 2381, 2521, 2665, 2813, 2965, 3121, 3281, 3445, 3613,
27    3785, 3961, 4141, 4325, 4513, 4705, 4901, 5101, 5305, 5513, 5725, 5941, 6161, 6385, 6613, 6845,
28    7081, 7321, 7565, 7813, 8065, 8321, 8581, 8845, 9113, 9385, 9661, 9941, 10225, 10513, 10805,
29    11101, 11401, 11705, 12013, 12325, 12641, 12961, 13285, 13613, 13945, 14281, 14621, 14965,
30    15313, 15665, 16021, 16381, 16745, 17113, 17485, 17861, 18241, 18625, 19013, 19405, 19801,
31    20201, 20605, 21013, 21425, 21841, 22261, 22685, 23113, 23545, 23981, 24421, 24865, 25313,
32    25765, 26221, 26681, 27145, 27613, 28085, 28561, 29041, 29525, 30013, 30505, 31001, 31501,
33    32005, 32513, 33025, 33541, 34061, 34585, 35113, 35645, 36181, 36721, 37265, 37813, 38365,
34    38921, 39481, 40045, 40613, 41185, 41761, 42341, 42925, 43513, 44105, 44701, 45301, 45905,
35    46513, 47125, 47741, 48361, 48985, 49613, 50245, 50881, 51521, 52165, 52813, 53465, 54121,
36    54781, 55445, 56113, 56785, 57461, 58141, 58825, 59513, 60205, 60901, 61601, 63, 129, 231, 377,
37    575, 833, 1159, 1561, 2047, 2625, 3303, 4089, 4991, 6017, 7175, 8473, 9919, 11521, 13287,
38    15225, 17343, 19649, 22151, 24857, 27775, 30913, 34279, 37881, 41727, 45825, 50183, 54809,
39    59711, 64897, 70375, 76153, 82239, 88641, 95367, 102425, 109823, 117569, 125671, 134137,
40    142975, 152193, 161799, 171801, 182207, 193025, 204263, 215929, 228031, 240577, 253575, 267033,
41    280959, 295361, 310247, 325625, 341503, 357889, 374791, 392217, 410175, 428673, 447719, 467321,
42    487487, 508225, 529543, 551449, 573951, 597057, 620775, 645113, 670079, 695681, 721927, 748825,
43    776383, 804609, 833511, 863097, 893375, 924353, 956039, 988441, 1021567, 1055425, 1090023,
44    1125369, 1161471, 1198337, 1235975, 1274393, 1313599, 1353601, 1394407, 1436025, 1478463,
45    1521729, 1565831, 1610777, 1656575, 1703233, 1750759, 1799161, 1848447, 1898625, 1949703,
46    2001689, 2054591, 2108417, 2163175, 2218873, 2275519, 2333121, 2391687, 2451225, 2511743,
47    2573249, 2635751, 2699257, 2763775, 2829313, 2895879, 2963481, 3032127, 3101825, 3172583,
48    3244409, 3317311, 3391297, 3466375, 3542553, 3619839, 3698241, 3777767, 3858425, 3940223,
49    4023169, 4107271, 4192537, 4278975, 4366593, 4455399, 4545401, 4636607, 4729025, 4822663,
50    4917529, 5013631, 5110977, 5209575, 5309433, 5410559, 5512961, 5616647, 5721625, 5827903,
51    5935489, 6044391, 6154617, 6266175, 6379073, 6493319, 6608921, 6725887, 6844225, 6963943,
52    7085049, 7207551, 321, 681, 1289, 2241, 3649, 5641, 8361, 11969, 16641, 22569, 29961, 39041,
53    50049, 63241, 78889, 97281, 118721, 143529, 172041, 204609, 241601, 283401, 330409, 383041,
54    441729, 506921, 579081, 658689, 746241, 842249, 947241, 1061761, 1186369, 1321641, 1468169,
55    1626561, 1797441, 1981449, 2179241, 2391489, 2618881, 2862121, 3121929, 3399041, 3694209,
56    4008201, 4341801, 4695809, 5071041, 5468329, 5888521, 6332481, 6801089, 7295241, 7815849,
57    8363841, 8940161, 9545769, 10181641, 10848769, 11548161, 12280841, 13047849, 13850241,
58    14689089, 15565481, 16480521, 17435329, 18431041, 19468809, 20549801, 21675201, 22846209,
59    24064041, 25329929, 26645121, 28010881, 29428489, 30899241, 32424449, 34005441, 35643561,
60    37340169, 39096641, 40914369, 42794761, 44739241, 46749249, 48826241, 50971689, 53187081,
61    55473921, 57833729, 60268041, 62778409, 65366401, 68033601, 70781609, 73612041, 76526529,
62    79526721, 82614281, 85790889, 89058241, 92418049, 95872041, 99421961, 103069569, 106816641,
63    110664969, 114616361, 118672641, 122835649, 127107241, 131489289, 135983681, 140592321,
64    145317129, 150160041, 155123009, 160208001, 165417001, 170752009, 176215041, 181808129,
65    187533321, 193392681, 199388289, 205522241, 211796649, 218213641, 224775361, 231483969,
66    238341641, 245350569, 252512961, 259831041, 267307049, 274943241, 282741889, 290705281,
67    298835721, 307135529, 315607041, 324252609, 333074601, 342075401, 351257409, 360623041,
68    370174729, 379914921, 389846081, 399970689, 410291241, 420810249, 431530241, 442453761,
69    453583369, 464921641, 476471169, 488234561, 500214441, 512413449, 524834241, 537479489,
70    550351881, 563454121, 576788929, 590359041, 604167209, 618216201, 632508801, 1683, 3653, 7183,
71    13073, 22363, 36365, 56695, 85305, 124515, 177045, 246047, 335137, 448427, 590557, 766727,
72    982729, 1244979, 1560549, 1937199, 2383409, 2908411, 3522221, 4235671, 5060441, 6009091,
73    7095093, 8332863, 9737793, 11326283, 13115773, 15124775, 17372905, 19880915, 22670725,
74    25765455, 29189457, 32968347, 37129037, 41699767, 46710137, 52191139, 58175189, 64696159,
75    71789409, 79491819, 87841821, 96879431, 106646281, 117185651, 128542501, 140763503, 153897073,
76    167993403, 183104493, 199284183, 216588185, 235074115, 254801525, 275831935, 298228865,
77    322057867, 347386557, 374284647, 402823977, 433078547, 465124549, 499040399, 534906769,
78    572806619, 612825229, 655050231, 699571641, 746481891, 795875861, 847850911, 902506913,
79    959946283, 1020274013, 1083597703, 1150027593, 1219676595, 1292660325, 1369097135, 1449108145,
80    1532817275, 1620351277, 1711839767, 1807415257, 1907213187, 2011371957, 2120032959, 8989,
81    19825, 40081, 75517, 134245, 227305, 369305, 579125, 880685, 1303777, 1884961, 2668525,
82    3707509, 5064793, 6814249, 9041957, 11847485, 15345233, 19665841, 24957661, 31388293, 39146185,
83    48442297, 59511829, 72616013, 88043969, 106114625, 127178701, 151620757, 179861305, 212358985,
84    249612805, 292164445, 340600625, 395555537, 457713341, 527810725, 606639529, 695049433,
85    793950709, 904317037, 1027188385, 1163673953, 1314955181, 1482288821, 1667010073, 1870535785,
86    2094367717, 48639, 108545, 224143, 433905, 795455, 1392065, 2340495, 3800305, 5984767, 9173505,
87    13726991, 20103025, 28875327, 40754369, 56610575, 77500017, 104692735, 139703809, 184327311,
88    240673265, 311207743, 398796225, 506750351, 638878193, 799538175, 993696769, 1226990095,
89    1505789553, 1837271615, 2229491905, 265729, 598417, 1256465, 2485825, 4673345, 8405905,
90    14546705, 24331777, 39490049, 62390545, 96220561, 145198913, 214828609, 312193553, 446304145,
91    628496897, 872893441, 1196924561, 1621925137, 2173806145, 1462563, 3317445, 7059735, 14218905,
92    27298155, 50250765, 89129247, 152951073, 254831667, 413442773, 654862247, 1014889769,
93    1541911931, 2300409629, 3375210671, 8097453, 18474633, 39753273, 81270333, 158819253,
94    298199265, 540279585, 948062325, 1616336765, 45046719, 103274625, 224298231, 464387817,
95    921406335, 1759885185, 3248227095, 251595969, 579168825, 1267854873, 2653649025, 1409933619,
96];
97
98const CELT_PVQ_U_ROW: [u32; 15] = [
99    0, 176, 351, 525, 698, 870, 1041, 1131, 1178, 1207, 1226, 1240, 1248, 1254, 1257,
100];
101
102#[inline(always)]
103pub fn celt_pvq_u_lookup(n: u32, k: u32) -> u32 {
104    let r = n.min(k) as usize;
105    let c = n.max(k) as usize;
106
107    if r >= CELT_PVQ_U_ROW.len() {
108        return compute_u(n, k);
109    }
110    unsafe {
111        let row_base = *CELT_PVQ_U_ROW.get_unchecked(r);
112        let idx = row_base as usize + c;
113        if idx >= CELT_PVQ_U_DATA.len() {
114            return compute_u(n, k);
115        }
116        *CELT_PVQ_U_DATA.get_unchecked(idx)
117    }
118}
119
120const MAX_PVQ_K: usize = 128;
121const MAX_PVQ_U: usize = MAX_PVQ_K + 2;
122pub const MAX_PVQ_N: usize = 352;
123
124pub fn ncwrs(n: u32, k: u32) -> u32 {
125    if n == 0 {
126        return 0;
127    }
128    if n == 1 {
129        return if k > 0 { 2 } else { 1 };
130    }
131    let mut u = [0u32; MAX_PVQ_U];
132    u[0] = 0;
133    u[1] = 1;
134    for ki in 2..=(k + 1) as usize {
135        u[ki] = (ki as u32 * 2).wrapping_sub(1);
136    }
137    let mut curr_n = n;
138    while curr_n > 2 {
139        unext(&mut u[1..], (k + 1) as usize, 1);
140        curr_n -= 1;
141    }
142    u[k as usize].wrapping_add(u[k as usize + 1])
143}
144
145fn compute_u(n: u32, k: u32) -> u32 {
146    if n == 0 {
147        return if k == 0 { 1 } else { 0 };
148    }
149    if n == 1 {
150        return 1;
151    }
152    let mut u = [0u32; MAX_PVQ_U];
153    u[0] = 0;
154    u[1] = 1;
155    for ki in 2..=(k + 1) as usize {
156        u[ki] = (ki as u32 * 2).wrapping_sub(1);
157    }
158    let mut curr_n = n;
159    while curr_n > 2 {
160        unext(&mut u[1..], (k + 1) as usize, 1);
161        curr_n -= 1;
162    }
163    u[k as usize]
164}
165
166#[inline(always)]
167pub fn celt_pvq_u(n: u32, k: u32) -> u32 {
168    celt_pvq_u_lookup(n, k)
169}
170
171#[inline(always)]
172pub fn celt_pvq_v(n: u32, k: u32) -> u32 {
173    celt_pvq_u_lookup(n, k).wrapping_add(celt_pvq_u_lookup(n, k + 1))
174}
175
176fn unext(u: &mut [u32], len: usize, mut u0: u32) {
177    let mut j = 1;
178    while j < len {
179        let u1 = u[j].wrapping_add(u[j - 1]).wrapping_add(u0);
180        u[j - 1] = u0;
181        u0 = u1;
182        j += 1;
183    }
184    u[j - 1] = u0;
185}
186
187#[inline(always)]
188pub fn icwrs(n: u32, _k: u32, y: &[i32]) -> u32 {
189    if n == 1 {
190        return if y[0] < 0 { 1 } else { 0 };
191    }
192    debug_assert!(n >= 2, "icwrs: n must be >= 2");
193    let mut j = (n - 1) as usize;
194
195    let mut i: u32 = if y[j] < 0 { 1 } else { 0 };
196    let mut k = y[j].unsigned_abs();
197
198    while j > 0 {
199        j -= 1;
200        let yj = y[j];
201        let m = n - j as u32;
202        i = i.wrapping_add(celt_pvq_u_lookup(m, k));
203        k += yj.unsigned_abs();
204
205        let sign_mask = yj >> 31;
206        let lookup = (sign_mask as u32) & celt_pvq_u_lookup(m, k + 1);
207        i = i.wrapping_add(lookup);
208    }
209    i
210}
211
212#[inline(always)]
213pub fn cwrsi(n: u32, k: u32, mut i: u32, y: &mut [i32]) {
214    debug_assert!(k > 0, "cwrsi: k must be > 0");
215
216    if n == 1 {
217        let s = -(i as i32);
218        y[0] = ((k as i32) + s) ^ s;
219        return;
220    }
221
222    let mut curr_n = n;
223
224    let mut curr_k = k as i32;
225    let mut j = 0usize;
226
227    while curr_n > 2 {
228        if curr_k >= curr_n as i32 {
229            let p_kp1 = celt_pvq_u_lookup(curr_n, (curr_k + 1) as u32);
230            let s: i32 = if i >= p_kp1 {
231                i -= p_kp1;
232                -1
233            } else {
234                0
235            };
236            let k0 = curr_k;
237            let q = celt_pvq_u_lookup(curr_n, curr_n);
238            let mut p;
239            if q > i {
240                curr_k = curr_n as i32;
241                loop {
242                    curr_k -= 1;
243                    p = celt_pvq_u_lookup(curr_n, curr_k.max(0) as u32);
244                    if p <= i || curr_k <= 0 {
245                        break;
246                    }
247                }
248            } else {
249                p = celt_pvq_u_lookup(curr_n, curr_k as u32);
250                while p > i && curr_k > 0 {
251                    curr_k -= 1;
252                    p = celt_pvq_u_lookup(curr_n, curr_k as u32);
253                }
254            }
255            i -= p;
256            let val = k0 - curr_k;
257            y[j] = (val + s) ^ s;
258        } else {
259            let p_k = celt_pvq_u_lookup(curr_k as u32, curr_n);
260            let p_kp1 = celt_pvq_u_lookup((curr_k + 1) as u32, curr_n);
261            if p_k <= i && i < p_kp1 {
262                i -= p_k;
263                y[j] = 0;
264                j += 1;
265                curr_n -= 1;
266                continue;
267            }
268            let s: i32 = if i >= p_kp1 {
269                i -= p_kp1;
270                -1
271            } else {
272                0
273            };
274            let k0 = curr_k;
275
276            let mut p;
277            loop {
278                curr_k -= 1;
279                p = celt_pvq_u_lookup(curr_k.max(0) as u32, curr_n);
280                if p <= i || curr_k <= 0 {
281                    break;
282                }
283            }
284            i -= p;
285            let val = k0 - curr_k;
286            y[j] = (val + s) ^ s;
287        }
288        j += 1;
289        curr_n -= 1;
290    }
291
292    let p2 = (2u32).wrapping_mul(curr_k as u32).wrapping_add(1);
293    let s2: i32 = if i >= p2 {
294        i -= p2;
295        -1
296    } else {
297        0
298    };
299    let k0 = curr_k;
300    curr_k = ((i + 1) >> 1) as i32;
301    if curr_k > 0 {
302        i -= 2 * curr_k as u32 - 1;
303    }
304    y[j] = ((k0 - curr_k) + s2) ^ s2;
305    j += 1;
306
307    let s1 = -(i as i32);
308    y[j] = (curr_k + s1) ^ s1;
309}
310
311#[inline(always)]
312pub fn encode_pulses(y: &[i32], n: u32, k: u32, rc: &mut RangeCoder) {
313    if k == 0 {
314        return;
315    }
316    let fl = icwrs(n, k, y);
317    let ft = celt_pvq_v(n, k);
318    debug_assert!(fl < ft, "encode_pulses: fl={fl} >= ft={ft}, n={n}, k={k}");
319    rc.enc_uint(fl, ft);
320}
321
322#[inline(always)]
323pub fn decode_pulses(y: &mut [i32], n: u32, k: u32, rc: &mut RangeCoder) {
324    if k == 0 {
325        for i in 0..n as usize {
326            y[i] = 0;
327        }
328        return;
329    }
330    let ft = celt_pvq_v(n, k);
331    let fl = rc.dec_uint(ft).min(ft.saturating_sub(1));
332    cwrsi(n, k, fl, y);
333}
334
335#[allow(non_snake_case)]
336fn op_pvq_refine(
337    Xn: &[f32],
338    iy: &mut [i32],
339    iy0: Option<&[i32]>,
340    K: i32,
341    up: i32,
342    margin: i32,
343    N: usize,
344) -> bool {
345    let K8 = (K as f32) * 256.0;
346
347    let mut iysum = 0i32;
348    for i in 0..N {
349        let tmp = K8 * Xn[i];
350        iy[i] = (tmp + 128.0) as i32 >> 8;
351        iysum += iy[i];
352    }
353
354    if let Some(iy0_ref) = iy0 {
355        for i in 0..N {
356            let min_val = up * iy0_ref[i] - (margin - 1);
357            let max_val = up * iy0_ref[i] + (margin - 1);
358            iy[i] = iy[i].clamp(min_val, max_val);
359        }
360        iysum = iy.iter().sum();
361    }
362
363    if (iysum - K).abs() > 32 {
364        return true;
365    }
366
367    let dir = if iysum < K { 1 } else { -1 };
368    let mut remaining = (K - iysum).abs();
369
370    let mut rounding: [f32; 32] = [0.0; 32];
371    for i in 0..N {
372        rounding[i] = K8 * Xn[i] - ((iy[i] as f32) * 256.0);
373    }
374
375    while remaining > 0 {
376        let mut best_i = 0;
377        let mut best_round = if dir == 1 { -1e30f32 } else { 1e30f32 };
378
379        for i in 0..N {
380            let can_adjust = if dir == 1 {
381                iy0.is_none_or(|iy0_ref| (iy[i] - up * iy0_ref[i]).abs() < (margin - 1))
382            } else {
383                iy[i] != 0
384                    && iy0.is_none_or(|iy0_ref| (iy[i] - up * iy0_ref[i]).abs() < (margin - 1))
385            };
386
387            if can_adjust
388                && ((dir == 1 && rounding[i] > best_round)
389                    || (dir == -1 && rounding[i] < best_round && iy[i] != 0))
390            {
391                best_round = rounding[i];
392                best_i = i;
393            }
394        }
395
396        iy[best_i] += dir;
397        rounding[best_i] -= dir as f32 * 256.0;
398        remaining -= 1;
399    }
400
401    false
402}
403
404pub fn pvq_search_qext(
405    x: &[f32],
406    y: &mut [i32],
407    up_y: &mut [i32],
408    refine: &mut [i32],
409    k: i32,
410    extra_bits: i32,
411    n: usize,
412) -> f32 {
413    debug_assert!(n <= 32);
414    debug_assert!(extra_bits >= 2);
415
416    let mut sum = 0.0f32;
417    for i in 0..n {
418        sum += x[i].abs();
419    }
420
421    if sum < 1e-15 {
422        y[0] = k;
423        up_y[0] = ((1 << extra_bits) - 1) * k;
424        for i in 1..n {
425            y[i] = 0;
426            up_y[i] = 0;
427            refine[i] = 0;
428        }
429        refine[0] = 0;
430        return (up_y[0] as f32) * (up_y[0] as f32);
431    }
432
433    #[allow(non_snake_case)]
434    let mut Xn: [f32; 32] = [0.0; 32];
435    let rcp_sum = 1.0 / sum;
436    for i in 0..n {
437        Xn[i] = x[i].abs() * rcp_sum;
438    }
439
440    let failed1 = op_pvq_refine(&Xn, y, None, k, 1, k + 1, n);
441
442    let up = (1 << extra_bits) - 1;
443    let up_k = up * k;
444    let margin = up;
445    let failed2 = op_pvq_refine(&Xn, up_y, Some(y), up_k, up, margin, n);
446
447    if failed1 || failed2 {
448        y[0] = k;
449        up_y[0] = up_k;
450        for i in 1..n {
451            y[i] = 0;
452            up_y[i] = 0;
453        }
454    }
455
456    for i in 0..n {
457        refine[i] = up_y[i] - up * y[i];
458    }
459
460    let mut yy = 0.0f32;
461    for i in 0..n {
462        yy += (up_y[i] as f32) * (up_y[i] as f32);
463    }
464
465    for i in 0..n {
466        if x[i] < 0.0 {
467            y[i] = -y[i];
468            up_y[i] = -up_y[i];
469            refine[i] = -refine[i];
470        }
471    }
472
473    yy
474}
475
476#[inline(always)]
477fn pvq_search_n2(x: &[f32], y: &mut [i32], k: i32) {
478    debug_assert!(x.len() >= 2 && y.len() >= 2);
479
480    let abs_x0 = x[0].abs();
481    let abs_x1 = x[1].abs();
482    let sum = abs_x0 + abs_x1;
483
484    if sum < 1e-15 {
485        y[0] = k;
486        y[1] = 0;
487        return;
488    }
489
490    let rcp_sum = 1.0 / sum;
491    let y0 = (k as f32 * abs_x0 * rcp_sum + 0.5).floor() as i32;
492    let y0 = y0.clamp(0, k);
493    let y1 = k - y0;
494
495    y[0] = if x[0] >= 0.0 { y0 } else { -y0 };
496    y[1] = if x[1] >= 0.0 { y1 } else { -y1 };
497}
498
499#[inline]
500fn pvq_search_n4(x: &[f32], y: &mut [i32], k: i32) {
501    debug_assert!(x.len() >= 4 && y.len() >= 4);
502
503    if k == 0 {
504        y[0] = 0;
505        y[1] = 0;
506        y[2] = 0;
507        y[3] = 0;
508        return;
509    }
510
511    #[cfg(target_arch = "x86_64")]
512    unsafe {
513        use std::arch::x86_64::*;
514
515        let sign_mask = _mm_castsi128_ps(_mm_set1_epi32(0x7FFF_FFFFu32 as i32));
516        let vx = _mm_loadu_ps(x.as_ptr());
517        let vabs = _mm_and_ps(vx, sign_mask);
518
519        let vzero_f = _mm_setzero_ps();
520
521        let vneg_mask = _mm_cmplt_ps(vx, vzero_f);
522
523        let vsigns = _mm_and_si128(_mm_castps_si128(vneg_mask), _mm_set1_epi32(1));
524
525        let vabs_x = vabs;
526        let mut vy2f = _mm_setzero_ps();
527        let mut vy = _mm_setzero_si128();
528        let mut xy = 0.0f32;
529        let mut yy = 0.0f32;
530
531        let vtwo = _mm_set1_ps(2.0);
532
533        let vone_i = _mm_set1_epi32(1);
534
535        for _ in 0..k {
536            let vxy = _mm_set1_ps(xy);
537            let vrxy = _mm_add_ps(vabs_x, vxy);
538            let vyy1 = _mm_add_ps(vy2f, _mm_set1_ps(yy + 1.0));
539
540            let vscore = _mm_mul_ps(vrxy, _mm_rsqrt_ps(vyy1));
541
542            let s0 = _mm_cvtss_f32(vscore);
543            let s1 = _mm_cvtss_f32(_mm_shuffle_ps(vscore, vscore, 0b01_01_01_01));
544            let s2 = _mm_cvtss_f32(_mm_shuffle_ps(vscore, vscore, 0b10_10_10_10));
545            let s3 = _mm_cvtss_f32(_mm_shuffle_ps(vscore, vscore, 0b11_11_11_11));
546            let mut best_score = s0;
547            let mut best_i: u32 = 0;
548            if s1 > best_score {
549                best_score = s1;
550                best_i = 1;
551            }
552            if s2 > best_score {
553                best_score = s2;
554                best_i = 2;
555            }
556            if s3 > best_score {
557                best_i = 3;
558            }
559            let _ = best_score;
560
561            let vbest = _mm_set1_epi32(best_i as i32);
562            let vlane = _mm_setr_epi32(0, 1, 2, 3);
563            let vmask = _mm_castsi128_ps(_mm_cmpeq_epi32(vlane, vbest));
564
565            let vpick_ax = _mm_and_ps(vabs_x, vmask);
566
567            let vpick_ax_hi = _mm_movehl_ps(vpick_ax, vpick_ax);
568            let vpick_ax2 = _mm_add_ps(vpick_ax, vpick_ax_hi);
569            let vpick_ax3 = _mm_add_ss(vpick_ax2, _mm_shuffle_ps(vpick_ax2, vpick_ax2, 1));
570            xy += _mm_cvtss_f32(vpick_ax3);
571
572            let vpick_ryy = _mm_and_ps(vyy1, vmask);
573            let vpick_ryy_hi = _mm_movehl_ps(vpick_ryy, vpick_ryy);
574            let vpick_ryy2 = _mm_add_ps(vpick_ryy, vpick_ryy_hi);
575            let vpick_ryy3 = _mm_add_ss(vpick_ryy2, _mm_shuffle_ps(vpick_ryy2, vpick_ryy2, 1));
576            yy = _mm_cvtss_f32(vpick_ryy3);
577
578            let vadd2 = _mm_and_ps(vtwo, vmask);
579            vy2f = _mm_add_ps(vy2f, vadd2);
580
581            let vadd1 = _mm_and_si128(vone_i, _mm_castps_si128(vmask));
582            vy = _mm_add_epi32(vy, vadd1);
583        }
584
585        let vneg_s = _mm_sub_epi32(_mm_setzero_si128(), vsigns);
586        let vy_xor = _mm_xor_si128(vy, vneg_s);
587        let vy_out = _mm_add_epi32(vy_xor, vsigns);
588        _mm_storeu_si128(y.as_mut_ptr() as *mut __m128i, vy_out);
589    }
590
591    #[cfg(not(target_arch = "x86_64"))]
592    {
593        let ax0 = x[0].abs();
594        let ax1 = x[1].abs();
595        let ax2 = x[2].abs();
596        let ax3 = x[3].abs();
597        let s0 = (x[0] < 0.0) as i32;
598        let s1 = (x[1] < 0.0) as i32;
599        let s2 = (x[2] < 0.0) as i32;
600        let s3 = (x[3] < 0.0) as i32;
601        let mut xy = 0.0f32;
602        let mut yy = 0.0f32;
603        let mut y2f0 = 0.0f32;
604        let mut y2f1 = 0.0f32;
605        let mut y2f2 = 0.0f32;
606        let mut y2f3 = 0.0f32;
607        let mut y0 = 0i32;
608        let mut y1 = 0i32;
609        let mut y2 = 0i32;
610        let mut y3 = 0i32;
611        for _ in 0..k {
612            let rxy0 = xy + ax0;
613            let sq0 = rxy0 * rxy0;
614            let ryy0 = yy + y2f0 + 1.0;
615            let rxy1 = xy + ax1;
616            let sq1 = rxy1 * rxy1;
617            let ryy1 = yy + y2f1 + 1.0;
618            let rxy2 = xy + ax2;
619            let sq2 = rxy2 * rxy2;
620            let ryy2 = yy + y2f2 + 1.0;
621            let rxy3 = xy + ax3;
622            let sq3 = rxy3 * rxy3;
623            let ryy3 = yy + y2f3 + 1.0;
624            let mut bsq = sq0;
625            let mut bden = ryy0;
626            let mut best_i: u32 = 0;
627            if bden * sq1 > ryy1 * bsq {
628                bsq = sq1;
629                bden = ryy1;
630                best_i = 1;
631            }
632            if bden * sq2 > ryy2 * bsq {
633                bsq = sq2;
634                bden = ryy2;
635                best_i = 2;
636            }
637            if bden * sq3 > ryy3 * bsq {
638                best_i = 3;
639            }
640            let _ = bsq;
641            match best_i {
642                0 => {
643                    xy += ax0;
644                    yy = ryy0;
645                    y2f0 += 2.0;
646                    y0 += 1;
647                }
648                1 => {
649                    xy += ax1;
650                    yy = ryy1;
651                    y2f1 += 2.0;
652                    y1 += 1;
653                }
654                2 => {
655                    xy += ax2;
656                    yy = ryy2;
657                    y2f2 += 2.0;
658                    y2 += 1;
659                }
660                _ => {
661                    xy += ax3;
662                    yy = ryy3;
663                    y2f3 += 2.0;
664                    y3 += 1;
665                }
666            }
667        }
668        y[0] = (y0 ^ -s0) + s0;
669        y[1] = (y1 ^ -s1) + s1;
670        y[2] = (y2 ^ -s2) + s2;
671        y[3] = (y3 ^ -s3) + s3;
672    }
673}
674
675#[inline(always)]
676pub fn pvq_search(x: &[f32], y: &mut [i32], k: i32, n: usize) {
677    if k == 1 {
678        let mut best_i = 0;
679        let mut best_abs = x[0].abs();
680        for i in 1..n {
681            let abs_xi = x[i].abs();
682            if abs_xi > best_abs {
683                best_abs = abs_xi;
684                best_i = i;
685            }
686        }
687        for j in 0..n {
688            y[j] = 0;
689        }
690        let sign: i32 = if x[best_i] >= 0.0 { 1 } else { -1 };
691        y[best_i] = sign;
692        return;
693    }
694
695    if n == 2 {
696        pvq_search_n2(x, y, k);
697        return;
698    }
699
700    if n == 4 {
701        pvq_search_n4(x, y, k);
702        return;
703    }
704
705    if n >= 32 {
706        pvq_search_fast_select(x, y, k, n);
707        return;
708    }
709
710    #[cfg(target_arch = "aarch64")]
711    if n <= 16 {
712        pvq_search_neon(x, y, k, n);
713        return;
714    }
715
716    #[cfg(target_arch = "x86_64")]
717    if k > 4 && std::arch::is_x86_feature_detected!("avx2") {
718        unsafe {
719            pvq_search_avx2(x, y, k, n);
720        }
721        return;
722    }
723
724    pvq_search_scalar(x, y, k, n);
725}
726
727#[cfg(target_arch = "aarch64")]
728#[inline(always)]
729#[allow(unsafe_op_in_unsafe_fn)]
730unsafe fn pvq_fast_select_init_neon(
731    x: &[f32],
732    n: usize,
733    abs_x: &mut [MaybeUninit<f32>; MAX_PVQ_N],
734    signs: &mut [MaybeUninit<i32>; MAX_PVQ_N],
735) -> f32 {
736    use std::arch::aarch64::*;
737
738    let mut sum_vec = vdupq_n_f32(0.0);
739    let mut i = 0;
740
741    while i + 16 <= n {
742        let vx0 = vld1q_f32(x.as_ptr().add(i));
743        let vx1 = vld1q_f32(x.as_ptr().add(i + 4));
744        let vx2 = vld1q_f32(x.as_ptr().add(i + 8));
745        let vx3 = vld1q_f32(x.as_ptr().add(i + 12));
746
747        let vabs0 = vabsq_f32(vx0);
748        let vabs1 = vabsq_f32(vx1);
749        let vabs2 = vabsq_f32(vx2);
750        let vabs3 = vabsq_f32(vx3);
751
752        vst1q_f32(abs_x.as_mut_ptr().add(i) as *mut f32, vabs0);
753        vst1q_f32(abs_x.as_mut_ptr().add(i + 4) as *mut f32, vabs1);
754        vst1q_f32(abs_x.as_mut_ptr().add(i + 8) as *mut f32, vabs2);
755        vst1q_f32(abs_x.as_mut_ptr().add(i + 12) as *mut f32, vabs3);
756
757        sum_vec = vaddq_f32(sum_vec, vabs0);
758        sum_vec = vaddq_f32(sum_vec, vabs1);
759        sum_vec = vaddq_f32(sum_vec, vabs2);
760        sum_vec = vaddq_f32(sum_vec, vabs3);
761
762        for j in 0..16 {
763            signs[i + j].write(if x[i + j] < 0.0 { -1i32 } else { 1i32 });
764        }
765
766        i += 16;
767    }
768
769    while i + 8 <= n {
770        let vx0 = vld1q_f32(x.as_ptr().add(i));
771        let vx1 = vld1q_f32(x.as_ptr().add(i + 4));
772
773        let vabs0 = vabsq_f32(vx0);
774        let vabs1 = vabsq_f32(vx1);
775
776        vst1q_f32(abs_x.as_mut_ptr().add(i) as *mut f32, vabs0);
777        vst1q_f32(abs_x.as_mut_ptr().add(i + 4) as *mut f32, vabs1);
778
779        sum_vec = vaddq_f32(sum_vec, vabs0);
780        sum_vec = vaddq_f32(sum_vec, vabs1);
781
782        for j in 0..8 {
783            signs[i + j].write(if x[i + j] < 0.0 { -1i32 } else { 1i32 });
784        }
785
786        i += 8;
787    }
788
789    while i + 4 <= n {
790        let vx = vld1q_f32(x.as_ptr().add(i));
791        let vabs = vabsq_f32(vx);
792        vst1q_f32(abs_x.as_mut_ptr().add(i) as *mut f32, vabs);
793        sum_vec = vaddq_f32(sum_vec, vabs);
794
795        for j in 0..4 {
796            signs[i + j].write(if x[i + j] < 0.0 { -1i32 } else { 1i32 });
797        }
798
799        i += 4;
800    }
801
802    let mut sum = vaddvq_f32(sum_vec);
803
804    for j in i..n {
805        let abs_xi = x[j].abs();
806        abs_x[j].write(abs_xi);
807        sum += abs_xi;
808        signs[j].write(if x[j] < 0.0 { -1i32 } else { 1i32 });
809    }
810
811    sum
812}
813
814#[inline]
815pub fn pvq_search_fast_select(x: &[f32], y: &mut [i32], k: i32, n: usize) -> f32 {
816    let mut k = k;
817    let mut yy = 0.0f32;
818    let mut xy = 0.0f32;
819
820    y[..n].fill(0);
821
822    if k <= 0 {
823        return 0.0;
824    }
825
826    let mut abs_x_mu = [MaybeUninit::<f32>::uninit(); MAX_PVQ_N];
827    let mut signs_mu = [MaybeUninit::<i32>::uninit(); MAX_PVQ_N];
828
829    #[cfg(target_arch = "aarch64")]
830    let sum = unsafe { pvq_fast_select_init_neon(x, n, &mut abs_x_mu, &mut signs_mu) };
831    #[cfg(not(target_arch = "aarch64"))]
832    let sum = {
833        let mut s = 0.0f32;
834        for i in 0..n {
835            abs_x_mu[i].write(x[i].abs());
836            signs_mu[i].write(if x[i] < 0.0 { -1i32 } else { 1i32 });
837            s += unsafe { abs_x_mu[i].assume_init() };
838        }
839        s
840    };
841
842    let abs_x = unsafe { std::slice::from_raw_parts(abs_x_mu.as_ptr() as *const f32, n) };
843    let signs = unsafe { std::slice::from_raw_parts(signs_mu.as_ptr() as *const i32, n) };
844
845    if k > (n >> 1) as i32 && sum > 1e-15 {
846        let rcp = (k as f32 + 0.8) / sum;
847
848        let abs_x_ptr = abs_x.as_ptr();
849        let y_ptr = y.as_mut_ptr();
850        unsafe {
851            for i in 0..n {
852                let yi = (*abs_x_ptr.add(i) * rcp) as i32;
853                *y_ptr.add(i) = yi;
854                let yf = yi as f32;
855                yy += yf * yf;
856                xy += yf * *abs_x_ptr.add(i);
857                k -= yi;
858            }
859        }
860
861        if k > n as i32 + 3 {
862            let tmp = k as f32;
863            unsafe {
864                yy += tmp * tmp + tmp * *y_ptr as f32;
865                *y_ptr += k;
866            }
867            k = 0;
868        }
869    }
870
871    const BATCH_SIZE: i32 = 4;
872
873    if k < BATCH_SIZE * 2 || n < 16 {
874        #[cfg(target_arch = "aarch64")]
875        {
876            use std::arch::aarch64::*;
877            let mut y2f_mu = [MaybeUninit::<f32>::uninit(); MAX_PVQ_N];
878            for i in 0..n {
879                y2f_mu[i].write(2.0 * y[i] as f32);
880            }
881            let y2f = unsafe {
882                std::slice::from_raw_parts_mut(y2f_mu.as_mut_ptr() as *mut f32, MAX_PVQ_N)
883            };
884
885            let abs_x_ptr = abs_x.as_ptr();
886            let y2f_ptr = y2f.as_mut_ptr();
887            let y_ptr = y.as_mut_ptr();
888            unsafe {
889                let n4 = n & !3;
890                while k > 0 {
891                    yy += 1.0;
892                    let vxy = vdupq_n_f32(xy);
893                    let vyy = vdupq_n_f32(yy);
894                    let mut vmax = vdupq_n_f32(0.0);
895                    let mut best_id: usize = 0;
896
897                    let mut i = 0;
898                    while i < n4 {
899                        let vx = vld1q_f32(abs_x_ptr.add(i));
900                        let vy = vld1q_f32(y2f_ptr.add(i));
901                        let rxy = vaddq_f32(vx, vxy);
902                        let ryy = vaddq_f32(vy, vyy);
903                        let inv_sqrt = vrsqrteq_f32(ryy);
904                        let score = vmulq_f32(rxy, inv_sqrt);
905                        vmax = vmaxq_f32(vmax, score);
906                        let sc = std::slice::from_raw_parts(
907                            &score as *const float32x4_t as *const f32,
908                            4,
909                        );
910                        let mx = vmaxvq_f32(vmax);
911                        for lane in 0..4 {
912                            if sc[lane] == mx {
913                                best_id = i + lane;
914                            }
915                        }
916                        i += 4;
917                    }
918
919                    while i < n {
920                        let rxy = xy + *abs_x_ptr.add(i);
921                        let ryy = yy + *y2f_ptr.add(i);
922                        let score = rxy * (1.0 / ryy.sqrt());
923                        let current_max = vmaxvq_f32(vmax);
924                        if score > current_max {
925                            best_id = i;
926                            vmax = vsetq_lane_f32(score, vmax, 0);
927                        }
928                        i += 1;
929                    }
930
931                    xy += *abs_x_ptr.add(best_id);
932                    yy += *y2f_ptr.add(best_id);
933                    *y2f_ptr.add(best_id) += 2.0;
934                    *y_ptr.add(best_id) += 1;
935                    k -= 1;
936                }
937            }
938        }
939        #[cfg(not(target_arch = "aarch64"))]
940        {
941            let mut y2f = [0.0f32; MAX_PVQ_N];
942            let abs_x_ptr = abs_x.as_ptr();
943            let y2f_ptr = y2f.as_mut_ptr();
944            let y_ptr = y.as_mut_ptr();
945            unsafe {
946                while k > 0 {
947                    yy += 1.0;
948                    let rxy0 = xy + *abs_x_ptr;
949                    let mut best_id = 0;
950                    let mut best_num = rxy0 * rxy0;
951                    let mut best_den = yy + *y2f_ptr;
952                    let mut i = 1;
953                    while i + 1 < n {
954                        let rxy1 = xy + *abs_x_ptr.add(i);
955                        let ryy1 = yy + *y2f_ptr.add(i);
956                        let rxy1_sq = rxy1 * rxy1;
957                        if best_den * rxy1_sq > ryy1 * best_num {
958                            best_id = i;
959                            best_num = rxy1_sq;
960                            best_den = ryy1;
961                        }
962                        let rxy2 = xy + *abs_x_ptr.add(i + 1);
963                        let ryy2 = yy + *y2f_ptr.add(i + 1);
964                        let rxy2_sq = rxy2 * rxy2;
965                        if best_den * rxy2_sq > ryy2 * best_num {
966                            best_id = i + 1;
967                            best_num = rxy2_sq;
968                            best_den = ryy2;
969                        }
970                        i += 2;
971                    }
972                    if i < n {
973                        let rxy = xy + *abs_x_ptr.add(i);
974                        let ryy = yy + *y2f_ptr.add(i);
975                        let rxy_sq = rxy * rxy;
976                        if best_den * rxy_sq > ryy * best_num {
977                            best_id = i;
978                        }
979                    }
980                    xy += *abs_x_ptr.add(best_id);
981                    yy += *y2f_ptr.add(best_id);
982                    *y2f_ptr.add(best_id) += 2.0;
983                    *y_ptr.add(best_id) += 1;
984                    k -= 1;
985                }
986            }
987        }
988    } else {
989        let mut y2f_mu = [MaybeUninit::<f32>::uninit(); MAX_PVQ_N];
990
991        let y_ptr = y.as_mut_ptr();
992        for i in 0..n {
993            unsafe {
994                y2f_mu[i].write(2.0 * *y_ptr.add(i) as f32);
995            }
996        }
997        let y2f =
998            unsafe { std::slice::from_raw_parts_mut(y2f_mu.as_mut_ptr() as *mut f32, MAX_PVQ_N) };
999        let mut scores_mu = [MaybeUninit::<(f32, usize)>::uninit(); MAX_PVQ_N];
1000
1001        let abs_x_ptr = abs_x.as_ptr();
1002        let y2f_ptr = y2f.as_mut_ptr();
1003        while k > 0 {
1004            let batch = BATCH_SIZE.min(k);
1005
1006            unsafe {
1007                for i in 0..n {
1008                    let rxy = xy + *abs_x_ptr.add(i);
1009                    let ryy = yy + *y2f_ptr.add(i) + 1.0;
1010                    let score = rxy * rxy / ryy;
1011                    scores_mu[i].write((score, i));
1012                }
1013            }
1014
1015            let scores = unsafe {
1016                std::slice::from_raw_parts_mut(scores_mu.as_mut_ptr() as *mut (f32, usize), n)
1017            };
1018
1019            let pos = batch as usize;
1020
1021            scores.select_nth_unstable_by(pos, |a, b| {
1022                if a.0 > b.0 {
1023                    std::cmp::Ordering::Less
1024                } else if a.0 < b.0 {
1025                    std::cmp::Ordering::Greater
1026                } else {
1027                    std::cmp::Ordering::Equal
1028                }
1029            });
1030
1031            unsafe {
1032                for b in 0..batch as usize {
1033                    let idx = scores[b].1;
1034                    xy += *abs_x_ptr.add(idx);
1035                    yy += *y2f_ptr.add(idx) + 1.0;
1036                    *y2f_ptr.add(idx) += 2.0;
1037                    *y_ptr.add(idx) += 1;
1038                }
1039            }
1040
1041            k -= batch;
1042        }
1043    }
1044
1045    unsafe {
1046        let y_ptr = y.as_mut_ptr();
1047        for i in 0..n {
1048            *y_ptr.add(i) *= *signs.as_ptr().add(i);
1049        }
1050    }
1051
1052    yy
1053}
1054
1055#[cfg(target_arch = "aarch64")]
1056#[inline(always)]
1057#[allow(unsafe_op_in_unsafe_fn)]
1058unsafe fn pvq_search_scalar_init_neon(
1059    x: &[f32],
1060    n: usize,
1061    abs_x: &mut [f32; 32],
1062    sign_x: &mut [i32; 32],
1063) -> f32 {
1064    use std::arch::aarch64::*;
1065
1066    let mut sum_vec = vdupq_n_f32(0.0);
1067    let mut i = 0;
1068
1069    while i + 16 <= n {
1070        let vx0 = vld1q_f32(x.as_ptr().add(i));
1071        let vx1 = vld1q_f32(x.as_ptr().add(i + 4));
1072        let vx2 = vld1q_f32(x.as_ptr().add(i + 8));
1073        let vx3 = vld1q_f32(x.as_ptr().add(i + 12));
1074
1075        let vabs0 = vabsq_f32(vx0);
1076        let vabs1 = vabsq_f32(vx1);
1077        let vabs2 = vabsq_f32(vx2);
1078        let vabs3 = vabsq_f32(vx3);
1079
1080        vst1q_f32(abs_x.as_mut_ptr().add(i), vabs0);
1081        vst1q_f32(abs_x.as_mut_ptr().add(i + 4), vabs1);
1082        vst1q_f32(abs_x.as_mut_ptr().add(i + 8), vabs2);
1083        vst1q_f32(abs_x.as_mut_ptr().add(i + 12), vabs3);
1084
1085        sum_vec = vaddq_f32(sum_vec, vabs0);
1086        sum_vec = vaddq_f32(sum_vec, vabs1);
1087        sum_vec = vaddq_f32(sum_vec, vabs2);
1088        sum_vec = vaddq_f32(sum_vec, vabs3);
1089
1090        for j in 0..16 {
1091            sign_x[i + j] = (x[i + j] < 0.0) as i32;
1092        }
1093
1094        i += 16;
1095    }
1096
1097    while i + 8 <= n {
1098        let vx0 = vld1q_f32(x.as_ptr().add(i));
1099        let vx1 = vld1q_f32(x.as_ptr().add(i + 4));
1100
1101        let vabs0 = vabsq_f32(vx0);
1102        let vabs1 = vabsq_f32(vx1);
1103
1104        vst1q_f32(abs_x.as_mut_ptr().add(i), vabs0);
1105        vst1q_f32(abs_x.as_mut_ptr().add(i + 4), vabs1);
1106
1107        sum_vec = vaddq_f32(sum_vec, vabs0);
1108        sum_vec = vaddq_f32(sum_vec, vabs1);
1109
1110        for j in 0..8 {
1111            sign_x[i + j] = (x[i + j] < 0.0) as i32;
1112        }
1113
1114        i += 8;
1115    }
1116
1117    while i + 4 <= n {
1118        let vx = vld1q_f32(x.as_ptr().add(i));
1119        let vabs = vabsq_f32(vx);
1120        vst1q_f32(abs_x.as_mut_ptr().add(i), vabs);
1121        sum_vec = vaddq_f32(sum_vec, vabs);
1122
1123        for j in 0..4 {
1124            sign_x[i + j] = (x[i + j] < 0.0) as i32;
1125        }
1126
1127        i += 4;
1128    }
1129
1130    let mut sum = vaddvq_f32(sum_vec);
1131
1132    for j in i..n {
1133        let xi = x[j];
1134        let abs_xi = xi.abs();
1135        abs_x[j] = abs_xi;
1136        sum += abs_xi;
1137        sign_x[j] = (xi < 0.0) as i32;
1138    }
1139
1140    sum
1141}
1142
1143#[inline(always)]
1144fn pvq_search_small_k(x: &[f32], y: &mut [i32], k: i32, n: usize) {
1145    debug_assert!(k <= 4 && k > 0);
1146    debug_assert!(n <= 31);
1147
1148    let mut abs_x = [0.0f32; 32];
1149    let mut y2f = [0.0f32; 32];
1150    let mut sign_x = [0i32; 32];
1151
1152    unsafe {
1153        let x_ptr = x.as_ptr();
1154        let abs_x_ptr = abs_x.as_mut_ptr();
1155        let sign_ptr = sign_x.as_mut_ptr();
1156        for i in 0..n {
1157            let xi = *x_ptr.add(i);
1158            *abs_x_ptr.add(i) = xi.abs();
1159            *sign_ptr.add(i) = (xi < 0.0) as i32;
1160        }
1161    }
1162
1163    let mut yy = 0.0f32;
1164    let mut xy = 0.0f32;
1165
1166    let abs_x_ptr = abs_x.as_ptr();
1167    let y2f_ptr = y2f.as_mut_ptr();
1168    let y_ptr = y.as_mut_ptr();
1169    unsafe {
1170        for _ in 0..k {
1171            yy += 1.0;
1172
1173            let rxy0 = xy + *abs_x_ptr;
1174            let mut best_id = 0usize;
1175            let mut best_num = rxy0 * rxy0;
1176            let mut best_den = yy + *y2f_ptr;
1177
1178            let mut i = 1;
1179            while i + 1 < n {
1180                let rxy1 = xy + *abs_x_ptr.add(i);
1181                let rxy2 = xy + *abs_x_ptr.add(i + 1);
1182                let den1 = yy + *y2f_ptr.add(i);
1183                let den2 = yy + *y2f_ptr.add(i + 1);
1184                let rxy1_sq = rxy1 * rxy1;
1185                let rxy2_sq = rxy2 * rxy2;
1186
1187                if best_den * rxy1_sq > den1 * best_num {
1188                    best_id = i;
1189                    best_num = rxy1_sq;
1190                    best_den = den1;
1191                }
1192                if best_den * rxy2_sq > den2 * best_num {
1193                    best_id = i + 1;
1194                    best_num = rxy2_sq;
1195                    best_den = den2;
1196                }
1197                i += 2;
1198            }
1199            if i < n {
1200                let rxy = xy + *abs_x_ptr.add(i);
1201                let rxy_sq = rxy * rxy;
1202                let den = yy + *y2f_ptr.add(i);
1203                if best_den * rxy_sq > den * best_num {
1204                    best_id = i;
1205                }
1206            }
1207
1208            xy += *abs_x_ptr.add(best_id);
1209            yy += *y2f_ptr.add(best_id);
1210            *y2f_ptr.add(best_id) += 2.0;
1211            *y_ptr.add(best_id) += 1;
1212        }
1213    }
1214
1215    unsafe {
1216        let y_ptr = y.as_mut_ptr();
1217        let sign_ptr = sign_x.as_ptr();
1218        for i in 0..n {
1219            let s = *sign_ptr.add(i);
1220            *y_ptr.add(i) = (*y_ptr.add(i) ^ -s) + s;
1221        }
1222    }
1223}
1224#[inline]
1225fn pvq_search_scalar(x: &[f32], y: &mut [i32], k: i32, n: usize) {
1226    debug_assert!(n <= 31);
1227    let mut k = k;
1228    let mut yy = 0.0f32;
1229    let mut xy = 0.0f32;
1230
1231    y[..n].fill(0);
1232
1233    if k <= 0 {
1234        return;
1235    }
1236
1237    if k <= 4 {
1238        pvq_search_small_k(x, y, k, n);
1239        return;
1240    }
1241
1242    let mut abs_x = [0.0f32; 32];
1243    let mut y2f = [0.0f32; 32];
1244    let mut sign_x = [0i32; 32];
1245
1246    #[cfg(target_arch = "aarch64")]
1247    let sum = unsafe { pvq_search_scalar_init_neon(x, n, &mut abs_x, &mut sign_x) };
1248    #[cfg(all(not(target_arch = "aarch64"), target_arch = "x86_64"))]
1249    let sum = unsafe {
1250        if std::arch::is_x86_feature_detected!("avx2") {
1251            pvq_search_scalar_init_avx2(x, n, &mut abs_x, &mut sign_x)
1252        } else {
1253            let mut s = 0.0f32;
1254            for i in 0..n {
1255                let xi = x[i];
1256                let abs_xi = xi.abs();
1257                abs_x[i] = abs_xi;
1258                s += abs_xi;
1259                sign_x[i] = (xi < 0.0) as i32;
1260            }
1261            s
1262        }
1263    };
1264    #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
1265    let sum = {
1266        let mut s = 0.0f32;
1267        for i in 0..n {
1268            let xi = x[i];
1269            let abs_xi = xi.abs();
1270            abs_x[i] = abs_xi;
1271            s += abs_xi;
1272            sign_x[i] = (xi < 0.0) as i32;
1273        }
1274        s
1275    };
1276
1277    if k > (n >> 1) as i32 && sum > 1e-15 {
1278        let rcp = (k as f32 + 0.8) / sum;
1279
1280        let abs_x_ptr = abs_x.as_ptr();
1281        let y2f_ptr = y2f.as_mut_ptr();
1282        let y_ptr = y.as_mut_ptr();
1283        unsafe {
1284            for i in 0..n {
1285                let yi = (*abs_x_ptr.add(i) * rcp) as i32;
1286                *y_ptr.add(i) = yi;
1287                let yf = yi as f32;
1288                yy += yf * yf;
1289                xy += yf * *abs_x_ptr.add(i);
1290                *y2f_ptr.add(i) = 2.0 * yf;
1291                k -= yi;
1292            }
1293
1294            if k > n as i32 + 3 {
1295                let tmp = k as f32;
1296                yy += tmp * tmp;
1297                yy += tmp * *y_ptr as f32;
1298                *y_ptr += k;
1299                *y2f_ptr = 2.0 * *y_ptr as f32;
1300                k = 0;
1301            }
1302        }
1303    }
1304
1305    let abs_x_ptr = abs_x.as_ptr();
1306    let y2f_ptr = y2f.as_mut_ptr();
1307    let y_ptr = y.as_mut_ptr();
1308    unsafe {
1309        while k > 0 {
1310            yy += 1.0;
1311
1312            let rxy0 = xy + *abs_x_ptr;
1313            let mut best_id = 0usize;
1314            let mut best_num = rxy0 * rxy0;
1315            let mut best_den = yy + *y2f_ptr;
1316
1317            let mut i = 1;
1318            while i < n {
1319                let rxy = xy + *abs_x_ptr.add(i);
1320                let ryy = yy + *y2f_ptr.add(i);
1321                let rxy_sq = rxy * rxy;
1322
1323                if best_den * rxy_sq > ryy * best_num {
1324                    best_num = rxy_sq;
1325                    best_den = ryy;
1326                    best_id = i;
1327                }
1328                i += 1;
1329            }
1330
1331            xy += *abs_x_ptr.add(best_id);
1332            yy += *y2f_ptr.add(best_id);
1333            *y2f_ptr.add(best_id) += 2.0;
1334            *y_ptr.add(best_id) += 1;
1335            k -= 1;
1336        }
1337    }
1338
1339    unsafe {
1340        let y_ptr = y.as_mut_ptr();
1341        let sign_ptr = sign_x.as_ptr();
1342        for i in 0..n {
1343            let s = *sign_ptr.add(i);
1344            *y_ptr.add(i) = (*y_ptr.add(i) ^ -s) + s;
1345        }
1346    }
1347}
1348
1349#[cfg(target_arch = "aarch64")]
1350#[inline]
1351fn pvq_search_neon(x: &[f32], y: &mut [i32], k: i32, n: usize) {
1352    use std::arch::aarch64::*;
1353
1354    debug_assert!(n <= 16);
1355    let mut k = k;
1356    let mut yy = 0.0f32;
1357    let mut xy = 0.0f32;
1358
1359    y[..n].fill(0);
1360
1361    if k <= 0 {
1362        return;
1363    }
1364
1365    if k <= 4 {
1366        let mut abs_x_arr = [0.0f32; 16];
1367        let mut y2f_arr = [0.0f32; 16];
1368        let mut sign_x_arr = [0i32; 16];
1369
1370        unsafe {
1371            let vzero = vdupq_n_f32(0.0);
1372            let n4 = n & !3;
1373            for i in (0..n4).step_by(4) {
1374                let vx = vld1q_f32(x.as_ptr().add(i));
1375                let vabs = vabsq_f32(vx);
1376                vst1q_f32(abs_x_arr.as_mut_ptr().add(i), vabs);
1377                let vneg = vcltq_f32(vx, vzero);
1378                let vsign = vandq_u32(vneg, vdupq_n_u32(1));
1379                vst1q_s32(sign_x_arr.as_mut_ptr().add(i), vreinterpretq_s32_u32(vsign));
1380            }
1381            for i in n4..n {
1382                let xi = x[i];
1383                abs_x_arr[i] = xi.abs();
1384                sign_x_arr[i] = (xi < 0.0) as i32;
1385            }
1386        }
1387
1388        let mut yy_local = 0.0f32;
1389        let mut xy_local = 0.0f32;
1390
1391        let abs_x_ptr = abs_x_arr.as_ptr();
1392        let y2f_ptr = y2f_arr.as_mut_ptr();
1393        let y_ptr = y.as_mut_ptr();
1394        unsafe {
1395            for _ in 0..k {
1396                yy_local += 1.0;
1397                let mut best_num = xy_local + *abs_x_ptr;
1398                let mut best_den = yy_local + *y2f_ptr;
1399                let mut best_id = 0;
1400
1401                let mut j = 1;
1402                while j < n {
1403                    let rxy = xy_local + *abs_x_ptr.add(j);
1404                    let ryy = yy_local + *y2f_ptr.add(j);
1405                    if best_den * rxy > best_num * ryy {
1406                        best_den = ryy;
1407                        best_num = rxy;
1408                        best_id = j;
1409                    }
1410                    j += 1;
1411                }
1412
1413                xy_local += *abs_x_ptr.add(best_id);
1414                yy_local += *y2f_ptr.add(best_id);
1415                *y2f_ptr.add(best_id) += 2.0;
1416                *y_ptr.add(best_id) += 1;
1417            }
1418        }
1419
1420        unsafe {
1421            let sign_ptr = sign_x_arr.as_ptr();
1422            for i in 0..n {
1423                let s = *sign_ptr.add(i);
1424                *y_ptr.add(i) = (*y_ptr.add(i) ^ -s) + s;
1425            }
1426        }
1427        return;
1428    }
1429
1430    let mut abs_x_mu = [MaybeUninit::<f32>::uninit(); 20];
1431    let mut y2f_mu = [MaybeUninit::<f32>::uninit(); 20];
1432    let mut sign_x_mu = [MaybeUninit::<i32>::uninit(); 16];
1433    let mut sum;
1434
1435    let n4 = n & !3;
1436    unsafe {
1437        let mut vsum = vdupq_n_f32(0.0);
1438        let vzero = vdupq_n_f32(0.0);
1439        for i in (0..n4).step_by(4) {
1440            let vx = vld1q_f32(x.as_ptr().add(i));
1441            let vabs = vabsq_f32(vx);
1442            vst1q_f32(abs_x_mu.as_mut_ptr().add(i) as *mut f32, vabs);
1443            vsum = vaddq_f32(vsum, vabs);
1444            let vneg = vcltq_f32(vx, vzero);
1445            let vsign = vandq_u32(vneg, vdupq_n_u32(1));
1446            vst1q_s32(
1447                sign_x_mu.as_mut_ptr().add(i) as *mut i32,
1448                vreinterpretq_s32_u32(vsign),
1449            );
1450        }
1451        sum = vaddvq_f32(vsum);
1452    }
1453
1454    for i in n4..n {
1455        let xi = x[i];
1456        let abs_xi = xi.abs();
1457        abs_x_mu[i].write(abs_xi);
1458        sum += abs_xi;
1459        sign_x_mu[i].write((xi < 0.0) as i32);
1460    }
1461
1462    let abs_x = unsafe { std::slice::from_raw_parts_mut(abs_x_mu.as_mut_ptr() as *mut f32, 20) };
1463    let y2f = unsafe { std::slice::from_raw_parts_mut(y2f_mu.as_mut_ptr() as *mut f32, 20) };
1464    let sign_x = unsafe { std::slice::from_raw_parts_mut(sign_x_mu.as_mut_ptr() as *mut i32, 16) };
1465
1466    let ran_presearch = k > (n >> 1) as i32 && sum > 1e-15;
1467    if ran_presearch {
1468        let rcp = (k as f32 + 0.8) / sum;
1469
1470        unsafe {
1471            let vrcp = vdupq_n_f32(rcp);
1472            let mut vyy = vdupq_n_f32(0.0);
1473            let mut vxy = vdupq_n_f32(0.0);
1474            let mut vk_sum = vdupq_n_s32(0);
1475
1476            for i in (0..n4).step_by(4) {
1477                let vabs = vld1q_f32(abs_x.as_ptr().add(i));
1478                let vyi_f = vmulq_f32(vabs, vrcp);
1479                let vyi = vcvtq_s32_f32(vyi_f);
1480                vst1q_s32(y.as_mut_ptr().add(i), vyi);
1481
1482                let vyi_f = vcvtq_f32_s32(vyi);
1483                vyy = vfmaq_f32(vyy, vyi_f, vyi_f);
1484                vxy = vfmaq_f32(vxy, vyi_f, vabs);
1485
1486                let vy2f = vaddq_f32(vyi_f, vyi_f);
1487                vst1q_f32(y2f.as_mut_ptr().add(i), vy2f);
1488
1489                vk_sum = vaddq_s32(vk_sum, vyi);
1490            }
1491
1492            yy = vaddvq_f32(vyy);
1493            xy = vaddvq_f32(vxy);
1494            k -= vaddvq_s32(vk_sum);
1495        }
1496
1497        for i in n4..n {
1498            let yi = (abs_x[i] * rcp) as i32;
1499            y[i] = yi;
1500            let yf = yi as f32;
1501            yy += yf * yf;
1502            xy += yf * abs_x[i];
1503            y2f[i] = 2.0 * yf;
1504            k -= yi;
1505        }
1506
1507        if k > n as i32 + 3 {
1508            let tmp = k as f32;
1509            yy += tmp * tmp + tmp * y[0] as f32;
1510            y[0] += k;
1511            y2f[0] = 2.0 * y[0] as f32;
1512            k = 0;
1513        }
1514    } else {
1515        for i in 0..n {
1516            y2f[i] = 0.0;
1517        }
1518    }
1519
1520    unsafe {
1521        let abs_x_ptr = abs_x.as_ptr();
1522        let y2f_ptr = y2f.as_mut_ptr();
1523        let y_ptr = y.as_mut_ptr();
1524        let n4 = n & !3;
1525
1526        for _ in 0..k {
1527            yy += 1.0;
1528
1529            let vxy = vdupq_n_f32(xy);
1530            let vyy = vdupq_n_f32(yy);
1531            let mut vmax = vdupq_n_f32(0.0);
1532            let mut best_id: usize = 0;
1533
1534            let mut j = 0;
1535            while j < n4 {
1536                let vx = vld1q_f32(abs_x_ptr.add(j));
1537                let vy = vld1q_f32(y2f_ptr.add(j));
1538                let rxy = vaddq_f32(vx, vxy);
1539                let ryy = vaddq_f32(vy, vyy);
1540                let inv_sqrt = vrsqrteq_f32(ryy);
1541                let score = vmulq_f32(rxy, inv_sqrt);
1542                vmax = vmaxq_f32(vmax, score);
1543                let sc = std::slice::from_raw_parts(&score as *const float32x4_t as *const f32, 4);
1544                let mx = vmaxvq_f32(vmax);
1545                for lane in 0..4 {
1546                    if sc[lane] == mx {
1547                        best_id = j + lane;
1548                    }
1549                }
1550                j += 4;
1551            }
1552
1553            while j < n {
1554                let rxy = xy + *abs_x_ptr.add(j);
1555                let ryy = yy + *y2f_ptr.add(j);
1556                let score = rxy * (1.0 / ryy.sqrt());
1557                let current_max = vmaxvq_f32(vmax);
1558                if score > current_max {
1559                    best_id = j;
1560                    vmax = vsetq_lane_f32(score, vmax, 0);
1561                }
1562                j += 1;
1563            }
1564
1565            xy += *abs_x_ptr.add(best_id);
1566            yy += *y2f_ptr.add(best_id);
1567            *y2f_ptr.add(best_id) += 2.0;
1568            *y_ptr.add(best_id) += 1;
1569        }
1570    }
1571
1572    unsafe {
1573        let y_ptr = y.as_mut_ptr();
1574        let sign_ptr = sign_x.as_ptr();
1575        for i in 0..n {
1576            let s = *sign_ptr.add(i);
1577            *y_ptr.add(i) = (*y_ptr.add(i) ^ -s) + s;
1578        }
1579    }
1580}
1581
1582#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
1583#[target_feature(enable = "avx2,fma")]
1584#[allow(unsafe_op_in_unsafe_fn)]
1585unsafe fn pvq_search_avx2(x: &[f32], y: &mut [i32], k: i32, n: usize) {
1586    use std::arch::x86_64::*;
1587
1588    debug_assert!(n <= 31);
1589    debug_assert!(k > 4);
1590
1591    let mut k = k;
1592    let mut yy = 0.0f32;
1593    let mut xy = 0.0f32;
1594
1595    y[..n].fill(0);
1596
1597    let mut abs_x = [0.0f32; 32];
1598    let mut y2f = [0.0f32; 32];
1599    let mut sign_x = [0i32; 32];
1600
1601    let sign_mask = _mm256_set1_ps(-0.0f32);
1602    let vzero_ps = _mm256_setzero_ps();
1603    let vone_i = _mm256_set1_epi32(1);
1604    let mut acc = _mm256_setzero_ps();
1605    let mut i = 0;
1606    while i + 8 <= n {
1607        let v = _mm256_loadu_ps(x.as_ptr().add(i));
1608        let a = _mm256_andnot_ps(sign_mask, v);
1609        _mm256_storeu_ps(abs_x.as_mut_ptr().add(i), a);
1610        acc = _mm256_add_ps(acc, a);
1611
1612        let neg_mask = _mm256_cmp_ps(v, vzero_ps, _CMP_LT_OS);
1613        let sign_i = _mm256_and_si256(_mm256_castps_si256(neg_mask), vone_i);
1614        _mm256_storeu_si256(sign_x.as_mut_ptr().add(i) as *mut __m256i, sign_i);
1615        i += 8;
1616    }
1617
1618    let lo4 = _mm256_castps256_ps128(acc);
1619    let hi4 = _mm256_extractf128_ps(acc, 1);
1620    let s4 = _mm_add_ps(lo4, hi4);
1621    let s2 = _mm_add_ps(s4, _mm_movehl_ps(s4, s4));
1622    let s1 = _mm_add_ss(s2, _mm_shuffle_ps(s2, s2, 1));
1623    let mut sum = _mm_cvtss_f32(s1);
1624    for j in i..n {
1625        let xi = x[j];
1626        let a = xi.abs();
1627        abs_x[j] = a;
1628        sum += a;
1629        sign_x[j] = (xi < 0.0) as i32;
1630    }
1631
1632    if k > (n >> 1) as i32 && sum > 1e-15 {
1633        let rcp = (k as f32 + 0.8) / sum;
1634        let vrcp = _mm256_set1_ps(rcp);
1635        let mut vyy_acc = _mm256_setzero_ps();
1636        let mut vxy_acc = _mm256_setzero_ps();
1637        let mut vk_acc = _mm256_setzero_ps();
1638        let mut i = 0;
1639        while i + 8 <= n {
1640            let vabs = _mm256_loadu_ps(abs_x.as_ptr().add(i));
1641            let vyi_f32 = _mm256_mul_ps(vabs, vrcp);
1642
1643            let vyi_i = _mm256_cvttps_epi32(vyi_f32);
1644            _mm256_storeu_si256(y.as_mut_ptr().add(i) as *mut __m256i, vyi_i);
1645            let vyi_f = _mm256_cvtepi32_ps(vyi_i);
1646
1647            vyy_acc = _mm256_fmadd_ps(vyi_f, vyi_f, vyy_acc);
1648            vxy_acc = _mm256_fmadd_ps(vyi_f, vabs, vxy_acc);
1649            vk_acc = _mm256_add_ps(vk_acc, vyi_f);
1650
1651            let vy2f = _mm256_add_ps(vyi_f, vyi_f);
1652            _mm256_storeu_ps(y2f.as_mut_ptr().add(i), vy2f);
1653            i += 8;
1654        }
1655
1656        let hsum = |v: __m256| -> f32 {
1657            let lo = _mm256_castps256_ps128(v);
1658            let hi = _mm256_extractf128_ps(v, 1);
1659            let s4 = _mm_add_ps(lo, hi);
1660            let s2 = _mm_add_ps(s4, _mm_movehl_ps(s4, s4));
1661            let s1 = _mm_add_ss(s2, _mm_shuffle_ps(s2, s2, 1));
1662            _mm_cvtss_f32(s1)
1663        };
1664        yy += hsum(vyy_acc);
1665        xy += hsum(vxy_acc);
1666        k -= hsum(vk_acc) as i32;
1667
1668        while i < n {
1669            let yi = (abs_x[i] * rcp) as i32;
1670            y[i] = yi;
1671            let yf = yi as f32;
1672            yy += yf * yf;
1673            xy += yf * abs_x[i];
1674            y2f[i] = 2.0 * yf;
1675            k -= yi;
1676            i += 1;
1677        }
1678        if k > n as i32 + 3 {
1679            let tmp = k as f32;
1680            yy += tmp * tmp + tmp * y[0] as f32;
1681            y[0] += k;
1682            y2f[0] = 2.0 * y[0] as f32;
1683            k = 0;
1684        }
1685    }
1686
1687    let abs_x_ptr = abs_x.as_ptr();
1688    let y2f_ptr = y2f.as_mut_ptr();
1689    let y_ptr = y.as_mut_ptr();
1690    let n8 = n & !7;
1691    let n_ceil8 = (n + 7) & !7;
1692    let mut scores = [0.0f32; 32];
1693
1694    while k > 0 {
1695        yy += 1.0;
1696        let vxy = _mm256_set1_ps(xy);
1697        let vyy = _mm256_set1_ps(yy);
1698
1699        let mut vmax = _mm256_setzero_ps();
1700        let mut j = 0;
1701        while j < n8 {
1702            let vabs = _mm256_loadu_ps(abs_x_ptr.add(j));
1703            let vy2f = _mm256_loadu_ps(y2f_ptr.add(j));
1704            let rxy = _mm256_add_ps(vabs, vxy);
1705            let ryy = _mm256_add_ps(vy2f, vyy);
1706            let score = _mm256_mul_ps(rxy, _mm256_rsqrt_ps(ryy));
1707            _mm256_storeu_ps(scores.as_mut_ptr().add(j), score);
1708            vmax = _mm256_max_ps(vmax, score);
1709            j += 8;
1710        }
1711
1712        while j < n {
1713            let rxy = xy + *abs_x_ptr.add(j);
1714            let ryy = yy + *y2f_ptr.add(j);
1715            scores[j] = rxy * (1.0 / ryy.sqrt());
1716            j += 1;
1717        }
1718
1719        let global_max = {
1720            let hi = _mm256_extractf128_ps(vmax, 1);
1721            let lo = _mm256_castps256_ps128(vmax);
1722            let m4 = _mm_max_ps(lo, hi);
1723            let m2 = _mm_max_ps(m4, _mm_movehl_ps(m4, m4));
1724            let m1 = _mm_max_ss(m2, _mm_shuffle_ps(m2, m2, 1));
1725            _mm_cvtss_f32(m1)
1726        };
1727
1728        let mut gmax = global_max;
1729        for j in n8..n {
1730            if scores[j] > gmax {
1731                gmax = scores[j];
1732            }
1733        }
1734
1735        let vgmax = _mm256_set1_ps(gmax);
1736        let mut best_id: usize = 0;
1737        let mut j = 0;
1738        while j < n_ceil8 {
1739            let vs = _mm256_loadu_ps(scores.as_ptr().add(j));
1740            let mask = _mm256_movemask_ps(_mm256_cmp_ps(vs, vgmax, _CMP_EQ_OQ)) as u32;
1741            if mask != 0 {
1742                best_id = j + mask.trailing_zeros() as usize;
1743                break;
1744            }
1745            j += 8;
1746        }
1747
1748        xy += *abs_x_ptr.add(best_id);
1749        yy += *y2f_ptr.add(best_id);
1750        *y2f_ptr.add(best_id) += 2.0;
1751        *y_ptr.add(best_id) += 1;
1752        k -= 1;
1753    }
1754
1755    for i in 0..n {
1756        let s = sign_x[i];
1757        y[i] = (y[i] ^ -s) + s;
1758    }
1759}
1760
1761#[inline]
1762fn exp_rotation1(x: &mut [f32], len: usize, stride: usize, c: f32, s: f32) {
1763    #[cfg(target_arch = "aarch64")]
1764    unsafe {
1765        exp_rotation1_neon(x, len, stride, c, s);
1766    }
1767    #[cfg(not(target_arch = "aarch64"))]
1768    {
1769        exp_rotation1_scalar(x, len, stride, c, s);
1770    }
1771}
1772
1773#[inline]
1774fn exp_rotation1_scalar(x: &mut [f32], len: usize, stride: usize, c: f32, s: f32) {
1775    let ms = -s;
1776    for i in 0..(len - stride) {
1777        let x1 = x[i];
1778        let x2 = x[i + stride];
1779        x[i + stride] = c * x2 + s * x1;
1780        x[i] = c * x1 + ms * x2;
1781    }
1782    if len >= 2 * stride {
1783        for i in (0..(len - 2 * stride)).rev() {
1784            let x1 = x[i];
1785            let x2 = x[i + stride];
1786            x[i + stride] = c * x2 + s * x1;
1787            x[i] = c * x1 + ms * x2;
1788        }
1789    }
1790}
1791
1792#[cfg(target_arch = "aarch64")]
1793#[inline(always)]
1794#[allow(unsafe_op_in_unsafe_fn)]
1795unsafe fn exp_rotation1_neon(x: &mut [f32], len: usize, stride: usize, c: f32, s: f32) {
1796    if stride < 4 {
1797        exp_rotation1_scalar(x, len, stride, c, s);
1798        return;
1799    }
1800
1801    use std::arch::aarch64::*;
1802
1803    let vc = vdupq_n_f32(c);
1804    let vs = vdupq_n_f32(s);
1805
1806    // Forward pass: SIMD when we can load 4 contiguous elements
1807    let mut i = 0;
1808    while i + 4 <= len - stride {
1809        let vx1 = vld1q_f32(x.as_ptr().add(i));
1810        let vx2 = vld1q_f32(x.as_ptr().add(i + stride));
1811
1812        // y1 = c*x1 - s*x2
1813        let vy1 = vfmsq_f32(vmulq_f32(vx1, vc), vs, vx2);
1814        // y2 = c*x2 + s*x1
1815        let vy2 = vfmaq_f32(vmulq_f32(vx2, vc), vs, vx1);
1816
1817        vst1q_f32(x.as_mut_ptr().add(i), vy1);
1818        vst1q_f32(x.as_mut_ptr().add(i + stride), vy2);
1819
1820        i += 4;
1821    }
1822    for j in i..(len - stride) {
1823        let x1 = x[j];
1824        let x2 = x[j + stride];
1825        x[j + stride] = c * x2 + s * x1;
1826        x[j] = c * x1 - s * x2;
1827    }
1828
1829    if len >= 2 * stride {
1830        for j in (0..(len - 2 * stride)).rev() {
1831            let x1 = x[j];
1832            let x2 = x[j + stride];
1833            x[j + stride] = c * x2 + s * x1;
1834            x[j] = c * x1 - s * x2;
1835        }
1836    }
1837}
1838
1839#[inline(always)]
1840pub fn exp_rotation(x: &mut [f32], length: usize, dir: i32, stride: usize, k: i32, spread: i32) {
1841    const SPREAD_FACTOR: [i32; 3] = [15, 10, 5];
1842    if 2 * k >= length as i32 || spread <= 0 || spread > 3 {
1843        return;
1844    }
1845    let factor = SPREAD_FACTOR[spread as usize - 1];
1846    let gain = (length as f32) / (length as f32 + factor as f32 * k as f32);
1847    let theta = 0.5 * gain * gain;
1848    let c = (0.5 * std::f32::consts::PI * theta).cos();
1849    let s = (0.5 * std::f32::consts::PI * theta).sin();
1850
1851    let mut stride2 = 0;
1852    if length >= 8 * stride {
1853        stride2 = 1;
1854        while (stride2 * stride2 + stride2) * stride + (stride >> 2) < length {
1855            stride2 += 1;
1856        }
1857    }
1858
1859    let block_len = length / stride;
1860    for i in 0..stride {
1861        let x_offset = i * block_len;
1862        let x_subset = &mut x[x_offset..x_offset + block_len];
1863        if dir < 0 {
1864            if stride2 != 0 {
1865                exp_rotation1(x_subset, block_len, stride2, s, c);
1866            }
1867            exp_rotation1(x_subset, block_len, 1, c, s);
1868        } else {
1869            exp_rotation1(x_subset, block_len, 1, c, -s);
1870            if stride2 != 0 {
1871                exp_rotation1(x_subset, block_len, stride2, s, -c);
1872            }
1873        }
1874    }
1875}
1876
1877#[cfg(target_arch = "aarch64")]
1878#[inline(always)]
1879#[allow(unsafe_op_in_unsafe_fn)]
1880unsafe fn extract_collapse_mask_neon(iy: &[i32], n: usize, b: usize) -> u32 {
1881    use std::arch::aarch64::*;
1882
1883    if b <= 1 {
1884        return 1;
1885    }
1886    let n0 = n / b;
1887    let mut collapse_mask = 0u32;
1888
1889    for i in 0..b {
1890        let base = i * n0;
1891        let slice = &iy[base..base + n0];
1892
1893        let mut any_nonzero = false;
1894        let n4 = n0 & !3;
1895        let mut j = 0;
1896
1897        while j < n4 {
1898            let v = vld1q_s32(slice.as_ptr().add(j));
1899
1900            let or_val = vorrq_s32(v, vextq_s32(v, v, 2));
1901            let or_val = vorrq_s32(or_val, vextq_s32(or_val, or_val, 1));
1902            if vgetq_lane_s32(or_val, 0) != 0 {
1903                any_nonzero = true;
1904                break;
1905            }
1906            j += 4;
1907        }
1908
1909        if !any_nonzero {
1910            for j in j..n0 {
1911                if slice[j] != 0 {
1912                    any_nonzero = true;
1913                    break;
1914                }
1915            }
1916        }
1917
1918        if any_nonzero {
1919            collapse_mask |= 1 << i;
1920        }
1921    }
1922    collapse_mask
1923}
1924
1925#[inline(always)]
1926pub fn extract_collapse_mask(iy: &[i32], n: usize, b: usize) -> u32 {
1927    if b <= 1 {
1928        return 1;
1929    }
1930
1931    #[cfg(target_arch = "aarch64")]
1932    unsafe {
1933        extract_collapse_mask_neon(iy, n, b)
1934    }
1935    #[cfg(not(target_arch = "aarch64"))]
1936    {
1937        let n0 = n / b;
1938        let mut collapse_mask = 0u32;
1939        for i in 0..b {
1940            let mut tmp = 0i32;
1941            let base = i * n0;
1942            for j in 0..n0 {
1943                tmp |= iy[base + j];
1944            }
1945            if tmp != 0 {
1946                collapse_mask |= 1 << i;
1947            }
1948        }
1949        collapse_mask
1950    }
1951}
1952
1953#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
1954#[target_feature(enable = "avx2,fma")]
1955#[allow(unsafe_op_in_unsafe_fn)]
1956unsafe fn renormalise_vector_avx2(x: &mut [f32], n: usize, gain: f32) {
1957    use std::arch::x86_64::*;
1958
1959    let mut acc0 = _mm256_setzero_ps();
1960    let mut acc1 = _mm256_setzero_ps();
1961    let mut i = 0;
1962
1963    while i + 16 <= n {
1964        let v0 = _mm256_loadu_ps(x.as_ptr().add(i));
1965        let v1 = _mm256_loadu_ps(x.as_ptr().add(i + 8));
1966        acc0 = _mm256_fmadd_ps(v0, v0, acc0);
1967        acc1 = _mm256_fmadd_ps(v1, v1, acc1);
1968        i += 16;
1969    }
1970    while i + 8 <= n {
1971        let v0 = _mm256_loadu_ps(x.as_ptr().add(i));
1972        acc0 = _mm256_fmadd_ps(v0, v0, acc0);
1973        i += 8;
1974    }
1975    let acc = _mm256_add_ps(acc0, acc1);
1976
1977    let lo = _mm256_castps256_ps128(acc);
1978    let hi = _mm256_extractf128_ps(acc, 1);
1979    let s4 = _mm_add_ps(lo, hi);
1980    let s2 = _mm_add_ps(s4, _mm_movehl_ps(s4, s4));
1981    let s1 = _mm_add_ss(s2, _mm_shuffle_ps(s2, s2, 1));
1982    let mut e = 1e-15f32 + _mm_cvtss_f32(s1);
1983    for j in i..n {
1984        e += x[j] * x[j];
1985    }
1986
1987    let g = gain * (1.0 / e.sqrt());
1988    let vnorm = _mm256_set1_ps(g);
1989    i = 0;
1990    while i + 16 <= n {
1991        let v0 = _mm256_loadu_ps(x.as_ptr().add(i));
1992        let v1 = _mm256_loadu_ps(x.as_ptr().add(i + 8));
1993        _mm256_storeu_ps(x.as_mut_ptr().add(i), _mm256_mul_ps(v0, vnorm));
1994        _mm256_storeu_ps(x.as_mut_ptr().add(i + 8), _mm256_mul_ps(v1, vnorm));
1995        i += 16;
1996    }
1997    while i + 8 <= n {
1998        let v0 = _mm256_loadu_ps(x.as_ptr().add(i));
1999        _mm256_storeu_ps(x.as_mut_ptr().add(i), _mm256_mul_ps(v0, vnorm));
2000        i += 8;
2001    }
2002    for j in i..n {
2003        x[j] *= g;
2004    }
2005}
2006
2007#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
2008#[target_feature(enable = "avx2,fma")]
2009#[allow(unsafe_op_in_unsafe_fn)]
2010unsafe fn alg_quant_resynth_avx2(y: &[i32], x: &mut [f32], n: usize, gain: f32) {
2011    use std::arch::x86_64::*;
2012
2013    let mut acc0 = _mm256_setzero_ps();
2014    let mut i = 0;
2015
2016    while i + 8 <= n {
2017        let yi = _mm256_loadu_si256(y.as_ptr().add(i) as *const __m256i);
2018        let yf = _mm256_cvtepi32_ps(yi);
2019        _mm256_storeu_ps(x.as_mut_ptr().add(i), yf);
2020        acc0 = _mm256_fmadd_ps(yf, yf, acc0);
2021        i += 8;
2022    }
2023
2024    let lo = _mm256_castps256_ps128(acc0);
2025    let hi = _mm256_extractf128_ps(acc0, 1);
2026    let s4 = _mm_add_ps(lo, hi);
2027    let s2 = _mm_add_ps(s4, _mm_movehl_ps(s4, s4));
2028    let s1 = _mm_add_ss(s2, _mm_shuffle_ps(s2, s2, 1));
2029    let mut ryy = _mm_cvtss_f32(s1);
2030
2031    for j in i..n {
2032        let v = y[j] as f32;
2033        x[j] = v;
2034        ryy += v * v;
2035    }
2036
2037    let g = gain / (1e-15f32 + ryy).sqrt();
2038    let vg = _mm256_set1_ps(g);
2039
2040    i = 0;
2041    while i + 8 <= n {
2042        let v = _mm256_loadu_ps(x.as_ptr().add(i));
2043        _mm256_storeu_ps(x.as_mut_ptr().add(i), _mm256_mul_ps(v, vg));
2044        i += 8;
2045    }
2046    for j in i..n {
2047        x[j] *= g;
2048    }
2049}
2050
2051#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
2052#[target_feature(enable = "avx2")]
2053#[allow(unsafe_op_in_unsafe_fn)]
2054unsafe fn pvq_search_scalar_init_avx2(
2055    x: &[f32],
2056    n: usize,
2057    abs_x: &mut [f32; 32],
2058    sign_x: &mut [i32; 32],
2059) -> f32 {
2060    use std::arch::x86_64::*;
2061    let sign_mask = _mm256_set1_ps(-0.0f32);
2062    let mut acc = _mm256_setzero_ps();
2063    let mut i = 0;
2064
2065    while i + 8 <= n {
2066        let v = _mm256_loadu_ps(x.as_ptr().add(i));
2067        let a = _mm256_andnot_ps(sign_mask, v);
2068        _mm256_storeu_ps(abs_x.as_mut_ptr().add(i), a);
2069        acc = _mm256_add_ps(acc, a);
2070
2071        for j in 0..8 {
2072            sign_x[i + j] = (x[i + j] < 0.0) as i32;
2073        }
2074        i += 8;
2075    }
2076
2077    let lo = _mm256_castps256_ps128(acc);
2078    let hi = _mm256_extractf128_ps(acc, 1);
2079    let s4 = _mm_add_ps(lo, hi);
2080    let s2 = _mm_add_ps(s4, _mm_movehl_ps(s4, s4));
2081    let s1 = _mm_add_ss(s2, _mm_shuffle_ps(s2, s2, 1));
2082    let mut sum = _mm_cvtss_f32(s1);
2083
2084    for j in i..n {
2085        let abs_xi = x[j].abs();
2086        abs_x[j] = abs_xi;
2087        sum += abs_xi;
2088        sign_x[j] = (x[j] < 0.0) as i32;
2089    }
2090    sum
2091}
2092
2093#[cfg(target_arch = "aarch64")]
2094#[inline(always)]
2095#[allow(unsafe_op_in_unsafe_fn)]
2096unsafe fn renormalise_vector_neon(x: &mut [f32], n: usize, gain: f32) {
2097    use std::arch::aarch64::*;
2098
2099    let mut sum_vec = vdupq_n_f32(0.0);
2100    let mut i = 0;
2101
2102    while i + 16 <= n {
2103        let x0 = vld1q_f32(x.as_ptr().add(i));
2104        let x1 = vld1q_f32(x.as_ptr().add(i + 4));
2105        let x2 = vld1q_f32(x.as_ptr().add(i + 8));
2106        let x3 = vld1q_f32(x.as_ptr().add(i + 12));
2107        sum_vec = vfmaq_f32(sum_vec, x0, x0);
2108        sum_vec = vfmaq_f32(sum_vec, x1, x1);
2109        sum_vec = vfmaq_f32(sum_vec, x2, x2);
2110        sum_vec = vfmaq_f32(sum_vec, x3, x3);
2111        i += 16;
2112    }
2113
2114    while i + 8 <= n {
2115        let x0 = vld1q_f32(x.as_ptr().add(i));
2116        let x1 = vld1q_f32(x.as_ptr().add(i + 4));
2117        sum_vec = vfmaq_f32(sum_vec, x0, x0);
2118        sum_vec = vfmaq_f32(sum_vec, x1, x1);
2119        i += 8;
2120    }
2121
2122    while i + 4 <= n {
2123        let x0 = vld1q_f32(x.as_ptr().add(i));
2124        sum_vec = vfmaq_f32(sum_vec, x0, x0);
2125        i += 4;
2126    }
2127
2128    let mut e = 1e-15f32 + vaddvq_f32(sum_vec);
2129    for j in i..n {
2130        e += x[j] * x[j];
2131    }
2132
2133    let g = gain * (1.0 / e.sqrt());
2134    let vg = vdupq_n_f32(g);
2135
2136    i = 0;
2137    while i + 16 <= n {
2138        let x0 = vld1q_f32(x.as_ptr().add(i));
2139        let x1 = vld1q_f32(x.as_ptr().add(i + 4));
2140        let x2 = vld1q_f32(x.as_ptr().add(i + 8));
2141        let x3 = vld1q_f32(x.as_ptr().add(i + 12));
2142        vst1q_f32(x.as_mut_ptr().add(i), vmulq_f32(x0, vg));
2143        vst1q_f32(x.as_mut_ptr().add(i + 4), vmulq_f32(x1, vg));
2144        vst1q_f32(x.as_mut_ptr().add(i + 8), vmulq_f32(x2, vg));
2145        vst1q_f32(x.as_mut_ptr().add(i + 12), vmulq_f32(x3, vg));
2146        i += 16;
2147    }
2148
2149    while i + 8 <= n {
2150        let x0 = vld1q_f32(x.as_ptr().add(i));
2151        let x1 = vld1q_f32(x.as_ptr().add(i + 4));
2152        vst1q_f32(x.as_mut_ptr().add(i), vmulq_f32(x0, vg));
2153        vst1q_f32(x.as_mut_ptr().add(i + 4), vmulq_f32(x1, vg));
2154        i += 8;
2155    }
2156
2157    while i + 4 <= n {
2158        let x0 = vld1q_f32(x.as_ptr().add(i));
2159        vst1q_f32(x.as_mut_ptr().add(i), vmulq_f32(x0, vg));
2160        i += 4;
2161    }
2162
2163    for j in i..n {
2164        x[j] *= g;
2165    }
2166}
2167
2168pub fn renormalise_vector(x: &mut [f32], n: usize, gain: f32) {
2169    #[cfg(target_arch = "aarch64")]
2170    unsafe {
2171        renormalise_vector_neon(x, n, gain);
2172    }
2173    #[cfg(target_arch = "x86_64")]
2174    unsafe {
2175        if n >= 8 && std::arch::is_x86_feature_detected!("avx2") {
2176            renormalise_vector_avx2(x, n, gain);
2177            return;
2178        }
2179    }
2180    #[cfg(not(target_arch = "aarch64"))]
2181    {
2182        let mut e = 1e-15f32;
2183        for i in 0..n {
2184            e += x[i] * x[i];
2185        }
2186        let g = gain * (1.0 / e.sqrt());
2187        for i in 0..n {
2188            x[i] *= g;
2189        }
2190    }
2191}
2192
2193#[cfg(target_arch = "aarch64")]
2194#[inline(always)]
2195#[allow(unsafe_op_in_unsafe_fn)]
2196unsafe fn alg_quant_resynth_neon(y: &[i32], x: &mut [f32], n: usize, gain: f32) {
2197    use std::arch::aarch64::*;
2198
2199    let mut sum_vec = vdupq_n_f32(0.0);
2200    let n8 = n & !7;
2201    let mut i = 0;
2202
2203    while i < n8 {
2204        let yi0 = vld1q_s32(y.as_ptr().add(i));
2205        let yi1 = vld1q_s32(y.as_ptr().add(i + 4));
2206
2207        let yf0 = vcvtq_f32_s32(yi0);
2208        let yf1 = vcvtq_f32_s32(yi1);
2209
2210        vst1q_f32(x.as_mut_ptr().add(i), yf0);
2211        vst1q_f32(x.as_mut_ptr().add(i + 4), yf1);
2212
2213        sum_vec = vfmaq_f32(sum_vec, yf0, yf0);
2214        sum_vec = vfmaq_f32(sum_vec, yf1, yf1);
2215
2216        i += 8;
2217    }
2218
2219    let mut ryy = vaddvq_f32(sum_vec);
2220    for j in i..n {
2221        let v = y[j] as f32;
2222        x[j] = v;
2223        ryy += v * v;
2224    }
2225
2226    let g = gain / (1e-15 + ryy).sqrt();
2227    let vg = vdupq_n_f32(g);
2228
2229    i = 0;
2230    while i < n8 {
2231        let vx0 = vld1q_f32(x.as_ptr().add(i));
2232        let vx1 = vld1q_f32(x.as_ptr().add(i + 4));
2233        let vr0 = vmulq_f32(vx0, vg);
2234        let vr1 = vmulq_f32(vx1, vg);
2235        vst1q_f32(x.as_mut_ptr().add(i), vr0);
2236        vst1q_f32(x.as_mut_ptr().add(i + 4), vr1);
2237        i += 8;
2238    }
2239
2240    for j in i..n {
2241        x[j] *= g;
2242    }
2243}
2244
2245#[cfg(not(target_arch = "aarch64"))]
2246#[inline(always)]
2247fn alg_quant_resynth_scalar(y: &[i32], x: &mut [f32], n: usize, gain: f32) {
2248    #[cfg(target_arch = "x86_64")]
2249    unsafe {
2250        if std::arch::is_x86_feature_detected!("avx2") {
2251            alg_quant_resynth_avx2(y, x, n, gain);
2252            return;
2253        }
2254    }
2255    let mut ryy = 0.0f32;
2256    for i in 0..n {
2257        let v = y[i] as f32;
2258        x[i] = v;
2259        ryy += v * v;
2260    }
2261    let g = gain / (1e-15 + ryy).sqrt();
2262    for i in 0..n {
2263        x[i] *= g;
2264    }
2265}
2266
2267fn ec_enc_refine(rc: &mut RangeCoder, refine: i32, up: i32, extra_bits: i32) {
2268    let half_up = up / 2;
2269    let large = refine.abs() > half_up;
2270
2271    rc.encode_bit_logp(large, 1);
2272
2273    if large {
2274        rc.enc_bits((refine < 0) as u32, 1);
2275        rc.enc_bits((refine.abs() - half_up - 1) as u32, (extra_bits - 1) as u32);
2276    } else {
2277        rc.enc_bits((refine + half_up) as u32, extra_bits as u32);
2278    }
2279}
2280
2281#[inline]
2282pub fn alg_quant(
2283    x: &mut [f32],
2284    n: usize,
2285    k: i32,
2286    spread: i32,
2287    stride: usize,
2288    rc: &mut RangeCoder,
2289    gain: f32,
2290    resynth: bool,
2291) -> u32 {
2292    if n <= 32 {
2293        let mut y_buf = [MaybeUninit::<i32>::uninit(); 32];
2294        let y = unsafe { std::slice::from_raw_parts_mut(y_buf.as_mut_ptr() as *mut i32, n) };
2295
2296        exp_rotation(x, n, 1, stride, k, spread);
2297        pvq_search(x, y, k, n);
2298        let mask = extract_collapse_mask(y, n, stride);
2299
2300        encode_pulses(y, n as u32, k as u32, rc);
2301
2302        if resynth {
2303            #[cfg(target_arch = "aarch64")]
2304            unsafe {
2305                alg_quant_resynth_neon(y, x, n, gain);
2306            }
2307            #[cfg(not(target_arch = "aarch64"))]
2308            alg_quant_resynth_scalar(y, x, n, gain);
2309            exp_rotation(x, n, -1, stride, k, spread);
2310        }
2311        mask
2312    } else {
2313        let mut y_mu = [MaybeUninit::<i32>::uninit(); MAX_PVQ_N];
2314        let y = unsafe { std::slice::from_raw_parts_mut(y_mu.as_mut_ptr() as *mut i32, MAX_PVQ_N) };
2315
2316        exp_rotation(x, n, 1, stride, k, spread);
2317        pvq_search(x, &mut y[..n], k, n);
2318        let mask = extract_collapse_mask(&y[..n], n, stride);
2319        encode_pulses(&y[..n], n as u32, k as u32, rc);
2320
2321        if resynth {
2322            #[cfg(target_arch = "aarch64")]
2323            unsafe {
2324                alg_quant_resynth_neon(y, x, n, gain);
2325            }
2326            #[cfg(not(target_arch = "aarch64"))]
2327            alg_quant_resynth_scalar(y, x, n, gain);
2328            exp_rotation(x, n, -1, stride, k, spread);
2329        }
2330        mask
2331    }
2332}
2333
2334pub fn alg_quant_qext(
2335    x: &mut [f32],
2336    n: usize,
2337    k: i32,
2338    spread: i32,
2339    stride: usize,
2340    rc: &mut RangeCoder,
2341    gain: f32,
2342    resynth: bool,
2343    extra_bits: Option<i32>,
2344) -> u32 {
2345    if n <= 32 {
2346        let mut y_buf = [0i32; 32];
2347        let y = &mut y_buf[..n];
2348
2349        exp_rotation(x, n, 1, stride, k, spread);
2350
2351        let use_qext = extra_bits.is_some_and(|eb| eb >= 2);
2352
2353        if use_qext && n == 2 {
2354            let eb = extra_bits.unwrap();
2355            pvq_search_n2(x, y, k);
2356            let mask = extract_collapse_mask(y, n, stride);
2357            encode_pulses(y, n as u32, k as u32, rc);
2358
2359            let up = (1 << eb) - 1;
2360            let abs_x0 = x[0].abs();
2361            let abs_x1 = x[1].abs();
2362            let sum = abs_x0 + abs_x1;
2363            if sum >= 1e-15 {
2364                let rcp_sum = 1.0 / sum;
2365                let ideal_y0 = k as f32 * abs_x0 * rcp_sum;
2366                let actual_y0 = y[0].abs() as f32;
2367                let refine = ((ideal_y0 - actual_y0) * up as f32).round() as i32;
2368                ec_enc_refine(rc, refine, up, eb);
2369            }
2370
2371            if resynth {
2372                #[cfg(target_arch = "aarch64")]
2373                unsafe {
2374                    alg_quant_resynth_neon(y, x, n, gain);
2375                }
2376                #[cfg(not(target_arch = "aarch64"))]
2377                alg_quant_resynth_scalar(y, x, n, gain);
2378                exp_rotation(x, n, -1, stride, k, spread);
2379            }
2380            return mask;
2381        }
2382
2383        if use_qext && n > 2 && n <= 32 {
2384            let eb = extra_bits.unwrap();
2385            let mut up_y = [0i32; 32];
2386            let mut refine = [0i32; 32];
2387            let _yy = pvq_search_qext(x, y, &mut up_y, &mut refine, k, eb, n);
2388            let mask = extract_collapse_mask(&up_y, n, stride);
2389            encode_pulses(y, n as u32, k as u32, rc);
2390
2391            let up = (1 << eb) - 1;
2392            for i in 0..n - 1 {
2393                ec_enc_refine(rc, refine[i], up, eb);
2394            }
2395
2396            if y[n - 1] == 0 {
2397                rc.enc_bits((up_y[n - 1] < 0) as u32, 1);
2398            }
2399
2400            if resynth {
2401                #[cfg(target_arch = "aarch64")]
2402                unsafe {
2403                    alg_quant_resynth_neon(&up_y, x, n, gain);
2404                }
2405                #[cfg(not(target_arch = "aarch64"))]
2406                alg_quant_resynth_scalar(&up_y, x, n, gain);
2407                exp_rotation(x, n, -1, stride, k, spread);
2408            }
2409            return mask;
2410        }
2411
2412        pvq_search(x, y, k, n);
2413        let mask = extract_collapse_mask(y, n, stride);
2414        encode_pulses(y, n as u32, k as u32, rc);
2415
2416        if resynth {
2417            #[cfg(target_arch = "aarch64")]
2418            unsafe {
2419                alg_quant_resynth_neon(y, x, n, gain);
2420            }
2421            #[cfg(not(target_arch = "aarch64"))]
2422            alg_quant_resynth_scalar(y, x, n, gain);
2423            exp_rotation(x, n, -1, stride, k, spread);
2424        }
2425        mask
2426    } else {
2427        let mut y_mu = [MaybeUninit::<i32>::uninit(); MAX_PVQ_N];
2428        let y = unsafe { std::slice::from_raw_parts_mut(y_mu.as_mut_ptr() as *mut i32, MAX_PVQ_N) };
2429
2430        exp_rotation(x, n, 1, stride, k, spread);
2431        pvq_search(x, &mut y[..n], k, n);
2432        let mask = extract_collapse_mask(&y[..n], n, stride);
2433        encode_pulses(&y[..n], n as u32, k as u32, rc);
2434
2435        if resynth {
2436            let mut ryy = 0.0f32;
2437            for i in 0..n {
2438                let v = y[i] as f32;
2439                x[i] = v;
2440                ryy += v * v;
2441            }
2442            let g = gain / (1e-15 + ryy).sqrt();
2443            for i in 0..n {
2444                x[i] *= g;
2445            }
2446            exp_rotation(x, n, -1, stride, k, spread);
2447        }
2448        mask
2449    }
2450}
2451
2452#[inline]
2453pub fn alg_unquant(
2454    x: &mut [f32],
2455    n: usize,
2456    k: i32,
2457    spread: i32,
2458    stride: usize,
2459    rc: &mut RangeCoder,
2460    gain: f32,
2461) -> u32 {
2462    let mut y_mu = [MaybeUninit::<i32>::uninit(); MAX_PVQ_N];
2463    let y = unsafe { std::slice::from_raw_parts_mut(y_mu.as_mut_ptr() as *mut i32, MAX_PVQ_N) };
2464    decode_pulses(&mut y[..n], n as u32, k as u32, rc);
2465
2466    let mask = extract_collapse_mask(&y[..n], n, stride);
2467
2468    #[cfg(target_arch = "aarch64")]
2469    unsafe {
2470        alg_quant_resynth_neon(&y[..n], x, n, gain);
2471    }
2472    #[cfg(not(target_arch = "aarch64"))]
2473    {
2474        alg_quant_resynth_scalar(&y[..n], x, n, gain);
2475    }
2476
2477    exp_rotation(x, n, -1, stride, k, spread);
2478
2479    mask
2480}