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