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 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 let vy1 = vfmsq_f32(vmulq_f32(vx1, vc), vs, vx2);
1814 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}