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