Skip to main content

orx_parallel/sort/
sort_slice.rs

1use crate::{IntoParIter, NumThreads, Par, Params, ThreadPool, runner::ParRunner};
2use alloc::vec::Vec;
3use core::{mem::MaybeUninit, num::NonZeroUsize};
4
5/// Task for sorting a disjoint chunk of the slice in parallel.
6struct SortChunkTask<T> {
7    ptr: *mut T,
8    len: usize,
9}
10unsafe impl<T: Send> Send for SortChunkTask<T> {}
11unsafe impl<T: Send> Sync for SortChunkTask<T> {}
12
13impl<T: Ord> SortChunkTask<T> {
14    #[inline]
15    fn execute(self) {
16        if self.len > 1 {
17            let s = unsafe { core::slice::from_raw_parts_mut(self.ptr, self.len) };
18            s.sort_unstable();
19        }
20    }
21}
22
23/// Binary searches the split point in two sorted slices A and B such that the first `r`
24/// elements of merged(A, B) consist of A[0..a_split] and B[0..b_split] with a_split + b_split == r.
25fn find_split<T: Ord>(a: &[T], b: &[T], r: usize) -> (usize, usize) {
26    let mut low = r.saturating_sub(b.len());
27    let mut high = a.len().min(r);
28    while low < high {
29        let mid = low + (high - low) / 2;
30        let b_idx = r - mid - 1;
31        match a[mid] < b[b_idx] {
32            true => low = mid + 1,
33            false => high = mid,
34        }
35    }
36    (low, r - low)
37}
38
39/// Task for merging two sorted contiguous sub-slices from `src_a` and `src_b` into `dst`.
40struct MergeSubTask<T> {
41    src_a: *const T,
42    len_a: usize,
43    src_b: *const T,
44    len_b: usize,
45    dst: *mut T,
46}
47unsafe impl<T: Send> Send for MergeSubTask<T> {}
48unsafe impl<T: Send> Sync for MergeSubTask<T> {}
49
50impl<T: Ord> MergeSubTask<T> {
51    #[inline]
52    fn execute(self) {
53        if self.len_a == 0 {
54            if self.len_b > 0 {
55                unsafe {
56                    core::ptr::copy_nonoverlapping(self.src_b, self.dst, self.len_b);
57                }
58            }
59            return;
60        }
61        if self.len_b == 0 {
62            unsafe {
63                core::ptr::copy_nonoverlapping(self.src_a, self.dst, self.len_a);
64            }
65            return;
66        }
67
68        // Fast-path: if all elements of A <= all elements of B
69        let last_a = unsafe { &*self.src_a.add(self.len_a - 1) };
70        let first_b = unsafe { &*self.src_b };
71        if last_a <= first_b {
72            unsafe {
73                core::ptr::copy_nonoverlapping(self.src_a, self.dst, self.len_a);
74                core::ptr::copy_nonoverlapping(self.src_b, self.dst.add(self.len_a), self.len_b);
75            }
76            return;
77        }
78
79        // Fast-path: if all elements of B < all elements of A
80        let last_b = unsafe { &*self.src_b.add(self.len_b - 1) };
81        let first_a = unsafe { &*self.src_a };
82        if last_b < first_a {
83            unsafe {
84                core::ptr::copy_nonoverlapping(self.src_b, self.dst, self.len_b);
85                core::ptr::copy_nonoverlapping(self.src_a, self.dst.add(self.len_b), self.len_a);
86            }
87            return;
88        }
89
90        // Standard merge
91        unsafe {
92            let mut ptr_a = self.src_a;
93            let end_a = self.src_a.add(self.len_a);
94            let mut ptr_b = self.src_b;
95            let end_b = self.src_b.add(self.len_b);
96            let mut out = self.dst;
97
98            while ptr_a < end_a && ptr_b < end_b {
99                match *ptr_b < *ptr_a {
100                    true => {
101                        core::ptr::copy_nonoverlapping(ptr_b, out, 1);
102                        ptr_b = ptr_b.add(1);
103                    }
104                    false => {
105                        core::ptr::copy_nonoverlapping(ptr_a, out, 1);
106                        ptr_a = ptr_a.add(1);
107                    }
108                }
109                out = out.add(1);
110            }
111
112            if ptr_a < end_a {
113                let rem = end_a.offset_from(ptr_a) as usize;
114                core::ptr::copy_nonoverlapping(ptr_a, out, rem);
115            } else if ptr_b < end_b {
116                let rem = end_b.offset_from(ptr_b) as usize;
117                core::ptr::copy_nonoverlapping(ptr_b, out, rem);
118            }
119        }
120    }
121}
122
123/// Task for parallel copying back to original slice if needed.
124struct CopyChunkTask<T> {
125    src: *const T,
126    dst: *mut T,
127    len: usize,
128}
129unsafe impl<T: Send> Send for CopyChunkTask<T> {}
130unsafe impl<T: Send> Sync for CopyChunkTask<T> {}
131
132impl<T> CopyChunkTask<T> {
133    #[inline]
134    fn execute(self) {
135        if self.len > 0 {
136            unsafe {
137                core::ptr::copy_nonoverlapping(self.src, self.dst, self.len);
138            }
139        }
140    }
141}
142
143/// Sorts the `slice` in parallel using the provided `runner` and parallelization `params`.
144///
145/// # Examples
146///
147/// ```
148/// use orx_parallel::*;
149///
150/// let mut data = vec![5, 2, 8, 1, 9, 3, 7, 4, 6];
151/// let mut runner = Runner::fixed();
152/// par_experimental_sort(&mut data, &mut runner, Params::default());
153/// assert_eq!(data, vec![1, 2, 3, 4, 5, 6, 7, 8, 9]);
154/// ```
155pub fn par_experimental_sort<T, R>(slice: &mut [T], runner: &mut R, params: Params)
156where
157    T: Ord + Send,
158    R: ParRunner,
159{
160    let n = slice.len();
161    if n <= 1 {
162        return;
163    }
164
165    if params.is_sequential() {
166        slice.sort_unstable();
167        return;
168    }
169
170    let par_num_threads = match params.num_threads {
171        NumThreads::Auto => NonZeroUsize::MAX,
172        NumThreads::Max(x) => x,
173    };
174    let max_threads: usize = runner.pool().max_num_threads().min(par_num_threads).into();
175
176    if max_threads <= 1 || n < 1024 {
177        slice.sort_unstable();
178        return;
179    }
180
181    // Determine number of initial chunks K (must be a power of 2)
182    let mut num_chunks = (max_threads * 2).next_power_of_two();
183    while num_chunks > 2 && n / num_chunks < 512 {
184        num_chunks /= 2;
185    }
186
187    if n / num_chunks < 64 {
188        slice.sort_unstable();
189        return;
190    }
191
192    let orig_ptr = slice.as_mut_ptr();
193
194    // Phase 1: Sort each chunk in parallel
195    let mut sort_tasks = Vec::with_capacity(num_chunks);
196    for i in 0..num_chunks {
197        let start = i * n / num_chunks;
198        let end = (i + 1) * n / num_chunks;
199        let len = end - start;
200        sort_tasks.push(SortChunkTask {
201            ptr: unsafe { orig_ptr.add(start) },
202            len,
203        });
204    }
205
206    let par = sort_tasks.into_par().runner(&mut *runner);
207    let par = params.apply(par);
208    par.for_each(|task| task.execute());
209
210    // Allocate auxiliary buffer for merging
211    let mut aux = Vec::<MaybeUninit<T>>::with_capacity(n);
212    let aux_ptr = aux.as_mut_ptr() as *mut T;
213
214    // Phase 2: Hierarchical merge passes
215    let mut current_chunks = num_chunks;
216    let mut pass = 0;
217
218    while current_chunks > 1 {
219        let (src_base, dst_base) = if pass % 2 == 0 {
220            (orig_ptr, aux_ptr)
221        } else {
222            (aux_ptr, orig_ptr)
223        };
224
225        let next_chunks = current_chunks / 2;
226        let stride = num_chunks / current_chunks;
227        let tasks_per_pair = (max_threads / next_chunks).max(1);
228        let mut merge_tasks = Vec::with_capacity(next_chunks * tasks_per_pair);
229
230        for pair_idx in 0..next_chunks {
231            let chunk_a_idx = pair_idx * 2;
232            let chunk_b_idx = chunk_a_idx + 1;
233
234            let start_a = (chunk_a_idx * stride) * n / num_chunks;
235            let end_a = ((chunk_a_idx + 1) * stride) * n / num_chunks;
236            let len_a = end_a - start_a;
237
238            let start_b = (chunk_b_idx * stride) * n / num_chunks;
239            let end_b = ((chunk_b_idx + 1) * stride) * n / num_chunks;
240            let len_b = end_b - start_b;
241
242            let slice_a = unsafe { core::slice::from_raw_parts(src_base.add(start_a), len_a) };
243            let slice_b = unsafe { core::slice::from_raw_parts(src_base.add(start_b), len_b) };
244
245            let pair_total = len_a + len_b;
246
247            let mut prev_a = 0;
248            let mut prev_b = 0;
249            let mut prev_r = 0;
250
251            for t in 1..=tasks_per_pair {
252                let r = t * pair_total / tasks_per_pair;
253                let (curr_a, curr_b) = if t == tasks_per_pair {
254                    (len_a, len_b)
255                } else {
256                    find_split(slice_a, slice_b, r)
257                };
258
259                let sub_len_a = curr_a - prev_a;
260                let sub_len_b = curr_b - prev_b;
261                let sub_dst = unsafe { dst_base.add(start_a + prev_r) };
262
263                merge_tasks.push(MergeSubTask {
264                    src_a: unsafe { slice_a.as_ptr().add(prev_a) },
265                    len_a: sub_len_a,
266                    src_b: unsafe { slice_b.as_ptr().add(prev_b) },
267                    len_b: sub_len_b,
268                    dst: sub_dst,
269                });
270
271                prev_a = curr_a;
272                prev_b = curr_b;
273                prev_r = r;
274            }
275        }
276
277        let par = merge_tasks.into_par().runner(&mut *runner);
278        let par = params.apply(par);
279        par.for_each(|task| task.execute());
280
281        current_chunks = next_chunks;
282        pass += 1;
283    }
284
285    // If odd number of passes, the sorted data is in `aux_ptr`, copy back to `orig_ptr`
286    if pass % 2 != 0 {
287        let copy_chunks = max_threads.min(n);
288        let mut copy_tasks = Vec::with_capacity(copy_chunks);
289        for i in 0..copy_chunks {
290            let start = i * n / copy_chunks;
291            let end = (i + 1) * n / copy_chunks;
292            let len = end - start;
293            copy_tasks.push(CopyChunkTask {
294                src: unsafe { aux_ptr.add(start) },
295                dst: unsafe { orig_ptr.add(start) },
296                len,
297            });
298        }
299
300        let par = copy_tasks.into_par().runner(&mut *runner);
301        let par = params.apply(par);
302        par.for_each(|task| task.execute());
303    }
304}
305
306#[cfg(test)]
307mod tests {
308    use super::*;
309    use crate::runner::default_runner;
310    use alloc::string::String;
311
312    #[test]
313    fn test_sort_empty_and_single() {
314        let mut runner = default_runner();
315        let mut empty: [i32; 0] = [];
316        par_experimental_sort(&mut empty, &mut runner, Params::default());
317
318        let mut single = [42];
319        par_experimental_sort(&mut single, &mut runner, Params::default());
320        assert_eq!(single, [42]);
321    }
322
323    #[test]
324    fn test_sort_small_slices() {
325        let mut runner = default_runner();
326        let mut data = [9, 3, 7, 1, 5, 2, 8, 4, 6];
327        par_experimental_sort(&mut data, &mut runner, Params::default());
328        assert_eq!(data, [1, 2, 3, 4, 5, 6, 7, 8, 9]);
329    }
330
331    #[test]
332    fn test_sort_medium_and_large_random() {
333        let mut runner = default_runner();
334
335        for size in [500, 1024, 2048, 5000, 20000, 50000] {
336            let mut data: Vec<i32> = (0..size as u64)
337                .map(|i| (i.wrapping_mul(1103515245).wrapping_add(12345) & 0x7FFFFFFF) as i32)
338                .collect();
339            let mut expected = data.clone();
340            expected.sort_unstable();
341
342            par_experimental_sort(&mut data, &mut runner, Params::default());
343            assert_eq!(data, expected, "Failed for size {}", size);
344        }
345    }
346
347    #[test]
348    fn test_sort_sorted_and_reversed() {
349        let mut runner = default_runner();
350        let size = 10000;
351
352        let mut sorted: Vec<i32> = (0..size).collect();
353        par_experimental_sort(&mut sorted, &mut runner, Params::default());
354        assert!(sorted.windows(2).all(|w| w[0] <= w[1]));
355
356        let mut reversed: Vec<i32> = (0..size).rev().collect();
357        par_experimental_sort(&mut reversed, &mut runner, Params::default());
358        assert!(reversed.windows(2).all(|w| w[0] <= w[1]));
359    }
360
361    #[test]
362    fn test_sort_high_duplicates() {
363        let mut runner = default_runner();
364        let size = 20000;
365        let mut data: Vec<i32> = (0..size).map(|i| i % 7).collect();
366        let mut expected = data.clone();
367        expected.sort_unstable();
368
369        par_experimental_sort(&mut data, &mut runner, Params::default());
370        assert_eq!(data, expected);
371    }
372
373    #[test]
374    fn test_sort_non_copy_types() {
375        let mut runner = default_runner();
376        let size = 5000;
377        let mut data: Vec<String> = (0..size)
378            .map(|i| alloc::format!("item_{:06}", (i * 7919) % size))
379            .collect();
380        let mut expected = data.clone();
381        expected.sort();
382
383        par_experimental_sort(&mut data, &mut runner, Params::default());
384        assert_eq!(data, expected);
385    }
386
387    #[test]
388    fn test_sort_num_threads_configs() {
389        for nt in [1, 2, 4, 8] {
390            let mut runner = default_runner();
391            let size = 10000;
392            let mut data: Vec<i32> = (0..size as u64)
393                .map(|i| (i.wrapping_mul(2654435761) & 0x7FFFFFFF) as i32)
394                .collect();
395            let mut expected = data.clone();
396            expected.sort_unstable();
397
398            let params = Params::default().with_num_threads(nt);
399            par_experimental_sort(&mut data, &mut runner, params);
400            assert_eq!(data, expected, "Failed for num_threads = {}", nt);
401        }
402    }
403}