1use 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
21fn 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
40pub 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
59pub 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
84pub 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
103fn 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
165fn 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
199fn 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 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
221fn 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 if dim_stride == 1 {
238 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
281fn 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 if dim_stride == 1 {
299 debug_assert_eq!(data.len(), indices.len());
302 debug_assert_eq!(data.len() % dim_size, 0);
303 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
358fn 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 if dim_stride == 1 {
375 debug_assert_eq!(data.len(), indices.len());
378 debug_assert_eq!(data.len() % dim_size, 0);
379 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
433fn 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
549fn 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#[cfg(test)]
573mod tests;
574
575pub mod dispatch;
576mod topk;
577pub use topk::argtopk;