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 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 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 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 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 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 self.apply_op1_no_bwd(&ArgSort { asc, last_dim })
310 }
311
312 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}