Skip to main content

ruprim_host/sort/
mod.rs

1//! Sort and argsort operations for FlexTensor.
2//!
3//! Operates directly on storage without TensorData round-trips.
4
5use alloc::vec;
6use alloc::vec::Vec;
7use ruda_core::tensor::{DType, element::Element};
8use ruda_core::{bytes::Bytes, tensor::Shape};
9use half::{bf16, f16};
10use bytemuck::Pod;
11
12#[cfg(feature = "rayon")]
13use rayon::prelude::*;
14
15use ruda_core::tensor::host::{HostTensor, Layout};
16
17use ruda_core::tensor::host::dtype::INDEX_DTYPE;
18#[cfg(feature = "rayon")]
19use ruda_core::tensor::host::parallel::PARALLEL_THRESHOLD;
20
21/// Validate sort dimension and check for empty tensors.
22/// Returns `true` if the tensor is empty (caller should return early).
23fn validate_sort_args(shape: &Shape, dim: usize) -> bool {
24    assert!(
25        dim < shape.num_dims(),
26        "sort: dim {} out of bounds for tensor with {} dimensions",
27        dim,
28        shape.num_dims()
29    );
30    let dim_size = shape[dim];
31    assert!(
32        dim_size <= isize::MAX as usize,
33        "sort: dimension {} has size {} which exceeds isize::MAX",
34        dim,
35        dim_size
36    );
37    shape.num_elements() == 0
38}
39
40/// Sort elements along a dimension, returning the sorted tensor.
41pub fn sort(tensor: HostTensor, dim: usize, descending: bool) -> HostTensor {
42    match tensor.dtype() {
43        DType::F32 => sort_typed::<f32>(tensor, dim, descending, f32::total_cmp),
44        DType::F64 => sort_typed::<f64>(tensor, dim, descending, f64::total_cmp),
45        DType::F16 => sort_half(tensor, dim, descending, f16::to_f32, f16::from_f32),
46        DType::BF16 => sort_half(tensor, dim, descending, bf16::to_f32, bf16::from_f32),
47        DType::I64 => sort_typed::<i64>(tensor, dim, descending, Ord::cmp),
48        DType::I32 => sort_typed::<i32>(tensor, dim, descending, Ord::cmp),
49        DType::I16 => sort_typed::<i16>(tensor, dim, descending, Ord::cmp),
50        DType::I8 => sort_typed::<i8>(tensor, dim, descending, Ord::cmp),
51        DType::U64 => sort_typed::<u64>(tensor, dim, descending, Ord::cmp),
52        DType::U32 => sort_typed::<u32>(tensor, dim, descending, Ord::cmp),
53        DType::U16 => sort_typed::<u16>(tensor, dim, descending, Ord::cmp),
54        DType::U8 => sort_typed::<u8>(tensor, dim, descending, Ord::cmp),
55        dt => panic!("sort: unsupported dtype {:?}", dt),
56    }
57}
58
59/// Sort elements along a dimension, returning (sorted_values, indices).
60pub fn sort_with_indices(
61    tensor: HostTensor,
62    dim: usize,
63    descending: bool,
64) -> (HostTensor, HostTensor) {
65    match tensor.dtype() {
66        DType::F32 => sort_with_indices_typed::<f32>(tensor, dim, descending, f32::total_cmp),
67        DType::F64 => sort_with_indices_typed::<f64>(tensor, dim, descending, f64::total_cmp),
68        DType::F16 => sort_with_indices_half(tensor, dim, descending, f16::to_f32, f16::from_f32),
69        DType::BF16 => {
70            sort_with_indices_half(tensor, dim, descending, bf16::to_f32, bf16::from_f32)
71        }
72        DType::I64 => sort_with_indices_typed::<i64>(tensor, dim, descending, Ord::cmp),
73        DType::I32 => sort_with_indices_typed::<i32>(tensor, dim, descending, Ord::cmp),
74        DType::I16 => sort_with_indices_typed::<i16>(tensor, dim, descending, Ord::cmp),
75        DType::I8 => sort_with_indices_typed::<i8>(tensor, dim, descending, Ord::cmp),
76        DType::U64 => sort_with_indices_typed::<u64>(tensor, dim, descending, Ord::cmp),
77        DType::U32 => sort_with_indices_typed::<u32>(tensor, dim, descending, Ord::cmp),
78        DType::U16 => sort_with_indices_typed::<u16>(tensor, dim, descending, Ord::cmp),
79        DType::U8 => sort_with_indices_typed::<u8>(tensor, dim, descending, Ord::cmp),
80        dt => panic!("sort_with_indices: unsupported dtype {:?}", dt),
81    }
82}
83
84/// Argsort along a dimension, returning indices that would sort the tensor.
85pub fn argsort(tensor: HostTensor, dim: usize, descending: bool) -> HostTensor {
86    match tensor.dtype() {
87        DType::F32 => argsort_typed::<f32>(tensor, dim, descending, f32::total_cmp),
88        DType::F64 => argsort_typed::<f64>(tensor, dim, descending, f64::total_cmp),
89        DType::F16 => argsort_half(tensor, dim, descending, f16::to_f32),
90        DType::BF16 => argsort_half(tensor, dim, descending, bf16::to_f32),
91        DType::I64 => argsort_typed::<i64>(tensor, dim, descending, Ord::cmp),
92        DType::I32 => argsort_typed::<i32>(tensor, dim, descending, Ord::cmp),
93        DType::I16 => argsort_typed::<i16>(tensor, dim, descending, Ord::cmp),
94        DType::I8 => argsort_typed::<i8>(tensor, dim, descending, Ord::cmp),
95        DType::U64 => argsort_typed::<u64>(tensor, dim, descending, Ord::cmp),
96        DType::U32 => argsort_typed::<u32>(tensor, dim, descending, Ord::cmp),
97        DType::U16 => argsort_typed::<u16>(tensor, dim, descending, Ord::cmp),
98        DType::U8 => argsort_typed::<u8>(tensor, dim, descending, Ord::cmp),
99        dt => panic!("argsort: unsupported dtype {:?}", dt),
100    }
101}
102
103// ---------------------------------------------------------------------------
104// Typed sort (operates directly on storage)
105// ---------------------------------------------------------------------------
106
107fn sort_typed<E: Element + Pod + Copy + Send>(
108    tensor: HostTensor,
109    dim: usize,
110    descending: bool,
111    cmp: fn(&E, &E) -> core::cmp::Ordering,
112) -> HostTensor {
113    let tensor = tensor.to_contiguous();
114    let shape = tensor.layout().shape().clone();
115    let dtype = tensor.dtype();
116    if validate_sort_args(&shape, dim) {
117        return tensor;
118    }
119
120    let mut data: Vec<E> = tensor.storage::<E>().to_vec();
121
122    if shape.num_dims() == 1 {
123        if descending {
124            data.sort_unstable_by(|a, b| cmp(b, a));
125        } else {
126            data.sort_unstable_by(cmp);
127        }
128    } else {
129        sort_along_dim(&mut data, &shape, dim, descending, cmp);
130    }
131
132    HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), dtype)
133}
134
135fn sort_with_indices_typed<E: Element + Pod + Copy + Send>(
136    tensor: HostTensor,
137    dim: usize,
138    descending: bool,
139    cmp: fn(&E, &E) -> core::cmp::Ordering,
140) -> (HostTensor, HostTensor) {
141    let tensor = tensor.to_contiguous();
142    let shape = tensor.layout().shape().clone();
143    let dtype = tensor.dtype();
144    let n = shape.num_elements();
145    if validate_sort_args(&shape, dim) {
146        let idx = make_index_tensor(Vec::new(), shape.clone());
147        return (tensor, idx);
148    }
149
150    let src: &[E] = tensor.storage();
151    let mut values: Vec<E> = src.to_vec();
152    let mut indices: Vec<isize> = vec![0; n];
153
154    if shape.num_dims() == 1 {
155        sort_1d_with_indices(&mut values, &mut indices, descending, cmp);
156    } else {
157        sort_along_dim_with_indices(&mut values, &mut indices, &shape, dim, descending, cmp);
158    }
159
160    let idx_tensor = make_index_tensor(indices, shape.clone());
161    let val_tensor = HostTensor::new(Bytes::from_elems(values), Layout::contiguous(shape), dtype);
162    (val_tensor, idx_tensor)
163}
164
165/// Argsort without materializing sorted values.
166fn argsort_typed<E: Element + Pod + Copy + Sync>(
167    tensor: HostTensor,
168    dim: usize,
169    descending: bool,
170    cmp: fn(&E, &E) -> core::cmp::Ordering,
171) -> HostTensor {
172    let tensor = tensor.to_contiguous();
173    let shape = tensor.layout().shape().clone();
174    let n = shape.num_elements();
175    if validate_sort_args(&shape, dim) {
176        return make_index_tensor(Vec::new(), shape);
177    }
178
179    let src: &[E] = tensor.storage();
180    let mut indices: Vec<isize> = vec![0; n];
181
182    if shape.num_dims() == 1 {
183        let mut idx_vec: Vec<usize> = (0..n).collect();
184        if descending {
185            idx_vec.sort_unstable_by(|&a, &b| cmp(&src[b], &src[a]));
186        } else {
187            idx_vec.sort_unstable_by(|&a, &b| cmp(&src[a], &src[b]));
188        }
189        for (out_i, &orig_i) in idx_vec.iter().enumerate() {
190            indices[out_i] = orig_i as isize;
191        }
192    } else {
193        argsort_along_dim(src, &mut indices, &shape, dim, descending, cmp);
194    }
195
196    make_index_tensor(indices, shape)
197}
198
199/// 1D sort with index tracking, shared by typed and half-precision paths.
200fn sort_1d_with_indices<E: Copy>(
201    values: &mut [E],
202    indices: &mut [isize],
203    descending: bool,
204    cmp: fn(&E, &E) -> core::cmp::Ordering,
205) {
206    let n = values.len();
207    let mut idx_vec: Vec<usize> = (0..n).collect();
208    if descending {
209        idx_vec.sort_unstable_by(|&a, &b| cmp(&values[b], &values[a]));
210    } else {
211        idx_vec.sort_unstable_by(|&a, &b| cmp(&values[a], &values[b]));
212    }
213    // Apply permutation in one pass using the sorted index order
214    let old_values = values.to_vec();
215    for (out_i, &orig_i) in idx_vec.iter().enumerate() {
216        values[out_i] = old_values[orig_i];
217        indices[out_i] = orig_i as isize;
218    }
219}
220
221/// Sort along a given dimension for N-D tensors.
222fn sort_along_dim<E: Copy + Send>(
223    data: &mut [E],
224    shape: &Shape,
225    dim: usize,
226    descending: bool,
227    cmp: fn(&E, &E) -> core::cmp::Ordering,
228) {
229    let strides = contiguous_strides(shape);
230    let dim_size = shape[dim];
231    let dim_stride = strides[dim];
232    let num_slices = data.len() / dim_size;
233
234    // Fast path: last dimension (stride==1). Rows are contiguous at
235    // offsets `slice_idx * dim_size`, so `chunks_exact_mut(dim_size)`
236    // walks them directly. Parallelized with rayon above the threshold.
237    if dim_stride == 1 {
238        // chunks_exact silently drops any remainder, so lock the
239        // "length is a multiple of dim_size" invariant that holds for
240        // contiguous tensors by construction.
241        debug_assert_eq!(data.len() % dim_size, 0);
242        let sort_row = |row: &mut [E]| {
243            if descending {
244                row.sort_unstable_by(|a, b| cmp(b, a));
245            } else {
246                row.sort_unstable_by(cmp);
247            }
248        };
249
250        #[cfg(feature = "rayon")]
251        if data.len() >= PARALLEL_THRESHOLD {
252            data.par_chunks_exact_mut(dim_size).for_each(sort_row);
253            return;
254        }
255
256        data.chunks_exact_mut(dim_size).for_each(sort_row);
257        return;
258    }
259
260    let mut slice_buf: Vec<E> = vec![data[0]; dim_size];
261
262    for slice_idx in 0..num_slices {
263        let base = slice_base_offset(slice_idx, shape, &strides, dim);
264
265        for i in 0..dim_size {
266            slice_buf[i] = data[base + i * dim_stride];
267        }
268
269        if descending {
270            slice_buf.sort_unstable_by(|a, b| cmp(b, a));
271        } else {
272            slice_buf.sort_unstable_by(cmp);
273        }
274
275        for i in 0..dim_size {
276            data[base + i * dim_stride] = slice_buf[i];
277        }
278    }
279}
280
281/// Sort along a dimension, tracking original indices.
282fn sort_along_dim_with_indices<E: Copy + Send>(
283    data: &mut [E],
284    indices: &mut [isize],
285    shape: &Shape,
286    dim: usize,
287    descending: bool,
288    cmp: fn(&E, &E) -> core::cmp::Ordering,
289) {
290    let strides = contiguous_strides(shape);
291    let dim_size = shape[dim];
292    let dim_stride = strides[dim];
293    let num_slices = data.len() / dim_size;
294
295    // Fast path: last dimension (stride==1). Values and indices rows
296    // are both contiguous at `slice_idx * dim_size`, so we can zip
297    // matching chunks and avoid the per-row stride arithmetic.
298    if dim_stride == 1 {
299        // zip silently truncates to the shorter iterator and chunks_exact
300        // silently drops remainders, so lock both invariants.
301        debug_assert_eq!(data.len(), indices.len());
302        debug_assert_eq!(data.len() % dim_size, 0);
303        // Buffer is reused across rows (one per thread under rayon via
304        // `for_each_init`) to avoid a heap alloc per row.
305        let sort_row = |pairs: &mut Vec<(usize, E)>, (row, idx_row): (&mut [E], &mut [isize])| {
306            pairs.clear();
307            pairs.extend((0..dim_size).map(|i| (i, row[i])));
308            if descending {
309                pairs.sort_unstable_by(|a, b| cmp(&b.1, &a.1));
310            } else {
311                pairs.sort_unstable_by(|a, b| cmp(&a.1, &b.1));
312            }
313            for (i, &(orig_idx, val)) in pairs.iter().enumerate() {
314                row[i] = val;
315                idx_row[i] = orig_idx as isize;
316            }
317        };
318
319        #[cfg(feature = "rayon")]
320        if data.len() >= PARALLEL_THRESHOLD {
321            data.par_chunks_exact_mut(dim_size)
322                .zip(indices.par_chunks_exact_mut(dim_size))
323                .for_each_init(|| Vec::with_capacity(dim_size), sort_row);
324            return;
325        }
326
327        let mut pairs: Vec<(usize, E)> = Vec::with_capacity(dim_size);
328        data.chunks_exact_mut(dim_size)
329            .zip(indices.chunks_exact_mut(dim_size))
330            .for_each(|row_and_idx| sort_row(&mut pairs, row_and_idx));
331        return;
332    }
333
334    let mut pairs: Vec<(usize, E)> = Vec::with_capacity(dim_size);
335
336    for slice_idx in 0..num_slices {
337        let base = slice_base_offset(slice_idx, shape, &strides, dim);
338
339        pairs.clear();
340        for i in 0..dim_size {
341            pairs.push((i, data[base + i * dim_stride]));
342        }
343
344        if descending {
345            pairs.sort_unstable_by(|a, b| cmp(&b.1, &a.1));
346        } else {
347            pairs.sort_unstable_by(|a, b| cmp(&a.1, &b.1));
348        }
349
350        for (i, &(orig_idx, val)) in pairs.iter().enumerate() {
351            let offset = base + i * dim_stride;
352            data[offset] = val;
353            indices[offset] = orig_idx as isize;
354        }
355    }
356}
357
358/// Argsort along a dimension without writing sorted values.
359fn argsort_along_dim<E: Copy + Sync>(
360    data: &[E],
361    indices: &mut [isize],
362    shape: &Shape,
363    dim: usize,
364    descending: bool,
365    cmp: fn(&E, &E) -> core::cmp::Ordering,
366) {
367    let strides = contiguous_strides(shape);
368    let dim_size = shape[dim];
369    let dim_stride = strides[dim];
370    let num_slices = data.len() / dim_size;
371
372    // Fast path: last dimension (stride==1). Both input rows and
373    // output index rows are contiguous at `slice_idx * dim_size`.
374    if dim_stride == 1 {
375        // zip silently truncates to the shorter iterator and chunks_exact
376        // silently drops remainders, so lock both invariants.
377        debug_assert_eq!(data.len(), indices.len());
378        debug_assert_eq!(data.len() % dim_size, 0);
379        // Buffer is reused across rows (one per thread under rayon via
380        // `for_each_init`) to avoid a heap alloc per row.
381        let sort_row = |idx_buf: &mut Vec<usize>, (row, idx_row): (&[E], &mut [isize])| {
382            idx_buf.clear();
383            idx_buf.extend(0..dim_size);
384            if descending {
385                idx_buf.sort_unstable_by(|&a, &b| cmp(&row[b], &row[a]));
386            } else {
387                idx_buf.sort_unstable_by(|&a, &b| cmp(&row[a], &row[b]));
388            }
389            for (i, &orig_idx) in idx_buf.iter().enumerate() {
390                idx_row[i] = orig_idx as isize;
391            }
392        };
393
394        #[cfg(feature = "rayon")]
395        if data.len() >= PARALLEL_THRESHOLD {
396            data.par_chunks_exact(dim_size)
397                .zip(indices.par_chunks_exact_mut(dim_size))
398                .for_each_init(|| Vec::with_capacity(dim_size), sort_row);
399            return;
400        }
401
402        let mut idx_buf: Vec<usize> = Vec::with_capacity(dim_size);
403        data.chunks_exact(dim_size)
404            .zip(indices.chunks_exact_mut(dim_size))
405            .for_each(|row_and_idx| sort_row(&mut idx_buf, row_and_idx));
406        return;
407    }
408
409    let mut idx_buf: Vec<usize> = (0..dim_size).collect();
410
411    for slice_idx in 0..num_slices {
412        let base = slice_base_offset(slice_idx, shape, &strides, dim);
413
414        idx_buf.clear();
415        idx_buf.extend(0..dim_size);
416
417        if descending {
418            idx_buf.sort_unstable_by(|&a, &b| {
419                cmp(&data[base + b * dim_stride], &data[base + a * dim_stride])
420            });
421        } else {
422            idx_buf.sort_unstable_by(|&a, &b| {
423                cmp(&data[base + a * dim_stride], &data[base + b * dim_stride])
424            });
425        }
426
427        for (i, &orig_idx) in idx_buf.iter().enumerate() {
428            indices[base + i * dim_stride] = orig_idx as isize;
429        }
430    }
431}
432
433// ---------------------------------------------------------------------------
434// Half-precision sort (convert to f32, sort, convert back)
435// ---------------------------------------------------------------------------
436
437fn sort_half<H: Element + Pod + Copy>(
438    tensor: HostTensor,
439    dim: usize,
440    descending: bool,
441    to_f32: fn(H) -> f32,
442    from_f32: fn(f32) -> H,
443) -> HostTensor {
444    let tensor = tensor.to_contiguous();
445    let shape = tensor.layout().shape().clone();
446    let dtype = tensor.dtype();
447    if validate_sort_args(&shape, dim) {
448        return tensor;
449    }
450    let src: &[H] = tensor.storage();
451    let mut f32_data: Vec<f32> = src.iter().map(|&v| to_f32(v)).collect();
452
453    if shape.num_dims() == 1 {
454        if descending {
455            f32_data.sort_unstable_by(|a, b| f32::total_cmp(b, a));
456        } else {
457            f32_data.sort_unstable_by(f32::total_cmp);
458        }
459    } else {
460        sort_along_dim(&mut f32_data, &shape, dim, descending, f32::total_cmp);
461    }
462
463    let result: Vec<H> = f32_data.iter().map(|&v| from_f32(v)).collect();
464    HostTensor::new(Bytes::from_elems(result), Layout::contiguous(shape), dtype)
465}
466
467fn sort_with_indices_half<H: Element + Pod + Copy>(
468    tensor: HostTensor,
469    dim: usize,
470    descending: bool,
471    to_f32: fn(H) -> f32,
472    from_f32: fn(f32) -> H,
473) -> (HostTensor, HostTensor) {
474    let tensor = tensor.to_contiguous();
475    let shape = tensor.layout().shape().clone();
476    let dtype = tensor.dtype();
477    let n = shape.num_elements();
478    if validate_sort_args(&shape, dim) {
479        let idx = make_index_tensor(Vec::new(), shape.clone());
480        return (tensor, idx);
481    }
482    let src: &[H] = tensor.storage();
483    let mut f32_data: Vec<f32> = src.iter().map(|&v| to_f32(v)).collect();
484    let mut indices: Vec<isize> = vec![0; n];
485
486    if shape.num_dims() == 1 {
487        sort_1d_with_indices(&mut f32_data, &mut indices, descending, f32::total_cmp);
488    } else {
489        sort_along_dim_with_indices(
490            &mut f32_data,
491            &mut indices,
492            &shape,
493            dim,
494            descending,
495            f32::total_cmp,
496        );
497    }
498
499    let result: Vec<H> = f32_data.iter().map(|&v| from_f32(v)).collect();
500    let val_tensor = HostTensor::new(
501        Bytes::from_elems(result),
502        Layout::contiguous(shape.clone()),
503        dtype,
504    );
505    let idx_tensor = make_index_tensor(indices, shape);
506    (val_tensor, idx_tensor)
507}
508
509fn argsort_half<H: Element + Pod + Copy>(
510    tensor: HostTensor,
511    dim: usize,
512    descending: bool,
513    to_f32: fn(H) -> f32,
514) -> HostTensor {
515    let tensor = tensor.to_contiguous();
516    let shape = tensor.layout().shape().clone();
517    let n = shape.num_elements();
518    if validate_sort_args(&shape, dim) {
519        return make_index_tensor(Vec::new(), shape);
520    }
521    let src: &[H] = tensor.storage();
522    let f32_data: Vec<f32> = src.iter().map(|&v| to_f32(v)).collect();
523    let mut indices: Vec<isize> = vec![0; n];
524
525    if shape.num_dims() == 1 {
526        let mut idx_vec: Vec<usize> = (0..n).collect();
527        if descending {
528            idx_vec.sort_unstable_by(|&a, &b| f32::total_cmp(&f32_data[b], &f32_data[a]));
529        } else {
530            idx_vec.sort_unstable_by(|&a, &b| f32::total_cmp(&f32_data[a], &f32_data[b]));
531        }
532        for (out_i, &orig_i) in idx_vec.iter().enumerate() {
533            indices[out_i] = orig_i as isize;
534        }
535    } else {
536        argsort_along_dim(
537            &f32_data,
538            &mut indices,
539            &shape,
540            dim,
541            descending,
542            f32::total_cmp,
543        );
544    }
545
546    make_index_tensor(indices, shape)
547}
548
549// ---------------------------------------------------------------------------
550// Helpers
551// ---------------------------------------------------------------------------
552
553fn contiguous_strides(shape: &Shape) -> Vec<usize> {
554    ruda_core::tensor::host::layout::contiguous_strides_usize(shape)
555}
556
557fn slice_base_offset(slice_idx: usize, shape: &Shape, strides: &[usize], dim: usize) -> usize {
558    ruda_core::tensor::host::layout::slice_base_offset(slice_idx, shape, strides, dim)
559}
560
561fn make_index_tensor(indices: Vec<isize>, shape: Shape) -> HostTensor {
562    let bytes = Bytes::from_elems(indices);
563    HostTensor::new(bytes, Layout::contiguous(shape), INDEX_DTYPE)
564}
565
566// Tests kept here exercise flex-specific behavior: the internal
567// `sort_along_dim` / `sort_along_dim_with_indices` / `argsort_along_dim`
568// helpers at sizes that straddle `PARALLEL_THRESHOLD`, so both the serial
569// and rayon-parallel branches are covered. End-to-end sort/argsort
570// correctness across backends lives in
571// crates/ruda-backend-tests/tests/tensor/float/ops/sort_argsort.rs.
572#[cfg(test)]
573mod tests;
574
575pub mod dispatch;
576mod topk;
577pub use topk::argtopk;