Skip to main content

burn_cubecl/ops/
tensor.rs

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