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 let all_steps_one = slices.iter().all(|info| info.step == 1);
65
66 if all_steps_one {
67 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 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 }