Skip to main content

hanzo_ml/
sort.rs

1use crate::{Result, Tensor};
2use rayon::prelude::*;
3
4#[derive(Debug, Clone, Copy)]
5struct ArgSort {
6    asc: bool,
7    last_dim: usize,
8}
9
10impl ArgSort {
11    fn asort<T: crate::WithDType>(&self, vs: &[T], layout: &crate::Layout) -> Result<Vec<u32>> {
12        let vs = match layout.contiguous_offsets() {
13            None => crate::bail!("input has to be contiguous"),
14            Some((o1, o2)) => &vs[o1..o2],
15        };
16        #[allow(clippy::uninit_vec)]
17        // Safety: indexes are set later in the parallelized section.
18        let mut sort_indexes = unsafe {
19            let el_count = layout.shape().elem_count();
20            let mut v = Vec::with_capacity(el_count);
21            v.set_len(el_count);
22            v
23        };
24        if self.asc {
25            sort_indexes
26                .par_chunks_exact_mut(self.last_dim)
27                .zip(vs.par_chunks_exact(self.last_dim))
28                .for_each(|(indexes, vs)| {
29                    indexes
30                        .iter_mut()
31                        .enumerate()
32                        .for_each(|(i, v)| *v = i as u32);
33                    indexes.sort_by(|&i, &j| {
34                        vs[i as usize]
35                            .partial_cmp(&vs[j as usize])
36                            .unwrap_or(std::cmp::Ordering::Greater)
37                    })
38                });
39        } else {
40            sort_indexes
41                .par_chunks_exact_mut(self.last_dim)
42                .zip(vs.par_chunks_exact(self.last_dim))
43                .for_each(|(indexes, vs)| {
44                    indexes
45                        .iter_mut()
46                        .enumerate()
47                        .for_each(|(i, v)| *v = i as u32);
48                    indexes.sort_by(|&j, &i| {
49                        vs[i as usize]
50                            .partial_cmp(&vs[j as usize])
51                            .unwrap_or(std::cmp::Ordering::Greater)
52                    })
53                });
54        }
55        Ok(sort_indexes)
56    }
57}
58
59#[cfg(feature = "cuda")]
60mod cuda {
61    use super::*;
62    use crate::cuda_backend::cudarc::driver::{
63        CudaSlice, DeviceRepr, LaunchConfig, ValidAsZeroBits,
64    };
65    use crate::cuda_backend::{kernel_name, kernels, CudaStorageSlice as S, WrapErr};
66    use crate::{CudaDevice, WithDType};
67
68    impl crate::cuda_backend::Map1Any for ArgSort {
69        fn f<T: DeviceRepr + WithDType + ValidAsZeroBits, W: Fn(CudaSlice<T>) -> S>(
70            &self,
71            src: &CudaSlice<T>,
72            dev: &CudaDevice,
73            layout: &crate::Layout,
74            _wrap: W,
75        ) -> Result<S> {
76            use cudarc::driver::PushKernelArg;
77
78            let slice = match layout.contiguous_offsets() {
79                None => crate::bail!("input has to be contiguous"),
80                Some((o1, o2)) => src.slice(o1..o2),
81            };
82            let elem_count = layout.shape().elem_count();
83            let dst = unsafe { dev.alloc::<u32>(elem_count)? };
84            let func = if self.asc {
85                dev.get_or_load_func(&kernel_name::<T>("asort_asc"), &kernels::SORT)?
86            } else {
87                dev.get_or_load_func(&kernel_name::<T>("asort_desc"), &kernels::SORT)?
88            };
89            let ncols = self.last_dim;
90            let nrows = elem_count / ncols;
91            let ncols_pad = next_power_of_2(ncols);
92            // Limit block dim to 1024 threads, which is the maximum on modern CUDA gpus.
93            let block_dim = ncols_pad.min(1024);
94            let cfg = LaunchConfig {
95                grid_dim: (nrows as u32, 1, 1),
96                block_dim: (block_dim as u32, 1, 1),
97                shared_mem_bytes: (ncols_pad * std::mem::size_of::<u32>()) as u32,
98            };
99            let stream = dev.cuda_stream();
100            let mut builder = stream.launch_builder(&func);
101            let ncols = ncols as i32;
102            let ncols_pad = ncols_pad as i32;
103            builder.arg(&slice).arg(&dst).arg(&ncols).arg(&ncols_pad);
104            unsafe { builder.launch(cfg) }.w()?;
105            Ok(S::U32(dst))
106        }
107    }
108}
109
110impl crate::CustomOp1 for ArgSort {
111    fn name(&self) -> &'static str {
112        "argsort"
113    }
114
115    fn cpu_fwd(
116        &self,
117        storage: &crate::CpuStorage,
118        layout: &crate::Layout,
119    ) -> Result<(crate::CpuStorage, crate::Shape)> {
120        let sort_indexes = match storage {
121            crate::CpuStorage::U8(vs) => self.asort(vs, layout)?,
122            crate::CpuStorage::U32(vs) => self.asort(vs, layout)?,
123            crate::CpuStorage::I16(vs) => self.asort(vs, layout)?,
124            crate::CpuStorage::I32(vs) => self.asort(vs, layout)?,
125            crate::CpuStorage::I64(vs) => self.asort(vs, layout)?,
126            crate::CpuStorage::BF16(vs) => self.asort(vs, layout)?,
127            crate::CpuStorage::F16(vs) => self.asort(vs, layout)?,
128            crate::CpuStorage::F32(vs) => self.asort(vs, layout)?,
129            crate::CpuStorage::F64(vs) => self.asort(vs, layout)?,
130            crate::CpuStorage::F8E4M3(vs) => self.asort(vs, layout)?,
131            // Dummy types don't support sorting
132            crate::CpuStorage::F6E2M3(_) => {
133                return Err(
134                    crate::Error::UnsupportedDTypeForOp(crate::DType::F6E2M3, "argsort").bt(),
135                )
136            }
137            crate::CpuStorage::F6E3M2(_) => {
138                return Err(
139                    crate::Error::UnsupportedDTypeForOp(crate::DType::F6E3M2, "argsort").bt(),
140                )
141            }
142            crate::CpuStorage::F4(_) => {
143                return Err(crate::Error::UnsupportedDTypeForOp(crate::DType::F4, "argsort").bt())
144            }
145            crate::CpuStorage::F8E8M0(_) => {
146                return Err(
147                    crate::Error::UnsupportedDTypeForOp(crate::DType::F8E8M0, "argsort").bt(),
148                )
149            }
150        };
151        let sort_indexes = crate::CpuStorage::U32(sort_indexes);
152        Ok((sort_indexes, layout.shape().into()))
153    }
154
155    #[cfg(feature = "cuda")]
156    fn cuda_fwd(
157        &self,
158        storage: &crate::CudaStorage,
159        layout: &crate::Layout,
160    ) -> Result<(crate::CudaStorage, crate::Shape)> {
161        use crate::backend::BackendStorage;
162        use crate::cuda_backend::Map1Any;
163        let dev = storage.device();
164        let slice = self.map(&storage.slice, dev, layout)?;
165        let dst = crate::cuda_backend::CudaStorage {
166            slice,
167            device: dev.clone(),
168        };
169        Ok((dst, layout.shape().clone()))
170    }
171
172    #[cfg(feature = "metal")]
173    fn metal_fwd(
174        &self,
175        storage: &crate::MetalStorage,
176        layout: &crate::Layout,
177    ) -> Result<(crate::MetalStorage, crate::Shape)> {
178        use crate::backend::BackendStorage;
179        use crate::DType;
180
181        let name = {
182            if self.asc {
183                match storage.dtype() {
184                    DType::BF16 => "asort_asc_bf16",
185                    DType::F16 => "asort_asc_f16",
186                    DType::F32 => "asort_asc_f32",
187                    DType::F64 => "asort_asc_f64",
188                    DType::U8 => "asort_asc_u8",
189                    DType::U32 => "asort_asc_u32",
190                    DType::I16 => "asort_asc_i16",
191                    DType::I32 => "asort_asc_i32",
192                    DType::I64 => "asort_asc_i64",
193                    DType::F8E4M3 => crate::bail!("Metal device does not yet support F8E4M3."),
194                    DType::F6E2M3 | DType::F6E3M2 | DType::F4 | DType::F8E8M0 => {
195                        return Err(
196                            crate::Error::UnsupportedDTypeForOp(storage.dtype(), "argsort").bt(),
197                        )
198                    }
199                }
200            } else {
201                match storage.dtype() {
202                    DType::BF16 => "asort_desc_bf16",
203                    DType::F16 => "asort_desc_f16",
204                    DType::F32 => "asort_desc_f32",
205                    DType::F64 => "asort_desc_f64",
206                    DType::U8 => "asort_desc_u8",
207                    DType::U32 => "asort_desc_u32",
208                    DType::I16 => "asort_desc_i16",
209                    DType::I32 => "asort_desc_i32",
210                    DType::I64 => "asort_desc_i64",
211                    DType::F8E4M3 => crate::bail!("Metal device does not yet support F8E4M3."),
212                    DType::F6E2M3 | DType::F6E3M2 | DType::F4 | DType::F8E8M0 => {
213                        return Err(
214                            crate::Error::UnsupportedDTypeForOp(storage.dtype(), "argsort").bt(),
215                        )
216                    }
217                }
218            }
219        };
220        let device = storage.device();
221        let kernels = device.kernels();
222        let command_encoder = device.command_encoder()?;
223        let el = layout.shape().elem_count();
224        let ncols = self.last_dim;
225        let nrows = el / ncols;
226        let src = crate::metal_backend::buffer_o(storage.buffer(), layout, storage.dtype());
227        let dst = device
228            .new_buffer_builder()
229            .with_size_for(el, DType::U32)
230            .with_label("asort")
231            .build()?;
232        let mut ncols_pad = 1;
233        while ncols_pad < ncols {
234            ncols_pad *= 2;
235        }
236        hanzo_metal_kernels::call_arg_sort(
237            device.metal_device(),
238            &command_encoder,
239            kernels,
240            name,
241            nrows,
242            ncols,
243            ncols_pad,
244            src,
245            &dst,
246        )
247        .map_err(crate::Error::wrap)?;
248        let dst = crate::MetalStorage::new(dst, device.clone(), el, DType::U32);
249        Ok((dst, layout.shape().clone()))
250    }
251
252    #[cfg(feature = "rocm")]
253    fn rocm_fwd(
254        &self,
255        storage: &crate::RocmStorage,
256        layout: &crate::Layout,
257    ) -> Result<(crate::RocmStorage, crate::Shape)> {
258        let dst = storage.asort(layout, self.asc, self.last_dim)?;
259        Ok((dst, layout.shape().clone()))
260    }
261
262    #[cfg(feature = "vulkan")]
263    fn vulkan_fwd(
264        &self,
265        storage: &crate::VulkanStorage,
266        layout: &crate::Layout,
267    ) -> Result<(crate::VulkanStorage, crate::Shape)> {
268        use crate::backend::{BackendDevice, BackendStorage};
269        // The bitonic argsort shader reads f32 rows and stays on the GPU while the padded row
270        // width fits its shared-index scratch (ARGSORT_MAX_COLS_PAD) -- the MoE routing case
271        // (cols == num_experts) always does. Any other dtype, or an over-wide row, takes the
272        // dtype-generic CPU argsort and re-uploads the u32 permutation to the device.
273        if storage.dtype() == crate::DType::F32 {
274            if let Some(dst) = storage.arg_sort_last_dim(layout, self.asc, self.last_dim)? {
275                return Ok((dst, layout.shape().clone()));
276            }
277        }
278        let (idx, shape) = self.cpu_fwd(&storage.to_cpu_storage()?, layout)?;
279        Ok((storage.device().storage_from_cpu_storage(&idx)?, shape))
280    }
281}
282
283#[allow(unused)]
284fn next_power_of_2(x: usize) -> usize {
285    let mut n = 1;
286    while n < x {
287        n *= 2
288    }
289    n
290}
291
292impl Tensor {
293    /// Returns the indices that sort the tensor along the last dimension.
294    ///
295    /// If `asc` is `true`, sorting is in ascending order. Otherwise sorting is performed in
296    /// descending order. The sort is unstable so there is no guarantees on the final order when it
297    /// comes to ties.
298    pub fn arg_sort_last_dim(&self, asc: bool) -> Result<Tensor> {
299        if !self.is_contiguous() {
300            return Err(crate::Error::RequiresContiguous {
301                op: "arg_sort_last_dim",
302            });
303        }
304        let last_dim = match self.dims().last() {
305            None => crate::bail!("empty last-dim in arg-sort"),
306            Some(last_dim) => *last_dim,
307        };
308        // No need for a backward pass for arg sort.
309        self.apply_op1_no_bwd(&ArgSort { asc, last_dim })
310    }
311
312    /// Sorts the tensor along the last dimension, returns the sorted tensor together with the
313    /// sorted indexes.
314    ///
315    /// If `asc` is `true`, sorting is in ascending order. Otherwise sorting is performed in
316    /// descending order. The sort is unstable so there is no guarantees on the final order when it
317    /// comes to ties.
318    pub fn sort_last_dim(&self, asc: bool) -> Result<(Tensor, Tensor)> {
319        if !self.is_contiguous() {
320            return Err(crate::Error::RequiresContiguous {
321                op: "sort_last_dim",
322            });
323        }
324        let asort = self.arg_sort_last_dim(asc)?;
325        let sorted = self.gather(&asort, crate::D::Minus1)?;
326        Ok((sorted, asort))
327    }
328}