Skip to main content

ruda_tensor_device/dispatch/
integer.rs

1use self::unary_basic_int::BasicIntUnaryKind;
2use super::{expand, numeric, permute, unfold};
3use crate::{DeviceBackend, DeviceRuntime, FloatElement, IntElement, element::BoolElement};
4use ruda_kernel::tensor::unary_numeric::{NumericUnaryOp, NumericUnaryOpFamily, launch_unary_numeric};
5use ruprim::elementwise::binary::int::{BitwiseShlOp, BitwiseShrOp, launch_binop_int, launch_scalar_binop_int};
6use ruprim::elementwise::unary::int::unary_basic_int;
7use ruprim::reduce::tensor as reduce;
8use rublas::tensor_matmul::{MatmulStrategy, matmul};
9use rurand::tensor::{random_bernoulli, random_normal, random_uniform};
10use ruda_tensor::tensor::{BoolTensor, Device, FloatTensor, IntTensor};
11use ruda_tensor::{DType, IntDType, Slice, ops::IntTensorOps};
12use ruda_tensor::{Distribution, ElementConversion, Shape, TensorData, get_device_settings};
13use ruda_tensor::{ExecutionError, Scalar};
14use ruda_core::tensor::{BoolDType, FloatDType};
15use ruda_kernel::dsl::frontend::Numeric;
16use ruda_kernel::dsl::{self as ruda, prelude::*};
17use ruprim::reduce::components::instructions::ReduceOperationConfig;
18use std::ops::Range;
19
20impl<R, F, I, BT> IntTensorOps<Self> for DeviceBackend<R, F, I, BT>
21where
22    R: DeviceRuntime,
23    F: FloatElement,
24    I: IntElement,
25    BT: BoolElement,
26{
27    fn int_empty(shape: Shape, device: &Device<Self>, dtype: IntDType) -> IntTensor<Self> {
28        let dtype = dtype.into();
29        super::empty(shape, device, dtype)
30    }
31
32    async fn int_into_data(tensor: IntTensor<Self>) -> Result<TensorData, ExecutionError> {
33        super::into_data(tensor).await
34    }
35
36    fn int_from_data(data: TensorData, device: &Device<Self>) -> IntTensor<Self> {
37        match data.dtype {
38            DType::I64
39            | DType::I32
40            | DType::I16
41            | DType::I8
42            | DType::U64
43            | DType::U32
44            | DType::U16
45            | DType::U8 => super::from_data(data, device),
46            _ => unimplemented!("Unsupported dtype for `int_from_data`"),
47        }
48    }
49
50    fn int_device(tensor: &IntTensor<Self>) -> Device<Self> {
51        tensor.device.clone()
52    }
53
54    fn int_to_device(tensor: IntTensor<Self>, device: &Device<Self>) -> IntTensor<Self> {
55        super::to_device(tensor, device)
56    }
57
58    fn int_reshape(tensor: IntTensor<Self>, shape: Shape) -> IntTensor<Self> {
59        super::reshape(tensor, shape)
60    }
61
62    fn int_slice(tensor: IntTensor<Self>, slices: &[Slice]) -> IntTensor<Self> {
63        // Check if all steps are 1
64        let all_steps_one = slices.iter().all(|info| info.step == 1);
65
66        if all_steps_one {
67            // Use optimized slice for step=1
68            let simple_ranges: Vec<Range<usize>> = slices
69                .iter()
70                .enumerate()
71                .map(|(i, slice)| slice.to_range(tensor.meta.shape()[i]))
72                .collect();
73
74            ruprim::indexing::slice(tensor, &simple_ranges)
75        } else {
76            // Use slice with steps kernel
77            ruprim::indexing::slice_with_steps(tensor, slices)
78        }
79    }
80
81    fn int_slice_assign(
82        tensor: IntTensor<Self>,
83        ranges: &[Slice],
84        value: IntTensor<Self>,
85    ) -> IntTensor<Self> {
86        ruprim::indexing::slice_assign(tensor, ranges, value)
87    }
88
89    fn int_matmul(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
90        let dtype = lhs.dtype;
91        matmul(lhs, rhs, None, MatmulStrategy::default(), dtype).unwrap()
92    }
93
94    fn int_mask_where(
95        tensor: IntTensor<Self>,
96        mask: BoolTensor<Self>,
97        value: IntTensor<Self>,
98    ) -> IntTensor<Self> {
99        let bool_dtype = mask.dtype;
100        ruprim::elementwise::mask::mask_where_auto(tensor, mask, value, bool_dtype)
101    }
102
103    fn int_mask_fill(
104        tensor: IntTensor<Self>,
105        mask: BoolTensor<Self>,
106        value: Scalar,
107    ) -> IntTensor<Self> {
108        let dtype = tensor.dtype;
109        let bool_dtype = mask.dtype;
110        ruprim::elementwise::mask::mask_fill_auto(tensor, mask, InputScalar::new(value, dtype), bool_dtype)
111    }
112
113    fn int_gather(
114        dim: usize,
115        tensor: IntTensor<Self>,
116        indices: IntTensor<Self>,
117    ) -> IntTensor<Self> {
118        ruprim::indexing::gather(dim, tensor, indices)
119    }
120
121    fn int_scatter_add(
122        dim: usize,
123        tensor: IntTensor<Self>,
124        indices: IntTensor<Self>,
125        value: IntTensor<Self>,
126    ) -> IntTensor<Self> {
127        ruprim::indexing::scatter(dim, tensor, indices, value, false)
128    }
129
130    fn int_scatter_nd(
131        data: IntTensor<Self>,
132        indices: IntTensor<Self>,
133        values: IntTensor<Self>,
134        reduction: ruda_tensor::tensor::IndexingUpdateOp,
135    ) -> IntTensor<Self> {
136        ruprim::indexing::scatter_nd(data, indices, values, reduction)
137    }
138
139    fn int_gather_nd(data: IntTensor<Self>, indices: IntTensor<Self>) -> IntTensor<Self> {
140        ruprim::indexing::gather_nd(data, indices)
141    }
142
143    fn int_select(
144        tensor: IntTensor<Self>,
145        dim: usize,
146        indices: IntTensor<Self>,
147    ) -> IntTensor<Self> {
148        ruprim::indexing::select(tensor, dim, indices)
149    }
150
151    fn int_select_add(
152        tensor: IntTensor<Self>,
153        dim: usize,
154        indices: IntTensor<Self>,
155        value: IntTensor<Self>,
156    ) -> IntTensor<Self> {
157        ruprim::indexing::select_assign(tensor, dim, indices, value, false)
158    }
159
160    fn int_equal(
161        lhs: IntTensor<Self>,
162        rhs: IntTensor<Self>,
163        out_dtype: BoolDType,
164    ) -> BoolTensor<Self> {
165        ruprim::elementwise::comparison::equal(lhs, rhs, out_dtype.into())
166    }
167
168    fn int_equal_elem(lhs: IntTensor<Self>, rhs: Scalar, out_dtype: BoolDType) -> BoolTensor<Self> {
169        let dtype = lhs.dtype;
170        ruprim::elementwise::comparison::equal_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
171    }
172
173    fn int_greater(
174        lhs: IntTensor<Self>,
175        rhs: IntTensor<Self>,
176        out_dtype: BoolDType,
177    ) -> BoolTensor<Self> {
178        ruprim::elementwise::comparison::greater(lhs, rhs, out_dtype.into())
179    }
180
181    fn int_greater_elem(
182        lhs: IntTensor<Self>,
183        rhs: Scalar,
184        out_dtype: BoolDType,
185    ) -> BoolTensor<Self> {
186        let dtype = lhs.dtype;
187        ruprim::elementwise::comparison::greater_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
188    }
189
190    fn int_greater_equal(
191        lhs: IntTensor<Self>,
192        rhs: IntTensor<Self>,
193        out_dtype: BoolDType,
194    ) -> BoolTensor<Self> {
195        ruprim::elementwise::comparison::greater_equal(lhs, rhs, out_dtype.into())
196    }
197
198    fn int_greater_equal_elem(
199        lhs: IntTensor<Self>,
200        rhs: Scalar,
201        out_dtype: BoolDType,
202    ) -> BoolTensor<Self> {
203        let dtype = lhs.dtype;
204        ruprim::elementwise::comparison::greater_equal_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
205    }
206
207    fn int_lower(
208        lhs: IntTensor<Self>,
209        rhs: IntTensor<Self>,
210        out_dtype: BoolDType,
211    ) -> BoolTensor<Self> {
212        ruprim::elementwise::comparison::lower(lhs, rhs, out_dtype.into())
213    }
214
215    fn int_lower_elem(lhs: IntTensor<Self>, rhs: Scalar, out_dtype: BoolDType) -> BoolTensor<Self> {
216        let dtype = lhs.dtype;
217        ruprim::elementwise::comparison::lower_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
218    }
219
220    fn int_lower_equal(
221        lhs: IntTensor<Self>,
222        rhs: IntTensor<Self>,
223        out_dtype: BoolDType,
224    ) -> BoolTensor<Self> {
225        ruprim::elementwise::comparison::lower_equal(lhs, rhs, out_dtype.into())
226    }
227
228    fn int_lower_equal_elem(
229        lhs: IntTensor<Self>,
230        rhs: Scalar,
231        out_dtype: BoolDType,
232    ) -> BoolTensor<Self> {
233        let dtype = lhs.dtype;
234        ruprim::elementwise::comparison::lower_equal_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
235    }
236
237    fn int_add(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
238        numeric::add(lhs, rhs)
239    }
240
241    fn int_add_scalar(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
242        let dtype = lhs.dtype;
243        numeric::add_scalar(lhs, InputScalar::new(rhs, dtype))
244    }
245
246    fn int_sub(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
247        numeric::sub(lhs, rhs)
248    }
249
250    fn int_sub_scalar(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
251        let dtype = lhs.dtype;
252        numeric::sub_scalar(lhs, InputScalar::new(rhs, dtype))
253    }
254
255    fn int_mul(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
256        numeric::mul(lhs, rhs)
257    }
258
259    fn int_mul_scalar(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
260        let dtype = lhs.dtype;
261        numeric::mul_scalar(lhs, InputScalar::new(rhs, dtype))
262    }
263
264    fn int_div(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
265        numeric::div(lhs, rhs)
266    }
267
268    fn int_div_scalar(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
269        let dtype = lhs.dtype;
270        numeric::div_scalar(lhs, InputScalar::new(rhs, dtype))
271    }
272
273    fn int_remainder(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
274        numeric::remainder(lhs, rhs)
275    }
276
277    fn int_remainder_scalar(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
278        let dtype = lhs.dtype;
279        numeric::remainder_scalar(lhs, InputScalar::new(rhs, dtype))
280    }
281
282    fn int_zeros(shape: Shape, device: &Device<Self>, dtype: IntDType) -> IntTensor<Self> {
283        let dtype = dtype.into();
284        numeric::zeros(device.clone(), shape, dtype)
285    }
286
287    fn int_ones(shape: Shape, device: &Device<Self>, dtype: IntDType) -> IntTensor<Self> {
288        let dtype = dtype.into();
289        numeric::ones(device.clone(), shape, dtype)
290    }
291
292    fn int_full(
293        shape: Shape,
294        fill_value: Scalar,
295        device: &Device<Self>,
296        dtype: IntDType,
297    ) -> IntTensor<Self> {
298        let dtype: DType = dtype.into();
299        let client = R::client(device);
300        numeric::full_device_dtype(
301            client,
302            shape,
303            device.clone(),
304            InputScalar::new(fill_value, dtype),
305            dtype,
306        )
307    }
308
309    fn int_sum(tensor: IntTensor<Self>) -> IntTensor<Self> {
310        reduce::sum_fallback(tensor, Default::default()).unwrap()
311    }
312
313    fn int_sum_dim(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
314        reduce::reduce_dim(
315            tensor,
316            None,
317            dim,
318            Default::default(),
319            ReduceOperationConfig::Sum,
320        )
321        .unwrap()
322    }
323
324    fn int_prod(tensor: IntTensor<Self>) -> IntTensor<Self> {
325        reduce::reduce(
326            tensor,
327            None,
328            Default::default(),
329            ReduceOperationConfig::Prod,
330        )
331        .unwrap()
332    }
333
334    fn int_prod_dim(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
335        reduce::reduce_dim(
336            tensor,
337            None,
338            dim,
339            Default::default(),
340            ReduceOperationConfig::Prod,
341        )
342        .unwrap()
343    }
344
345    fn int_max(tensor: IntTensor<Self>) -> IntTensor<Self> {
346        reduce::reduce(tensor, None, Default::default(), ReduceOperationConfig::Max).unwrap()
347    }
348
349    fn int_max_dim(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
350        reduce::reduce_dim(
351            tensor,
352            None,
353            dim,
354            Default::default(),
355            ReduceOperationConfig::Max,
356        )
357        .unwrap()
358    }
359
360    fn int_topk(tensor: IntTensor<Self>, dim: usize, k: usize) -> IntTensor<Self> {
361        reduce::reduce_dim(
362            tensor,
363            None,
364            dim,
365            Default::default(),
366            ReduceOperationConfig::TopK(k),
367        )
368        .unwrap()
369    }
370
371    fn int_max_abs(tensor: IntTensor<Self>) -> IntTensor<Self> {
372        reduce::reduce(
373            tensor,
374            None,
375            Default::default(),
376            ReduceOperationConfig::MaxAbs,
377        )
378        .unwrap()
379    }
380
381    fn int_max_abs_dim(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
382        reduce::reduce_dim(
383            tensor,
384            None,
385            dim,
386            Default::default(),
387            ReduceOperationConfig::MaxAbs,
388        )
389        .unwrap()
390    }
391
392    fn int_min(tensor: IntTensor<Self>) -> IntTensor<Self> {
393        reduce::reduce(tensor, None, Default::default(), ReduceOperationConfig::Min).unwrap()
394    }
395
396    fn int_min_dim(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
397        reduce::reduce_dim(
398            tensor,
399            None,
400            dim,
401            Default::default(),
402            ReduceOperationConfig::Min,
403        )
404        .unwrap()
405    }
406
407    fn int_mean_dim(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
408        reduce::reduce_dim(
409            tensor,
410            None,
411            dim,
412            Default::default(),
413            ReduceOperationConfig::Mean,
414        )
415        .unwrap()
416    }
417
418    fn int_cumsum(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
419        numeric::cumsum(tensor, dim)
420    }
421
422    fn int_cumprod(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
423        numeric::cumprod(tensor, dim)
424    }
425
426    fn int_cummin(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
427        numeric::cummin(tensor, dim)
428    }
429
430    fn int_cummax(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
431        numeric::cummax(tensor, dim)
432    }
433
434    fn int_argmax(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
435        let dtype = tensor.dtype;
436        reduce::reduce_dim(
437            tensor,
438            Some(dtype),
439            dim,
440            Default::default(),
441            ReduceOperationConfig::ArgMax,
442        )
443        .unwrap()
444    }
445
446    fn int_argtopk(tensor: IntTensor<Self>, dim: usize, k: usize) -> IntTensor<Self> {
447        let dtype = tensor.dtype;
448        reduce::reduce_dim(
449            tensor,
450            Some(dtype),
451            dim,
452            Default::default(),
453            ReduceOperationConfig::ArgTopK(k),
454        )
455        .unwrap()
456    }
457
458    fn int_argmin(tensor: IntTensor<Self>, dim: usize) -> IntTensor<Self> {
459        let dtype = tensor.dtype;
460        reduce::reduce_dim(
461            tensor,
462            Some(dtype),
463            dim,
464            Default::default(),
465            ReduceOperationConfig::ArgMin,
466        )
467        .unwrap()
468    }
469
470    fn int_clamp(tensor: IntTensor<Self>, min: Scalar, max: Scalar) -> IntTensor<Self> {
471        let dtype = tensor.dtype;
472        ruprim::elementwise::unary::clamp::clamp(
473            tensor,
474            InputScalar::new(min, dtype),
475            InputScalar::new(max, dtype),
476        )
477    }
478
479    fn int_abs(tensor: IntTensor<Self>) -> IntTensor<Self> {
480        struct Abs;
481
482        #[ruda]
483        impl<T: Numeric, N: Size> NumericUnaryOp<T, N> for Abs {
484            type Options = ();
485
486            fn execute(input: Vector<T, N>, _options: &Self::Options) -> Vector<T, N> {
487                Vector::abs(input)
488            }
489        }
490
491        impl NumericUnaryOpFamily for Abs {
492            type Options = ();
493            type Unary<T: Numeric, N: Size> = Self;
494        }
495
496        launch_unary_numeric::<R, Abs, _>(tensor, |_| ())
497    }
498
499    fn int_sign(tensor: IntTensor<Self>) -> IntTensor<Self> {
500        unary_basic_int::launch::<R, _>(tensor, |_| BasicIntUnaryKind::Sign)
501    }
502
503    fn int_into_float(tensor: IntTensor<Self>, out_dtype: FloatDType) -> FloatTensor<Self> {
504        ruprim::elementwise::cast::cast(tensor, out_dtype.into())
505    }
506
507    fn int_swap_dims(mut tensor: IntTensor<Self>, dim1: usize, dim2: usize) -> IntTensor<Self> {
508        tensor.meta.swap(dim1, dim2);
509
510        tensor
511    }
512
513    fn int_repeat_dim(tensor: IntTensor<Self>, dim: usize, times: usize) -> IntTensor<Self> {
514        ruprim::indexing::repeat_dim(tensor, dim, times)
515    }
516
517    fn int_random(
518        shape: Shape,
519        distribution: Distribution,
520        device: &Device<Self>,
521        dtype: IntDType,
522    ) -> IntTensor<Self> {
523        let dtype = dtype.into();
524        match distribution {
525            Distribution::Default => random_uniform(shape, device, 0., 255., dtype),
526            Distribution::Uniform(low, high) => {
527                random_uniform(shape, device, low.elem(), high.elem(), dtype)
528            }
529            Distribution::Bernoulli(prob) => random_bernoulli(shape, device, prob as f32, dtype),
530            Distribution::Normal(mean, std) => {
531                random_normal(shape, device, mean.elem(), std.elem(), dtype)
532            }
533        }
534    }
535
536    fn int_permute(tensor: IntTensor<Self>, axes: &[usize]) -> IntTensor<Self> {
537        permute(tensor, axes)
538    }
539
540    fn int_expand(tensor: IntTensor<Self>, shape: Shape) -> IntTensor<Self> {
541        expand(tensor, shape)
542    }
543
544    fn int_flip(tensor: IntTensor<Self>, axes: &[usize]) -> IntTensor<Self> {
545        let bool_dtype = get_device_settings::<Self>(&tensor.device).bool_dtype;
546        ruprim::indexing::flip(tensor, axes, bool_dtype.into())
547    }
548
549    fn bitwise_and(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
550        numeric::bitwise_and(lhs, rhs)
551    }
552
553    fn bitwise_and_scalar(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
554        let dtype = lhs.dtype;
555        numeric::bitwise_and_scalar(lhs, InputScalar::new(rhs, dtype))
556    }
557
558    fn bitwise_or(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
559        numeric::bitwise_or(lhs, rhs)
560    }
561
562    fn bitwise_or_scalar(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
563        let dtype = lhs.dtype;
564        numeric::bitwise_or_scalar(lhs, InputScalar::new(rhs, dtype))
565    }
566
567    fn bitwise_xor(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
568        numeric::bitwise_xor(lhs, rhs)
569    }
570
571    fn bitwise_xor_scalar(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
572        let dtype = lhs.dtype;
573        numeric::bitwise_xor_scalar(lhs, InputScalar::new(rhs, dtype))
574    }
575
576    fn bitwise_not(tensor: IntTensor<Self>) -> IntTensor<Self> {
577        unary_basic_int::launch::<R, _>(tensor, |_| BasicIntUnaryKind::BitwiseNot)
578    }
579
580    fn bitwise_left_shift(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
581        launch_binop_int::<R, ruprim::elementwise::binary::int::BitwiseShlOp>(lhs, rhs)
582    }
583
584    fn bitwise_left_shift_scalar(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
585        let dtype = lhs.dtype;
586        launch_scalar_binop_int::<R, BitwiseShlOp>(lhs, InputScalar::new(rhs, dtype))
587    }
588
589    fn bitwise_right_shift(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
590        launch_binop_int::<R, BitwiseShrOp>(lhs, rhs)
591    }
592
593    fn bitwise_right_shift_scalar(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
594        let dtype = lhs.dtype;
595        launch_scalar_binop_int::<R, BitwiseShrOp>(lhs, InputScalar::new(rhs, dtype))
596    }
597
598    fn int_cast(tensor: IntTensor<Self>, dtype: IntDType) -> IntTensor<Self> {
599        ruprim::elementwise::cast::cast(tensor, dtype.into())
600    }
601
602    fn int_unfold(
603        tensor: FloatTensor<Self>,
604        dim: usize,
605        size: usize,
606        step: usize,
607    ) -> FloatTensor<Self> {
608        unfold(tensor, dim, size, step)
609    }
610
611    // TODO
612    // fn int_powi(lhs: IntTensor<Self>, rhs: IntTensor<Self>) -> IntTensor<Self> {
613    //     todo!()
614    // }
615
616    // fn int_powi_scalar_impl(lhs: IntTensor<Self>, rhs: Scalar) -> IntTensor<Self> {
617    //     todo!()
618    // }
619}