1use assume::assume;
5use rostl_primitives::{
6 traits::{Cmov, CswapIndex},
7 utils::get_smaller_or_equal_power_of_two,
8};
9pub fn compute_prefix_sum<T, F>(arr: &[T], is_dummy: F) -> Vec<usize>
18where
19 F: Fn(&T) -> bool,
20{
21 let size = arr.len();
22 let mut sarr = vec![0; size + 1];
23 for i in 0..size {
24 let mut adder = 1usize;
25 adder.cmov(&0, is_dummy(&arr[i]));
26 sarr[i + 1] = sarr[i] + adder;
27 }
28 sarr
29}
30
31#[deprecated(note = "use compact instead, it's faster")]
41pub fn compact_goodrich<T, F>(arr: &mut [T], is_dummy: F) -> usize
42where
43 F: Fn(&T) -> bool,
44 T: Cmov + Copy,
45{
46 if arr.is_empty() {
47 return 0;
48 }
49 let l2len = arr.len().next_power_of_two().trailing_zeros() as usize;
50 let mut csum = vec![0; arr.len()];
51 let mut dummy_count = 0;
52
53 csum[0] = 0;
54 let pred = is_dummy(&arr[0]);
55 dummy_count.cmov(&1, pred);
56 for i in 1..arr.len() {
57 csum[i] = 0;
58 let pred = is_dummy(&arr[i]);
59 dummy_count.cmov(&(dummy_count + 1), pred);
60 csum[i].cmov(&(dummy_count), !pred);
61 }
62 let ret = arr.len() - dummy_count;
63
64 for i in 0..l2len {
65 let offset = 1 << i;
66 for j in 0..(arr.len() - offset) {
67 let a = j;
68 let b = j + offset;
69 let pred = (csum[b] & offset) != 0;
70 arr.cswap(a, b, pred);
71 let newacsum = csum[b].wrapping_sub(offset);
72 csum[a].cmov(&newacsum, pred);
73 csum[b].cmov(&0, pred);
74 }
75 }
76
77 ret
78}
79
80fn compact_payload_offset<T>(arr: &mut [T], payload: &[usize], z: usize)
87where
88 T: Cmov + Copy,
89{
90 assume!(unsafe: arr.len()+1 == payload.len());
91 let n = arr.len();
92 let half_n = n / 2;
93 let m = payload[half_n] - payload[0];
94 if n == 2 {
95 let should_swap = ((!m) & (payload[2] - payload[1])) != z;
96 arr.cswap(0, 1, should_swap);
97 return;
98 }
99 let zleft = z % half_n;
100 let zright = (z + m) % half_n;
101 compact_payload_offset(&mut arr[..half_n], &payload[..half_n + 1], zleft);
102 compact_payload_offset(&mut arr[half_n..], &payload[half_n..], zright);
103
104 let s_a = zleft + m >= half_n;
105 let s_b = z >= half_n;
106 let s = s_a ^ s_b;
107
108 for i in 0..half_n {
109 let left = i;
110 let right = i + half_n;
111 let cond = s ^ (i >= zright);
112 assume!(unsafe: left < arr.len());
113 assume!(unsafe: right < arr.len());
114 arr.cswap(left, right, cond);
115 }
116}
117
118pub fn compact_payload<T>(arr: &mut [T], payload: &[usize])
130where
131 T: Cmov + Copy,
132{
133 assume!(unsafe: arr.len() + 1 == payload.len());
134 let n = arr.len();
135 if n <= 1 {
136 return;
137 }
138
139 let n1 = get_smaller_or_equal_power_of_two(n);
140 let n2 = n - n1;
141
142 if n2 == 0 {
143 compact_payload_offset(arr, payload, 0);
144 return;
145 }
146
147 let m = payload[n2] - payload[0];
148 compact_payload(arr[..n2].as_mut(), &payload[..n2 + 1]);
149 compact_payload_offset(arr[n2..].as_mut(), &payload[n2..], (n1 - n2 + m) % n1);
150
151 for i in 0..n2 {
152 let left = i;
153 let right = i + n1;
154 assume!(unsafe: left < arr.len());
155 assume!(unsafe: right < arr.len());
156 arr.cswap(left, right, i >= m);
157 }
158}
159
160pub fn compact<T, F>(arr: &mut [T], is_dummy: F) -> usize
173where
174 F: Fn(&T) -> bool,
175 T: Cmov + Copy,
176{
177 let payload = compute_prefix_sum(arr, is_dummy);
178 compact_payload(arr, &payload);
179 payload[payload.len() - 1]
180}
181
182fn distribute_payload_offset<T>(arr: &mut [T], payload: &[usize], z: usize)
183where
184 T: Cmov + Copy,
185{
186 assume!(unsafe: arr.len()+1 == payload.len());
187 let n = arr.len();
188 let half_n = n / 2;
189 let m = payload[half_n] - payload[0];
190 if n == 2 {
191 let should_swap = ((!m) & (payload[2] - payload[1])) != z;
192 arr.cswap(0, 1, should_swap);
193 return;
194 }
195 let zleft = z % half_n;
196 let zright = (z + m) % half_n;
197 let s_a = zleft + m >= half_n;
198 let s_b = z >= half_n;
199 let s = s_a ^ s_b;
200
201 for i in 0..half_n {
202 let left = i;
203 let right = i + half_n;
204 let cond = s ^ (i >= zright);
205 assume!(unsafe: left < arr.len());
206 assume!(unsafe: right < arr.len());
207 arr.cswap(left, right, cond);
208 }
209 distribute_payload_offset(&mut arr[..half_n], &payload[..half_n + 1], zleft);
210 distribute_payload_offset(&mut arr[half_n..], &payload[half_n..], zright);
211}
212
213pub fn distribute_payload<T>(arr: &mut [T], payload: &[usize])
217where
218 T: Cmov + Copy,
219{
220 assume!(unsafe: arr.len() + 1 == payload.len());
221 let n = arr.len();
222 if n <= 1 {
223 return;
224 }
225
226 let n1 = get_smaller_or_equal_power_of_two(n);
227 let n2 = n - n1;
228
229 if n2 == 0 {
230 distribute_payload_offset(arr, payload, 0);
231 return;
232 }
233
234 let m = payload[n2] - payload[0];
235
236 for i in 0..n2 {
237 let left = i;
238 let right = i + n1;
239 assume!(unsafe: left < arr.len());
240 assume!(unsafe: right < arr.len());
241 arr.cswap(left, right, i >= m);
242 }
243
244 distribute_payload(arr[..n2].as_mut(), &payload[..n2 + 1]);
245 distribute_payload_offset(arr[n2..].as_mut(), &payload[n2..], (n1 - n2 + m) % n1);
246}
247
248#[cfg(test)]
249#[allow(deprecated)]
250mod tests {
251 use rand::Rng;
252
253 use super::*;
254
255 #[test]
256 fn test_compact() {
257 let mut arr = [1, 2, 3, 4, 5];
258 let new_len = compact(&mut arr, |x| *x % 2 == 0);
259 assert_eq!(new_len, 3);
260 assert_eq!(&arr[..new_len], &[1, 3, 5]);
261
262 let mut arr = [1, 2, 3, 4, 5];
263 compact_goodrich(&mut arr, |x| *x % 2 == 0);
264 assert_eq!(&arr[..3], &[1, 3, 5]);
265 }
266
267 #[test]
268 fn test_small() {
269 let mut arr: Vec<i32> = vec![1];
270 let new_len = compact(&mut arr, |x| *x % 2 == 0);
271 assert_eq!(new_len, 1);
272 assert_eq!(&arr[..new_len], &[1]);
273 let mut arr: Vec<i32> = vec![1];
274 compact_goodrich(&mut arr, |x| *x % 2 == 0);
275 assert_eq!(&arr[..1], &[1]);
276
277 let mut arr: Vec<i32> = vec![2];
278 let new_len = compact(&mut arr, |x| *x % 2 == 0);
279 assert_eq!(new_len, 0);
280 assert_eq!(&arr[..new_len], &[]);
281
282 let mut arr: Vec<i32> = vec![1, 2];
283 let new_len = compact(&mut arr, |x| *x % 2 == 0);
284 assert_eq!(new_len, 1);
285 assert_eq!(&arr[..new_len], &[1]);
286 let mut arr: Vec<i32> = vec![1, 2];
287 compact_goodrich(&mut arr, |x| *x % 2 == 0);
288 assert_eq!(&arr[..1], &[1]);
289
290 let mut arr: Vec<i32> = vec![];
291 let new_len = compact(&mut arr, |x| *x % 2 == 0);
292 assert_eq!(new_len, 0);
293 assert_eq!(&arr[..new_len], &[]);
294 let mut arr: Vec<i32> = vec![];
295 compact_goodrich(&mut arr, |x| *x % 2 == 0);
296 assert_eq!(&arr[..0], &[]);
297 }
298
299 #[test]
300 fn test_many_sizes() {
301 let mut rng = rand::rng();
303 for _i in 0..100 {
304 let size = rng.random_range(0..2050);
305 let arr: Vec<i32> = (0..size).map(|_| rng.random_range(0..100)).collect();
306 let mut arr1 = arr.clone();
307 let new_len = compact(&mut arr1, |x| *x % 2 == 0);
308 for itm in arr1.iter().take(new_len) {
309 assert!(itm % 2 != 0);
310 }
311 for itm in arr1.iter().skip(new_len) {
312 assert!(itm % 2 == 0);
313 }
314 let mut arr2 = arr.clone();
315 compact_goodrich(&mut arr2, |x| *x % 2 == 0);
316 for itm in arr2.iter().take(new_len) {
317 assert!(itm % 2 != 0);
318 }
319 for itm in arr2.iter().skip(new_len) {
320 assert!(itm % 2 == 0);
321 }
322 }
323 }
324
325 #[test]
326 fn test_distribute() {
327 let mut arr = [1, 3, 5, 0, 2, 4];
328 let payload = [0, 1, 2, 3, 3, 4, 5];
329 distribute_payload(&mut arr, &payload);
330 assert_eq!(&arr, &[1, 3, 5, 4, 0, 2]);
331
332 let mut arr = [1, 2, 3, 4, 5];
333 let payload = [0, 1, 1, 2, 2, 3];
334 compact_payload(&mut arr, &payload);
335 assert_eq!(&arr[..3], &[1, 3, 5]);
336 distribute_payload(&mut arr, &payload);
337 assert_eq!(&arr, &[1, 2, 3, 4, 5]);
338 }
339
340 #[test]
341 fn test_distribute_after_compact_rands() {
342 let mut rng = rand::rng();
343 for _i in 0..100 {
344 let size = rng.random_range(0..2050);
345 let arr: Vec<i32> = (0..size).map(|_| rng.random_range(0..100)).collect();
346 let mut arr1 = arr.clone();
347 let mut payload = vec![0; size + 1];
348 for i in 0..size {
349 let mut adder = 1usize;
350 adder.cmov(&0, arr[i] % 2 == 0);
351 payload[i + 1] = payload[i] + adder;
352 }
353 compact_payload(&mut arr1, &payload);
354 distribute_payload(&mut arr1, &payload);
355 assert_eq!(&arr1, &arr);
356 }
357 }
358}