Skip to main content

ruda_tensor_device/dispatch/
float.rs

1use super::{expand, numeric, permute, unfold};
2use crate::{DeviceBackend, DeviceRuntime, FloatElement, IntElement, element::BoolElement};
3use rurand::tensor::{random_bernoulli, random_normal, random_uniform};
4use ruprim::elementwise::unary::float::{FloatUnaryOp, FloatUnaryOpFamily, launch_unary_float, unary_basic};
5use ruprim::elementwise::unary::float::unary_basic::BasicFloatUnaryKind;
6use ruprim::reduce::tensor as reduce;
7use rublas::tensor_matmul::{MatmulStrategy, matmul};
8use ruda_tensor::ops::GridSampleOptions;
9use ruda_tensor::tensor::{BoolTensor, Device, FloatTensor, IntTensor};
10use ruda_tensor::{DType, ElementConversion, FloatDType, Slice};
11use ruda_tensor::{Distribution, Shape, TensorData, ops::FloatTensorOps};
12use ruda_tensor::{ExecutionError, Scalar, get_device_settings};
13use ruda_core::tensor::{BoolDType, IntDType};
14use ruda_kernel::dsl::{self as ruda, prelude::*};
15use ruprim::reduce::components::instructions::ReduceOperationConfig;
16use std::ops::Range;
17
18impl<R, F, I, BT> FloatTensorOps<Self> for DeviceBackend<R, F, I, BT>
19where
20    R: DeviceRuntime,
21    F: FloatElement,
22    I: IntElement,
23    BT: BoolElement,
24{
25    #[cfg_attr(feature = "tracing", tracing::instrument(
26        level="trace",
27        skip(data),
28        fields(?data.shape, ?data.dtype)
29    ))]
30    fn float_from_data(data: TensorData, device: &Device<Self>) -> FloatTensor<Self> {
31        match data.dtype {
32            DType::F64 | DType::F32 | DType::F16 | DType::BF16 => super::from_data(data, device),
33            _ => unimplemented!("Unsupported dtype for `float_from_data`"),
34        }
35    }
36
37    fn float_random(
38        shape: Shape,
39        distribution: Distribution,
40        device: &Device<Self>,
41        dtype: FloatDType,
42    ) -> FloatTensor<Self> {
43        let dtype = dtype.into();
44        match distribution {
45            Distribution::Default => random_uniform(shape, device, 0., 1., dtype),
46            Distribution::Uniform(low, high) => {
47                random_uniform(shape, device, low.elem(), high.elem(), dtype)
48            }
49            Distribution::Bernoulli(prob) => random_bernoulli(shape, device, prob as f32, dtype),
50            Distribution::Normal(mean, std) => {
51                random_normal(shape, device, mean.elem(), std.elem(), dtype)
52            }
53        }
54    }
55
56    #[cfg_attr(feature = "tracing", tracing::instrument(
57        level="trace",
58        skip(tensor),
59        fields(from = ?tensor.device, meta = ?tensor.meta, dtype = ?tensor.dtype)
60    ))]
61    async fn float_into_data(tensor: FloatTensor<Self>) -> Result<TensorData, ExecutionError> {
62        super::into_data(tensor).await
63    }
64
65    fn float_device(tensor: &FloatTensor<Self>) -> Device<Self> {
66        tensor.device.clone()
67    }
68
69    #[cfg_attr(feature = "tracing", tracing::instrument(
70        level="trace",
71        skip(tensor),
72        fields(from = ?tensor.device, meta = ?tensor.meta, dtype = ?tensor.dtype)
73    ))]
74    fn float_to_device(tensor: FloatTensor<Self>, device: &Device<Self>) -> FloatTensor<Self> {
75        super::to_device(tensor, device)
76    }
77
78    fn float_empty(shape: Shape, device: &Device<Self>, dtype: FloatDType) -> FloatTensor<Self> {
79        let dtype = dtype.into();
80        super::empty(shape, device, dtype)
81    }
82
83    fn float_add(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
84        numeric::add(lhs, rhs)
85    }
86
87    fn float_add_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
88        let dtype = lhs.dtype;
89        numeric::add_scalar(lhs, InputScalar::new(rhs, dtype))
90    }
91
92    fn float_zeros(shape: Shape, device: &Device<Self>, dtype: FloatDType) -> FloatTensor<Self> {
93        let dtype = dtype.into();
94        numeric::zeros(device.clone(), shape, dtype)
95    }
96
97    fn float_full(
98        shape: Shape,
99        fill_value: Scalar,
100        device: &R::Device,
101        dtype: FloatDType,
102    ) -> FloatTensor<Self> {
103        let dtype: DType = dtype.into();
104        let client = R::client(device);
105        numeric::full_device_dtype(
106            client,
107            shape,
108            device.clone(),
109            InputScalar::new(fill_value, dtype),
110            dtype,
111        )
112    }
113
114    fn float_ones(shape: Shape, device: &Device<Self>, dtype: FloatDType) -> FloatTensor<Self> {
115        let dtype = dtype.into();
116        numeric::ones(device.clone(), shape, dtype)
117    }
118
119    fn float_sub(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
120        numeric::sub(lhs, rhs)
121    }
122
123    fn float_sub_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
124        let dtype = lhs.dtype;
125        numeric::sub_scalar(lhs, InputScalar::new(rhs, dtype))
126    }
127
128    fn float_mul(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
129        numeric::mul(lhs, rhs)
130    }
131
132    fn float_mul_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
133        let dtype = lhs.dtype;
134        numeric::mul_scalar(lhs, InputScalar::new(rhs, dtype))
135    }
136
137    fn float_div(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
138        numeric::div(lhs, rhs)
139    }
140
141    fn float_div_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
142        let dtype = lhs.dtype;
143        numeric::div_scalar(lhs, InputScalar::new(rhs, dtype))
144    }
145
146    fn float_remainder(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
147        numeric::remainder(lhs, rhs)
148    }
149
150    fn float_remainder_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
151        let dtype = lhs.dtype;
152        numeric::remainder_scalar(lhs, InputScalar::new(rhs, dtype))
153    }
154
155    fn float_matmul(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
156        let dtype = lhs.dtype;
157        matmul(lhs, rhs, None, MatmulStrategy::default(), dtype).unwrap()
158    }
159
160    fn float_cross(
161        lhs: FloatTensor<Self>,
162        rhs: FloatTensor<Self>,
163        dim: usize,
164    ) -> FloatTensor<Self> {
165        rublas::tensor_vector::cross(lhs, rhs, dim)
166    }
167
168    fn float_swap_dims(tensor: FloatTensor<Self>, dim1: usize, dim2: usize) -> FloatTensor<Self> {
169        super::swap_dims(tensor, dim1, dim2)
170    }
171
172    fn float_reshape(tensor: FloatTensor<Self>, shape: Shape) -> FloatTensor<Self> {
173        super::reshape(tensor, shape)
174    }
175
176    fn float_gather(
177        dim: usize,
178        tensor: FloatTensor<Self>,
179        indices: IntTensor<Self>,
180    ) -> FloatTensor<Self> {
181        ruprim::indexing::gather(dim, tensor, indices)
182    }
183
184    fn float_scatter_add(
185        dim: usize,
186        tensor: FloatTensor<Self>,
187        indices: IntTensor<Self>,
188        value: FloatTensor<Self>,
189    ) -> FloatTensor<Self> {
190        ruprim::indexing::scatter(dim, tensor, indices, value, false)
191    }
192
193    fn float_scatter_nd(
194        data: FloatTensor<Self>,
195        indices: IntTensor<Self>,
196        values: FloatTensor<Self>,
197        reduction: ruda_tensor::tensor::IndexingUpdateOp,
198    ) -> FloatTensor<Self> {
199        ruprim::indexing::scatter_nd(data, indices, values, reduction)
200    }
201
202    fn float_gather_nd(data: FloatTensor<Self>, indices: IntTensor<Self>) -> FloatTensor<Self> {
203        ruprim::indexing::gather_nd(data, indices)
204    }
205
206    fn float_select(
207        tensor: FloatTensor<Self>,
208        dim: usize,
209        indices: IntTensor<Self>,
210    ) -> FloatTensor<Self> {
211        ruprim::indexing::select(tensor, dim, indices)
212    }
213
214    fn float_select_add(
215        tensor: FloatTensor<Self>,
216        dim: usize,
217        indices: IntTensor<Self>,
218        value: FloatTensor<Self>,
219    ) -> FloatTensor<Self> {
220        ruprim::indexing::select_assign(tensor, dim, indices, value, false)
221    }
222
223    fn float_slice(tensor: FloatTensor<Self>, slices: &[Slice]) -> FloatTensor<Self> {
224        // Check if all steps are 1
225        let all_steps_one = slices.iter().all(|info| info.step == 1);
226
227        if all_steps_one {
228            // Use optimized slice for step=1
229            let simple_ranges: Vec<Range<usize>> = slices
230                .iter()
231                .enumerate()
232                .map(|(i, slice)| slice.to_range(tensor.meta.shape()[i]))
233                .collect();
234
235            ruprim::indexing::slice(tensor, &simple_ranges)
236        } else {
237            // Use slice with steps kernel
238            ruprim::indexing::slice_with_steps(tensor, slices)
239        }
240    }
241
242    fn float_slice_assign(
243        tensor: FloatTensor<Self>,
244        ranges: &[Slice],
245        value: FloatTensor<Self>,
246    ) -> FloatTensor<Self> {
247        ruprim::indexing::slice_assign(tensor, ranges, value)
248    }
249
250    fn float_mask_where(
251        tensor: FloatTensor<Self>,
252        mask: BoolTensor<Self>,
253        value: FloatTensor<Self>,
254    ) -> FloatTensor<Self> {
255        let bool_dtype = mask.dtype;
256        ruprim::elementwise::mask::mask_where_auto(tensor, mask, value, bool_dtype)
257    }
258
259    fn float_mask_fill(
260        tensor: FloatTensor<Self>,
261        mask: BoolTensor<Self>,
262        value: Scalar,
263    ) -> FloatTensor<Self> {
264        let dtype = tensor.dtype;
265        let bool_dtype = mask.dtype;
266        ruprim::elementwise::mask::mask_fill_auto(tensor, mask, InputScalar::new(value, dtype), bool_dtype)
267    }
268
269    fn float_equal(
270        lhs: FloatTensor<Self>,
271        rhs: FloatTensor<Self>,
272        out_dtype: BoolDType,
273    ) -> BoolTensor<Self> {
274        ruprim::elementwise::comparison::equal(lhs, rhs, out_dtype.into())
275    }
276
277    fn float_equal_elem(
278        lhs: FloatTensor<Self>,
279        rhs: Scalar,
280        out_dtype: BoolDType,
281    ) -> BoolTensor<Self> {
282        let dtype = lhs.dtype;
283        ruprim::elementwise::comparison::equal_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
284    }
285
286    fn float_greater(
287        lhs: FloatTensor<Self>,
288        rhs: FloatTensor<Self>,
289        out_dtype: BoolDType,
290    ) -> BoolTensor<Self> {
291        ruprim::elementwise::comparison::greater(lhs, rhs, out_dtype.into())
292    }
293
294    fn float_greater_elem(
295        lhs: FloatTensor<Self>,
296        rhs: Scalar,
297        out_dtype: BoolDType,
298    ) -> BoolTensor<Self> {
299        let dtype = lhs.dtype;
300        ruprim::elementwise::comparison::greater_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
301    }
302
303    fn float_greater_equal(
304        lhs: FloatTensor<Self>,
305        rhs: FloatTensor<Self>,
306        out_dtype: BoolDType,
307    ) -> BoolTensor<Self> {
308        ruprim::elementwise::comparison::greater_equal(lhs, rhs, out_dtype.into())
309    }
310
311    fn float_greater_equal_elem(
312        lhs: FloatTensor<Self>,
313        rhs: Scalar,
314        out_dtype: BoolDType,
315    ) -> BoolTensor<Self> {
316        let dtype = lhs.dtype;
317        ruprim::elementwise::comparison::greater_equal_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
318    }
319
320    fn float_lower(
321        lhs: FloatTensor<Self>,
322        rhs: FloatTensor<Self>,
323        out_dtype: BoolDType,
324    ) -> BoolTensor<Self> {
325        ruprim::elementwise::comparison::lower(lhs, rhs, out_dtype.into())
326    }
327
328    fn float_lower_elem(
329        lhs: FloatTensor<Self>,
330        rhs: Scalar,
331        out_dtype: BoolDType,
332    ) -> BoolTensor<Self> {
333        let dtype = lhs.dtype;
334        ruprim::elementwise::comparison::lower_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
335    }
336
337    fn float_lower_equal(
338        lhs: FloatTensor<Self>,
339        rhs: FloatTensor<Self>,
340        out_dtype: BoolDType,
341    ) -> BoolTensor<Self> {
342        ruprim::elementwise::comparison::lower_equal(lhs, rhs, out_dtype.into())
343    }
344
345    fn float_lower_equal_elem(
346        lhs: FloatTensor<Self>,
347        rhs: Scalar,
348        out_dtype: BoolDType,
349    ) -> BoolTensor<Self> {
350        let dtype = lhs.dtype;
351        ruprim::elementwise::comparison::lower_equal_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
352    }
353
354    fn float_sum(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
355        reduce::sum_fallback(tensor, Default::default()).unwrap()
356    }
357
358    fn float_max(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
359        reduce::reduce(tensor, None, Default::default(), ReduceOperationConfig::Max).unwrap()
360    }
361
362    fn float_max_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
363        reduce::reduce_dim(
364            tensor,
365            None,
366            dim,
367            Default::default(),
368            ReduceOperationConfig::Max,
369        )
370        .unwrap()
371    }
372
373    fn float_min(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
374        reduce::reduce(tensor, None, Default::default(), ReduceOperationConfig::Min).unwrap()
375    }
376
377    fn float_min_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
378        reduce::reduce_dim(
379            tensor,
380            None,
381            dim,
382            Default::default(),
383            ReduceOperationConfig::Min,
384        )
385        .unwrap()
386    }
387
388    fn float_max_abs(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
389        reduce::reduce(
390            tensor,
391            None,
392            Default::default(),
393            ReduceOperationConfig::MaxAbs,
394        )
395        .unwrap()
396    }
397
398    fn float_max_abs_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
399        reduce::reduce_dim(
400            tensor,
401            None,
402            dim,
403            Default::default(),
404            ReduceOperationConfig::MaxAbs,
405        )
406        .unwrap()
407    }
408
409    fn float_sum_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
410        reduce::reduce_dim(
411            tensor,
412            None,
413            dim,
414            Default::default(),
415            ReduceOperationConfig::Sum,
416        )
417        .unwrap()
418    }
419
420    fn float_mean_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
421        reduce::reduce_dim(
422            tensor,
423            None,
424            dim,
425            Default::default(),
426            ReduceOperationConfig::Mean,
427        )
428        .unwrap()
429    }
430
431    fn float_mean(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
432        reduce::reduce(
433            tensor,
434            None,
435            Default::default(),
436            ReduceOperationConfig::Mean,
437        )
438        .unwrap()
439    }
440
441    fn float_cumsum(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
442        numeric::cumsum(tensor, dim)
443    }
444
445    fn float_cumprod(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
446        numeric::cumprod(tensor, dim)
447    }
448
449    fn float_cummin(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
450        numeric::cummin(tensor, dim)
451    }
452
453    fn float_cummax(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
454        numeric::cummax(tensor, dim)
455    }
456
457    fn float_prod(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
458        reduce::reduce(
459            tensor,
460            None,
461            Default::default(),
462            ReduceOperationConfig::Prod,
463        )
464        .unwrap()
465    }
466
467    fn float_prod_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
468        reduce::reduce_dim(
469            tensor,
470            None,
471            dim,
472            Default::default(),
473            ReduceOperationConfig::Prod,
474        )
475        .unwrap()
476    }
477
478    fn float_exp(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
479        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Exp)
480    }
481
482    fn float_log(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
483        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Log)
484    }
485
486    fn float_log1p(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
487        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Log1p)
488    }
489
490    fn float_powi_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
491        if matches!(rhs, Scalar::UInt(value) if value > i64::MAX as u64) {
492            return Self::float_powi_scalar_impl(lhs, rhs);
493        }
494        match rhs.elem::<i64>() {
495            0 => Self::float_ones(lhs.meta.shape().clone(), &lhs.device, lhs.dtype.into()),
496            1 => lhs,
497            2 => Self::float_mul(lhs.clone(), lhs),
498            -1 => Self::float_recip(lhs),
499            -2 => Self::float_recip(Self::float_mul(lhs.clone(), lhs)),
500            _ => Self::float_powi_scalar_impl(lhs, rhs),
501        }
502    }
503
504    fn float_powi_scalar_impl(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
505        ruprim::elementwise::binary::integer_power::scalar(lhs, rhs)
506    }
507
508    fn float_powf_scalar_impl(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
509        struct Powf;
510
511        #[ruda]
512        impl<F: Float, N: Size> FloatUnaryOp<F, N> for Powf {
513            type Options = InputScalar;
514
515            fn execute(input: Vector<F, N>, options: &Self::Options) -> Vector<F, N> {
516                Vector::powf(input, Vector::new(options.get::<F>()))
517            }
518        }
519
520        impl FloatUnaryOpFamily for Powf {
521            type Options = InputScalar;
522            type Unary<F: Float, N: Size> = Self;
523        }
524
525        let dtype = lhs.dtype;
526        launch_unary_float::<R, Powf, _>(lhs, |_| InputScalar::new(rhs, dtype))
527    }
528
529    fn float_sqrt(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
530        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Sqrt)
531    }
532
533    fn float_rsqrt(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
534        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::InverseSqrt)
535    }
536
537    fn float_abs(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
538        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Abs)
539    }
540
541    fn float_sign(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
542        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Sign)
543    }
544
545    fn float_cos(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
546        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Cos)
547    }
548
549    fn float_sin(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
550        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Sin)
551    }
552
553    fn float_tan(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
554        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Tan)
555    }
556
557    fn float_cosh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
558        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Cosh)
559    }
560
561    fn float_sinh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
562        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Sinh)
563    }
564
565    fn float_tanh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
566        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Tanh)
567    }
568
569    fn float_acos(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
570        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcCos)
571    }
572
573    fn float_acosh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
574        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcCosh)
575    }
576
577    fn float_asin(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
578        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcSin)
579    }
580
581    fn float_asinh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
582        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcSinh)
583    }
584
585    fn float_atan(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
586        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcTan)
587    }
588
589    fn float_atanh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
590        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcTanh)
591    }
592
593    fn float_atan2(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
594        ruprim::elementwise::binary::float::atan2::<R>(lhs, rhs)
595    }
596
597    fn float_round(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
598        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Round)
599    }
600
601    fn float_floor(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
602        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Floor)
603    }
604
605    fn float_ceil(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
606        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Ceil)
607    }
608
609    fn float_trunc(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
610        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Trunc)
611    }
612
613    fn float_erf(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
614        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Erf)
615    }
616
617    fn float_argmax(tensor: FloatTensor<Self>, dim: usize, out_dtype: IntDType) -> IntTensor<Self> {
618        reduce::reduce_dim(
619            tensor,
620            Some(out_dtype.into()),
621            dim,
622            Default::default(),
623            ReduceOperationConfig::ArgMax,
624        )
625        .unwrap()
626    }
627
628    fn float_argtopk(
629        tensor: FloatTensor<Self>,
630        dim: usize,
631        k: usize,
632        out_dtype: IntDType,
633    ) -> IntTensor<Self> {
634        reduce::reduce_dim(
635            tensor,
636            Some(out_dtype.into()),
637            dim,
638            Default::default(),
639            ReduceOperationConfig::ArgTopK(k),
640        )
641        .unwrap()
642    }
643
644    fn float_topk(tensor: FloatTensor<Self>, dim: usize, k: usize) -> FloatTensor<Self> {
645        reduce::reduce_dim(
646            tensor,
647            None,
648            dim,
649            Default::default(),
650            ReduceOperationConfig::TopK(k),
651        )
652        .unwrap()
653    }
654
655    fn float_argmin(tensor: FloatTensor<Self>, dim: usize, out_dtype: IntDType) -> IntTensor<Self> {
656        reduce::reduce_dim(
657            tensor,
658            Some(out_dtype.into()),
659            dim,
660            Default::default(),
661            ReduceOperationConfig::ArgMin,
662        )
663        .unwrap()
664    }
665
666    fn float_into_int(tensor: FloatTensor<Self>, out_dtype: IntDType) -> IntTensor<Self> {
667        ruprim::elementwise::cast::cast(tensor, out_dtype.into())
668    }
669
670    fn float_clamp(tensor: FloatTensor<Self>, min: Scalar, max: Scalar) -> FloatTensor<Self> {
671        let dtype = tensor.dtype;
672        ruprim::elementwise::unary::clamp::clamp(
673            tensor,
674            InputScalar::new(min, dtype),
675            InputScalar::new(max, dtype),
676        )
677    }
678
679    fn float_recip(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
680        unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Recip)
681    }
682
683    fn float_repeat_dim(tensor: FloatTensor<Self>, dim: usize, times: usize) -> FloatTensor<Self> {
684        ruprim::indexing::repeat_dim(tensor, dim, times)
685    }
686
687    fn float_powf(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
688        numeric::pow(lhs, rhs)
689    }
690
691    fn float_powi(lhs: FloatTensor<Self>, rhs: IntTensor<Self>) -> FloatTensor<Self> {
692        ruprim::elementwise::binary::integer_power::tensor(lhs, rhs)
693    }
694
695    fn float_permute(tensor: FloatTensor<Self>, axes: &[usize]) -> FloatTensor<Self> {
696        permute(tensor, axes)
697    }
698
699    fn float_expand(tensor: FloatTensor<Self>, shape: Shape) -> FloatTensor<Self> {
700        expand(tensor, shape)
701    }
702
703    fn float_flip(tensor: FloatTensor<Self>, axes: &[usize]) -> FloatTensor<Self> {
704        let bool_dtype = get_device_settings::<Self>(&tensor.device).bool_dtype;
705        ruprim::indexing::flip(tensor, axes, bool_dtype.into())
706    }
707
708    fn float_cast(tensor: FloatTensor<Self>, dtype: FloatDType) -> FloatTensor<Self> {
709        ruprim::elementwise::cast::cast(tensor, dtype.into())
710    }
711
712    fn float_unfold(
713        tensor: FloatTensor<Self>,
714        dim: usize,
715        size: usize,
716        step: usize,
717    ) -> FloatTensor<Self> {
718        unfold(tensor, dim, size, step)
719    }
720
721    fn float_is_nan(tensor: FloatTensor<Self>, out_dtype: BoolDType) -> BoolTensor<Self> {
722        ruprim::elementwise::comparison::is_nan(tensor, out_dtype.into())
723    }
724
725    fn float_is_inf(tensor: FloatTensor<Self>, out_dtype: BoolDType) -> BoolTensor<Self> {
726        ruprim::elementwise::comparison::is_inf(tensor, out_dtype.into())
727    }
728
729    fn float_grid_sample_2d(
730        tensor: FloatTensor<Self>,
731        grid: FloatTensor<Self>,
732        options: GridSampleOptions,
733    ) -> FloatTensor<Self> {
734        rudnn::grid_sample::grid_sample(tensor, grid, options)
735    }
736}