core/slice/sort/unstable/quicksort.rs
1//! This module contains an unstable quicksort and two partition implementations.
2
3#[cfg(not(feature = "optimize_for_size"))]
4use crate::mem;
5use crate::mem::ManuallyDrop;
6#[cfg(not(feature = "optimize_for_size"))]
7use crate::slice::sort::shared::pivot::choose_pivot;
8#[cfg(not(feature = "optimize_for_size"))]
9use crate::slice::sort::shared::smallsort::UnstableSmallSortTypeImpl;
10#[cfg(not(feature = "optimize_for_size"))]
11use crate::slice::sort::unstable::heapsort;
12use crate::{cfg_select, intrinsics, ptr};
13
14/// Sorts `v` recursively.
15///
16/// If the slice had a predecessor in the original array, it is specified as `ancestor_pivot`.
17///
18/// `limit` is the number of allowed imbalanced partitions before switching to `heapsort`. If zero,
19/// this function will immediately switch to heapsort.
20#[cfg(not(feature = "optimize_for_size"))]
21pub(crate) fn quicksort<'a, T, F>(
22 mut v: &'a mut [T],
23 mut ancestor_pivot: Option<&'a T>,
24 mut limit: u32,
25 is_less: &mut F,
26) where
27 F: FnMut(&T, &T) -> bool,
28{
29 loop {
30 if v.len() <= T::small_sort_threshold() {
31 T::small_sort(v, is_less);
32 return;
33 }
34
35 // If too many bad pivot choices were made, simply fall back to heapsort in order to
36 // guarantee `O(N x log(N))` worst-case.
37 if limit == 0 {
38 heapsort::heapsort(v, is_less);
39 return;
40 }
41
42 limit -= 1;
43
44 // Choose a pivot and try guessing whether the slice is already sorted.
45 let pivot_pos = choose_pivot(v, is_less);
46
47 // If the chosen pivot is equal to the predecessor, then it's the smallest element in the
48 // slice. Partition the slice into elements equal to and elements greater than the pivot.
49 // This case is usually hit when the slice contains many duplicate elements.
50 if let Some(p) = ancestor_pivot {
51 if !is_less(p, &v[pivot_pos]) {
52 let num_lt = partition(v, pivot_pos, &mut |a, b| !is_less(b, a));
53
54 // Continue sorting elements greater than the pivot. We know that `num_lt` contains
55 // the pivot. So we can continue after `num_lt`.
56 v = &mut v[(num_lt + 1)..];
57 ancestor_pivot = None;
58 continue;
59 }
60 }
61
62 // Partition the slice.
63 let num_lt = partition(v, pivot_pos, is_less);
64 // SAFETY: partition ensures that `num_lt` will be in-bounds.
65 unsafe { intrinsics::assume(num_lt < v.len()) };
66
67 // Split the slice into `left`, `pivot`, and `right`.
68 let (left, right) = v.split_at_mut(num_lt);
69 let (pivot, right) = right.split_at_mut(1);
70 let pivot = &pivot[0];
71
72 // Recurse into the left side. We have a fixed recursion limit, testing shows no real
73 // benefit for recursing into the shorter side.
74 quicksort(left, ancestor_pivot, limit, is_less);
75
76 // Continue with the right side.
77 v = right;
78 ancestor_pivot = Some(pivot);
79 }
80}
81
82/// Takes the input slice `v` and re-arranges elements such that when the call returns normally
83/// all elements that compare true for `is_less(elem, pivot)` where `pivot == v[pivot_pos]` are
84/// on the left side of `v` followed by the other elements, notionally considered greater or
85/// equal to `pivot`.
86///
87/// Returns the number of elements that are compared true for `is_less(elem, pivot)`.
88///
89/// If `is_less` does not implement a total order the resulting order and return value are
90/// unspecified. All original elements will remain in `v` and any possible modifications via
91/// interior mutability will be observable. Same is true if `is_less` panics or `v.len()`
92/// exceeds `scratch.len()`.
93pub(crate) fn partition<T, F>(v: &mut [T], pivot: usize, is_less: &mut F) -> usize
94where
95 F: FnMut(&T, &T) -> bool,
96{
97 let len = v.len();
98
99 // Allows for panic-free code-gen by proving this property to the compiler.
100 if len == 0 {
101 return 0;
102 }
103
104 if pivot >= len {
105 intrinsics::abort();
106 }
107
108 // SAFETY: We checked that `pivot` is in-bounds.
109 unsafe {
110 // Place the pivot at the beginning of slice.
111 v.swap_unchecked(0, pivot);
112 }
113 let (pivot, v_without_pivot) = v.split_at_mut(1);
114
115 // Assuming that Rust generates noalias LLVM IR we can be sure that a partition function
116 // signature of the form `(v: &mut [T], pivot: &T)` guarantees that pivot and v can't alias.
117 // Having this guarantee is crucial for optimizations. It's possible to copy the pivot value
118 // into a stack value, but this creates issues for types with interior mutability mandating
119 // a drop guard.
120 let pivot = &mut pivot[0];
121
122 // This construct is used to limit the LLVM IR generated, which saves large amounts of
123 // compile-time by only instantiating the code that is needed. Idea by Frank Steffahn.
124 let num_lt = (const { inst_partition::<T, F>() })(v_without_pivot, pivot, is_less);
125
126 if num_lt >= len {
127 intrinsics::abort();
128 }
129
130 // SAFETY: We checked that `num_lt` is in-bounds.
131 unsafe {
132 // Place the pivot between the two partitions.
133 v.swap_unchecked(0, num_lt);
134 }
135
136 num_lt
137}
138
139const fn inst_partition<T, F: FnMut(&T, &T) -> bool>() -> fn(&mut [T], &T, &mut F) -> usize {
140 const MAX_BRANCHLESS_PARTITION_SIZE: usize = 96;
141 if size_of::<T>() <= MAX_BRANCHLESS_PARTITION_SIZE {
142 // Specialize for types that are relatively cheap to copy, where branchless optimizations
143 // have large leverage e.g. `u64` and `String`.
144 cfg_select! {
145 feature = "optimize_for_size" => partition_lomuto_branchless_simple::<T, F>,
146 _ => partition_lomuto_branchless_cyclic::<T, F>,
147 }
148 } else {
149 partition_hoare_branchy_cyclic::<T, F>
150 }
151}
152
153/// See [`partition`].
154fn partition_hoare_branchy_cyclic<T, F>(v: &mut [T], pivot: &T, is_less: &mut F) -> usize
155where
156 F: FnMut(&T, &T) -> bool,
157{
158 let len = v.len();
159
160 if len == 0 {
161 return 0;
162 }
163
164 // Optimized for large types that are expensive to move. Not optimized for integers. Optimized
165 // for small code-gen, assuming that is_less is an expensive operation that generates
166 // substantial amounts of code or a call. And that copying elements will likely be a call to
167 // memcpy. Using 2 `ptr::copy_nonoverlapping` has the chance to be faster than
168 // `ptr::swap_nonoverlapping` because `memcpy` can use wide SIMD based on runtime feature
169 // detection. Benchmarks support this analysis.
170
171 let mut gap_opt: Option<GapGuard<T>> = None;
172
173 // SAFETY: The left-to-right scanning loop performs a bounds check, where we know that `left >=
174 // v_base && left < right && right <= v_base.add(len)`. The right-to-left scanning loop performs
175 // a bounds check ensuring that `right` is in-bounds. We checked that `len` is more than zero,
176 // which means that unconditional `right = right.sub(1)` is safe to do. The exit check makes
177 // sure that `left` and `right` never alias, making `ptr::copy_nonoverlapping` safe. The
178 // drop-guard `gap` ensures that should `is_less` panic we always overwrite the duplicate in the
179 // input. `gap.pos` stores the previous value of `right` and starts at `right` and so it too is
180 // in-bounds. We never pass the saved `gap.value` to `is_less` while it is inside the `GapGuard`
181 // thus any changes via interior mutability will be observed.
182 unsafe {
183 let v_base = v.as_mut_ptr();
184
185 let mut left = v_base;
186 let mut right = v_base.add(len);
187
188 loop {
189 // Find the first element greater than the pivot.
190 while left < right && is_less(&*left, pivot) {
191 left = left.add(1);
192 }
193
194 // Find the last element equal to the pivot.
195 loop {
196 right = right.sub(1);
197 if left >= right || is_less(&*right, pivot) {
198 break;
199 }
200 }
201
202 if left >= right {
203 break;
204 }
205
206 // Swap the found pair of out-of-order elements via cyclic permutation.
207 let is_first_swap_pair = gap_opt.is_none();
208
209 if is_first_swap_pair {
210 gap_opt = Some(GapGuard { pos: right, value: ManuallyDrop::new(ptr::read(left)) });
211 }
212
213 let gap = gap_opt.as_mut().unwrap_unchecked();
214
215 // Single place where we instantiate ptr::copy_nonoverlapping in the partition.
216 if !is_first_swap_pair {
217 ptr::copy_nonoverlapping(left, gap.pos, 1);
218 }
219 gap.pos = right;
220 ptr::copy_nonoverlapping(right, left, 1);
221
222 left = left.add(1);
223 }
224
225 left.offset_from_unsigned(v_base)
226
227 // `gap_opt` goes out of scope and overwrites the last wrong-side element on the right side
228 // with the first wrong-side element of the left side that was initially overwritten by the
229 // first wrong-side element on the right side element.
230 }
231}
232
233#[cfg(not(feature = "optimize_for_size"))]
234struct PartitionState<T> {
235 // The current element that is being looked at, scans left to right through slice.
236 right: *mut T,
237 // Counts the number of elements that compared less-than, also works around:
238 // https://github.com/rust-lang/rust/issues/117128
239 num_lt: usize,
240 // Gap guard that tracks the temporary duplicate in the input.
241 gap: GapGuardRaw<T>,
242}
243
244#[cfg(not(feature = "optimize_for_size"))]
245fn partition_lomuto_branchless_cyclic<T, F>(v: &mut [T], pivot: &T, is_less: &mut F) -> usize
246where
247 F: FnMut(&T, &T) -> bool,
248{
249 // Novel partition implementation by Lukas Bergdoll and Orson Peters. Branchless Lomuto
250 // partition paired with a cyclic permutation.
251 // https://github.com/Voultapher/sort-research-rs/blob/main/writeup/lomcyc_partition/text.md
252
253 let len = v.len();
254 let v_base = v.as_mut_ptr();
255
256 if len == 0 {
257 return 0;
258 }
259
260 // SAFETY: We checked that `len` is more than zero, which means that reading `v_base` is safe to
261 // do. From there we have a bounded loop where `v_base.add(i)` is guaranteed in-bounds. `v` and
262 // `pivot` can't alias because of type system rules. The drop-guard `gap` ensures that should
263 // `is_less` panic we always overwrite the duplicate in the input. `gap.pos` stores the previous
264 // value of `right` and starts at `v_base` and so it too is in-bounds. Given `UNROLL_LEN == 2`
265 // after the main loop we either have A) the last element in `v` that has not yet been processed
266 // because `len % 2 != 0`, or B) all elements have been processed except the gap value that was
267 // saved at the beginning with `ptr::read(v_base)`. In the case A) the loop will iterate twice,
268 // first performing loop_body to take care of the last element that didn't fit into the unroll.
269 // After that the behavior is the same as for B) where we use the saved value as `right` to
270 // overwrite the duplicate. If this very last call to `is_less` panics the saved value will be
271 // copied back including all possible changes via interior mutability. If `is_less` does not
272 // panic and the code continues we overwrite the duplicate and do `right = right.add(1)`, this
273 // is safe to do with `&mut *gap.value` because `T` is the same as `[T; 1]` and generating a
274 // pointer one past the allocation is safe.
275 unsafe {
276 let mut loop_body = |state: &mut PartitionState<T>| {
277 let right_is_lt = is_less(&*state.right, pivot);
278 let left = v_base.add(state.num_lt);
279
280 ptr::copy(left, state.gap.pos, 1);
281 ptr::copy_nonoverlapping(state.right, left, 1);
282
283 state.gap.pos = state.right;
284 state.num_lt += right_is_lt as usize;
285
286 state.right = state.right.add(1);
287 };
288
289 // Ideally we could just use GapGuard in PartitionState, but the reference that is
290 // materialized with `&mut state` when calling `loop_body` would create a mutable reference
291 // to the parent struct that contains the gap value, invalidating the reference pointer
292 // created from a reference to the gap value in the cleanup loop. This is only an issue
293 // under Stacked Borrows, Tree Borrows accepts the intuitive code using GapGuard as valid.
294 let mut gap_value = ManuallyDrop::new(ptr::read(v_base));
295
296 let mut state = PartitionState {
297 num_lt: 0,
298 right: v_base.add(1),
299
300 gap: GapGuardRaw { pos: v_base, value: &mut *gap_value },
301 };
302
303 // Manual unrolling that works well on x86, Arm and with opt-level=s without murdering
304 // compile-times. Leaving this to the compiler yields ok to bad results.
305 let unroll_len = const { if size_of::<T>() <= 16 { 2 } else { 1 } };
306
307 let unroll_end = v_base.add(len - (unroll_len - 1));
308 while state.right < unroll_end {
309 if unroll_len == 2 {
310 loop_body(&mut state);
311 loop_body(&mut state);
312 } else {
313 loop_body(&mut state);
314 }
315 }
316
317 // Single instantiate `loop_body` for both the unroll cleanup and cyclic permutation
318 // cleanup. Optimizes binary-size and compile-time.
319 let end = v_base.add(len);
320 loop {
321 let is_done = state.right == end;
322 state.right = if is_done { state.gap.value } else { state.right };
323
324 loop_body(&mut state);
325
326 if is_done {
327 mem::forget(state.gap);
328 break;
329 }
330 }
331
332 state.num_lt
333 }
334}
335
336#[cfg(feature = "optimize_for_size")]
337fn partition_lomuto_branchless_simple<T, F: FnMut(&T, &T) -> bool>(
338 v: &mut [T],
339 pivot: &T,
340 is_less: &mut F,
341) -> usize {
342 let mut left = 0;
343
344 for right in 0..v.len() {
345 // SAFETY: `left` can at max be incremented by 1 each loop iteration, which implies that
346 // left <= right and that both are in-bounds.
347 unsafe {
348 let right_is_lt = is_less(v.get_unchecked(right), pivot);
349 v.swap_unchecked(left, right);
350 left += right_is_lt as usize;
351 }
352 }
353
354 left
355}
356
357struct GapGuard<T> {
358 pos: *mut T,
359 value: ManuallyDrop<T>,
360}
361
362impl<T> Drop for GapGuard<T> {
363 fn drop(&mut self) {
364 // SAFETY: `self` MUST be constructed in a way that makes copying the gap value into
365 // `self.pos` sound.
366 unsafe {
367 ptr::copy_nonoverlapping(&*self.value, self.pos, 1);
368 }
369 }
370}
371
372/// Ideally this wouldn't be needed and we could just use the regular GapGuard.
373/// See comment in [`partition_lomuto_branchless_cyclic`].
374#[cfg(not(feature = "optimize_for_size"))]
375struct GapGuardRaw<T> {
376 pos: *mut T,
377 value: *mut T,
378}
379
380#[cfg(not(feature = "optimize_for_size"))]
381impl<T> Drop for GapGuardRaw<T> {
382 fn drop(&mut self) {
383 // SAFETY: `self` MUST be constructed in a way that makes copying the gap value into
384 // `self.pos` sound.
385 unsafe {
386 ptr::copy_nonoverlapping(self.value, self.pos, 1);
387 }
388 }
389}