use super::{expand, numeric, permute, unfold};
use crate::{DeviceBackend, DeviceRuntime, FloatElement, IntElement, element::BoolElement};
use rurand::tensor::{random_bernoulli, random_normal, random_uniform};
use ruprim::elementwise::unary::float::{FloatUnaryOp, FloatUnaryOpFamily, launch_unary_float, unary_basic};
use ruprim::elementwise::unary::float::unary_basic::BasicFloatUnaryKind;
use ruprim::reduce::tensor as reduce;
use rublas::tensor_matmul::{MatmulStrategy, matmul};
use ruda_tensor::ops::GridSampleOptions;
use ruda_tensor::tensor::{BoolTensor, Device, FloatTensor, IntTensor};
use ruda_tensor::{DType, ElementConversion, FloatDType, Slice};
use ruda_tensor::{Distribution, Shape, TensorData, ops::FloatTensorOps};
use ruda_tensor::{ExecutionError, Scalar, get_device_settings};
use ruda_core::tensor::{BoolDType, IntDType};
use ruda_kernel::dsl::{self as ruda, prelude::*};
use ruprim::reduce::components::instructions::ReduceOperationConfig;
use std::ops::Range;
impl<R, F, I, BT> FloatTensorOps<Self> for DeviceBackend<R, F, I, BT>
where
R: DeviceRuntime,
F: FloatElement,
I: IntElement,
BT: BoolElement,
{
#[cfg_attr(feature = "tracing", tracing::instrument(
level="trace",
skip(data),
fields(?data.shape, ?data.dtype)
))]
fn float_from_data(data: TensorData, device: &Device<Self>) -> FloatTensor<Self> {
match data.dtype {
DType::F64 | DType::F32 | DType::F16 | DType::BF16 => super::from_data(data, device),
_ => unimplemented!("Unsupported dtype for `float_from_data`"),
}
}
fn float_random(
shape: Shape,
distribution: Distribution,
device: &Device<Self>,
dtype: FloatDType,
) -> FloatTensor<Self> {
let dtype = dtype.into();
match distribution {
Distribution::Default => random_uniform(shape, device, 0., 1., dtype),
Distribution::Uniform(low, high) => {
random_uniform(shape, device, low.elem(), high.elem(), dtype)
}
Distribution::Bernoulli(prob) => random_bernoulli(shape, device, prob as f32, dtype),
Distribution::Normal(mean, std) => {
random_normal(shape, device, mean.elem(), std.elem(), dtype)
}
}
}
#[cfg_attr(feature = "tracing", tracing::instrument(
level="trace",
skip(tensor),
fields(from = ?tensor.device, meta = ?tensor.meta, dtype = ?tensor.dtype)
))]
async fn float_into_data(tensor: FloatTensor<Self>) -> Result<TensorData, ExecutionError> {
super::into_data(tensor).await
}
fn float_device(tensor: &FloatTensor<Self>) -> Device<Self> {
tensor.device.clone()
}
#[cfg_attr(feature = "tracing", tracing::instrument(
level="trace",
skip(tensor),
fields(from = ?tensor.device, meta = ?tensor.meta, dtype = ?tensor.dtype)
))]
fn float_to_device(tensor: FloatTensor<Self>, device: &Device<Self>) -> FloatTensor<Self> {
super::to_device(tensor, device)
}
fn float_empty(shape: Shape, device: &Device<Self>, dtype: FloatDType) -> FloatTensor<Self> {
let dtype = dtype.into();
super::empty(shape, device, dtype)
}
fn float_add(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
numeric::add(lhs, rhs)
}
fn float_add_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
let dtype = lhs.dtype;
numeric::add_scalar(lhs, InputScalar::new(rhs, dtype))
}
fn float_zeros(shape: Shape, device: &Device<Self>, dtype: FloatDType) -> FloatTensor<Self> {
let dtype = dtype.into();
numeric::zeros(device.clone(), shape, dtype)
}
fn float_full(
shape: Shape,
fill_value: Scalar,
device: &R::Device,
dtype: FloatDType,
) -> FloatTensor<Self> {
let dtype: DType = dtype.into();
let client = R::client(device);
numeric::full_device_dtype(
client,
shape,
device.clone(),
InputScalar::new(fill_value, dtype),
dtype,
)
}
fn float_ones(shape: Shape, device: &Device<Self>, dtype: FloatDType) -> FloatTensor<Self> {
let dtype = dtype.into();
numeric::ones(device.clone(), shape, dtype)
}
fn float_sub(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
numeric::sub(lhs, rhs)
}
fn float_sub_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
let dtype = lhs.dtype;
numeric::sub_scalar(lhs, InputScalar::new(rhs, dtype))
}
fn float_mul(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
numeric::mul(lhs, rhs)
}
fn float_mul_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
let dtype = lhs.dtype;
numeric::mul_scalar(lhs, InputScalar::new(rhs, dtype))
}
fn float_div(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
numeric::div(lhs, rhs)
}
fn float_div_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
let dtype = lhs.dtype;
numeric::div_scalar(lhs, InputScalar::new(rhs, dtype))
}
fn float_remainder(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
numeric::remainder(lhs, rhs)
}
fn float_remainder_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
let dtype = lhs.dtype;
numeric::remainder_scalar(lhs, InputScalar::new(rhs, dtype))
}
fn float_matmul(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
let dtype = lhs.dtype;
matmul(lhs, rhs, None, MatmulStrategy::default(), dtype).unwrap()
}
fn float_cross(
lhs: FloatTensor<Self>,
rhs: FloatTensor<Self>,
dim: usize,
) -> FloatTensor<Self> {
rublas::tensor_vector::cross(lhs, rhs, dim)
}
fn float_swap_dims(tensor: FloatTensor<Self>, dim1: usize, dim2: usize) -> FloatTensor<Self> {
super::swap_dims(tensor, dim1, dim2)
}
fn float_reshape(tensor: FloatTensor<Self>, shape: Shape) -> FloatTensor<Self> {
super::reshape(tensor, shape)
}
fn float_gather(
dim: usize,
tensor: FloatTensor<Self>,
indices: IntTensor<Self>,
) -> FloatTensor<Self> {
ruprim::indexing::gather(dim, tensor, indices)
}
fn float_scatter_add(
dim: usize,
tensor: FloatTensor<Self>,
indices: IntTensor<Self>,
value: FloatTensor<Self>,
) -> FloatTensor<Self> {
ruprim::indexing::scatter(dim, tensor, indices, value, false)
}
fn float_scatter_nd(
data: FloatTensor<Self>,
indices: IntTensor<Self>,
values: FloatTensor<Self>,
reduction: ruda_tensor::tensor::IndexingUpdateOp,
) -> FloatTensor<Self> {
ruprim::indexing::scatter_nd(data, indices, values, reduction)
}
fn float_gather_nd(data: FloatTensor<Self>, indices: IntTensor<Self>) -> FloatTensor<Self> {
ruprim::indexing::gather_nd(data, indices)
}
fn float_select(
tensor: FloatTensor<Self>,
dim: usize,
indices: IntTensor<Self>,
) -> FloatTensor<Self> {
ruprim::indexing::select(tensor, dim, indices)
}
fn float_select_add(
tensor: FloatTensor<Self>,
dim: usize,
indices: IntTensor<Self>,
value: FloatTensor<Self>,
) -> FloatTensor<Self> {
ruprim::indexing::select_assign(tensor, dim, indices, value, false)
}
fn float_slice(tensor: FloatTensor<Self>, slices: &[Slice]) -> FloatTensor<Self> {
let all_steps_one = slices.iter().all(|info| info.step == 1);
if all_steps_one {
let simple_ranges: Vec<Range<usize>> = slices
.iter()
.enumerate()
.map(|(i, slice)| slice.to_range(tensor.meta.shape()[i]))
.collect();
ruprim::indexing::slice(tensor, &simple_ranges)
} else {
ruprim::indexing::slice_with_steps(tensor, slices)
}
}
fn float_slice_assign(
tensor: FloatTensor<Self>,
ranges: &[Slice],
value: FloatTensor<Self>,
) -> FloatTensor<Self> {
ruprim::indexing::slice_assign(tensor, ranges, value)
}
fn float_mask_where(
tensor: FloatTensor<Self>,
mask: BoolTensor<Self>,
value: FloatTensor<Self>,
) -> FloatTensor<Self> {
let bool_dtype = mask.dtype;
ruprim::elementwise::mask::mask_where_auto(tensor, mask, value, bool_dtype)
}
fn float_mask_fill(
tensor: FloatTensor<Self>,
mask: BoolTensor<Self>,
value: Scalar,
) -> FloatTensor<Self> {
let dtype = tensor.dtype;
let bool_dtype = mask.dtype;
ruprim::elementwise::mask::mask_fill_auto(tensor, mask, InputScalar::new(value, dtype), bool_dtype)
}
fn float_equal(
lhs: FloatTensor<Self>,
rhs: FloatTensor<Self>,
out_dtype: BoolDType,
) -> BoolTensor<Self> {
ruprim::elementwise::comparison::equal(lhs, rhs, out_dtype.into())
}
fn float_equal_elem(
lhs: FloatTensor<Self>,
rhs: Scalar,
out_dtype: BoolDType,
) -> BoolTensor<Self> {
let dtype = lhs.dtype;
ruprim::elementwise::comparison::equal_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
}
fn float_greater(
lhs: FloatTensor<Self>,
rhs: FloatTensor<Self>,
out_dtype: BoolDType,
) -> BoolTensor<Self> {
ruprim::elementwise::comparison::greater(lhs, rhs, out_dtype.into())
}
fn float_greater_elem(
lhs: FloatTensor<Self>,
rhs: Scalar,
out_dtype: BoolDType,
) -> BoolTensor<Self> {
let dtype = lhs.dtype;
ruprim::elementwise::comparison::greater_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
}
fn float_greater_equal(
lhs: FloatTensor<Self>,
rhs: FloatTensor<Self>,
out_dtype: BoolDType,
) -> BoolTensor<Self> {
ruprim::elementwise::comparison::greater_equal(lhs, rhs, out_dtype.into())
}
fn float_greater_equal_elem(
lhs: FloatTensor<Self>,
rhs: Scalar,
out_dtype: BoolDType,
) -> BoolTensor<Self> {
let dtype = lhs.dtype;
ruprim::elementwise::comparison::greater_equal_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
}
fn float_lower(
lhs: FloatTensor<Self>,
rhs: FloatTensor<Self>,
out_dtype: BoolDType,
) -> BoolTensor<Self> {
ruprim::elementwise::comparison::lower(lhs, rhs, out_dtype.into())
}
fn float_lower_elem(
lhs: FloatTensor<Self>,
rhs: Scalar,
out_dtype: BoolDType,
) -> BoolTensor<Self> {
let dtype = lhs.dtype;
ruprim::elementwise::comparison::lower_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
}
fn float_lower_equal(
lhs: FloatTensor<Self>,
rhs: FloatTensor<Self>,
out_dtype: BoolDType,
) -> BoolTensor<Self> {
ruprim::elementwise::comparison::lower_equal(lhs, rhs, out_dtype.into())
}
fn float_lower_equal_elem(
lhs: FloatTensor<Self>,
rhs: Scalar,
out_dtype: BoolDType,
) -> BoolTensor<Self> {
let dtype = lhs.dtype;
ruprim::elementwise::comparison::lower_equal_elem(lhs, InputScalar::new(rhs, dtype), out_dtype.into())
}
fn float_sum(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
reduce::sum_fallback(tensor, Default::default()).unwrap()
}
fn float_max(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
reduce::reduce(tensor, None, Default::default(), ReduceOperationConfig::Max).unwrap()
}
fn float_max_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
reduce::reduce_dim(
tensor,
None,
dim,
Default::default(),
ReduceOperationConfig::Max,
)
.unwrap()
}
fn float_min(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
reduce::reduce(tensor, None, Default::default(), ReduceOperationConfig::Min).unwrap()
}
fn float_min_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
reduce::reduce_dim(
tensor,
None,
dim,
Default::default(),
ReduceOperationConfig::Min,
)
.unwrap()
}
fn float_max_abs(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
reduce::reduce(
tensor,
None,
Default::default(),
ReduceOperationConfig::MaxAbs,
)
.unwrap()
}
fn float_max_abs_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
reduce::reduce_dim(
tensor,
None,
dim,
Default::default(),
ReduceOperationConfig::MaxAbs,
)
.unwrap()
}
fn float_sum_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
reduce::reduce_dim(
tensor,
None,
dim,
Default::default(),
ReduceOperationConfig::Sum,
)
.unwrap()
}
fn float_mean_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
reduce::reduce_dim(
tensor,
None,
dim,
Default::default(),
ReduceOperationConfig::Mean,
)
.unwrap()
}
fn float_mean(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
reduce::reduce(
tensor,
None,
Default::default(),
ReduceOperationConfig::Mean,
)
.unwrap()
}
fn float_cumsum(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
numeric::cumsum(tensor, dim)
}
fn float_cumprod(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
numeric::cumprod(tensor, dim)
}
fn float_cummin(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
numeric::cummin(tensor, dim)
}
fn float_cummax(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
numeric::cummax(tensor, dim)
}
fn float_prod(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
reduce::reduce(
tensor,
None,
Default::default(),
ReduceOperationConfig::Prod,
)
.unwrap()
}
fn float_prod_dim(tensor: FloatTensor<Self>, dim: usize) -> FloatTensor<Self> {
reduce::reduce_dim(
tensor,
None,
dim,
Default::default(),
ReduceOperationConfig::Prod,
)
.unwrap()
}
fn float_exp(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Exp)
}
fn float_log(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Log)
}
fn float_log1p(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Log1p)
}
fn float_powi_scalar(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
if matches!(rhs, Scalar::UInt(value) if value > i64::MAX as u64) {
return Self::float_powi_scalar_impl(lhs, rhs);
}
match rhs.elem::<i64>() {
0 => Self::float_ones(lhs.meta.shape().clone(), &lhs.device, lhs.dtype.into()),
1 => lhs,
2 => Self::float_mul(lhs.clone(), lhs),
-1 => Self::float_recip(lhs),
-2 => Self::float_recip(Self::float_mul(lhs.clone(), lhs)),
_ => Self::float_powi_scalar_impl(lhs, rhs),
}
}
fn float_powi_scalar_impl(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
ruprim::elementwise::binary::integer_power::scalar(lhs, rhs)
}
fn float_powf_scalar_impl(lhs: FloatTensor<Self>, rhs: Scalar) -> FloatTensor<Self> {
struct Powf;
#[ruda]
impl<F: Float, N: Size> FloatUnaryOp<F, N> for Powf {
type Options = InputScalar;
fn execute(input: Vector<F, N>, options: &Self::Options) -> Vector<F, N> {
Vector::powf(input, Vector::new(options.get::<F>()))
}
}
impl FloatUnaryOpFamily for Powf {
type Options = InputScalar;
type Unary<F: Float, N: Size> = Self;
}
let dtype = lhs.dtype;
launch_unary_float::<R, Powf, _>(lhs, |_| InputScalar::new(rhs, dtype))
}
fn float_sqrt(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Sqrt)
}
fn float_rsqrt(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::InverseSqrt)
}
fn float_abs(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Abs)
}
fn float_sign(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Sign)
}
fn float_cos(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Cos)
}
fn float_sin(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Sin)
}
fn float_tan(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Tan)
}
fn float_cosh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Cosh)
}
fn float_sinh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Sinh)
}
fn float_tanh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Tanh)
}
fn float_acos(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcCos)
}
fn float_acosh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcCosh)
}
fn float_asin(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcSin)
}
fn float_asinh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcSinh)
}
fn float_atan(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcTan)
}
fn float_atanh(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::ArcTanh)
}
fn float_atan2(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
ruprim::elementwise::binary::float::atan2::<R>(lhs, rhs)
}
fn float_round(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Round)
}
fn float_floor(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Floor)
}
fn float_ceil(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Ceil)
}
fn float_trunc(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Trunc)
}
fn float_erf(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Erf)
}
fn float_argmax(tensor: FloatTensor<Self>, dim: usize, out_dtype: IntDType) -> IntTensor<Self> {
reduce::reduce_dim(
tensor,
Some(out_dtype.into()),
dim,
Default::default(),
ReduceOperationConfig::ArgMax,
)
.unwrap()
}
fn float_argtopk(
tensor: FloatTensor<Self>,
dim: usize,
k: usize,
out_dtype: IntDType,
) -> IntTensor<Self> {
reduce::reduce_dim(
tensor,
Some(out_dtype.into()),
dim,
Default::default(),
ReduceOperationConfig::ArgTopK(k),
)
.unwrap()
}
fn float_topk(tensor: FloatTensor<Self>, dim: usize, k: usize) -> FloatTensor<Self> {
reduce::reduce_dim(
tensor,
None,
dim,
Default::default(),
ReduceOperationConfig::TopK(k),
)
.unwrap()
}
fn float_argmin(tensor: FloatTensor<Self>, dim: usize, out_dtype: IntDType) -> IntTensor<Self> {
reduce::reduce_dim(
tensor,
Some(out_dtype.into()),
dim,
Default::default(),
ReduceOperationConfig::ArgMin,
)
.unwrap()
}
fn float_into_int(tensor: FloatTensor<Self>, out_dtype: IntDType) -> IntTensor<Self> {
ruprim::elementwise::cast::cast(tensor, out_dtype.into())
}
fn float_clamp(tensor: FloatTensor<Self>, min: Scalar, max: Scalar) -> FloatTensor<Self> {
let dtype = tensor.dtype;
ruprim::elementwise::unary::clamp::clamp(
tensor,
InputScalar::new(min, dtype),
InputScalar::new(max, dtype),
)
}
fn float_recip(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
unary_basic::launch::<R, _>(tensor, |_| BasicFloatUnaryKind::Recip)
}
fn float_repeat_dim(tensor: FloatTensor<Self>, dim: usize, times: usize) -> FloatTensor<Self> {
ruprim::indexing::repeat_dim(tensor, dim, times)
}
fn float_powf(lhs: FloatTensor<Self>, rhs: FloatTensor<Self>) -> FloatTensor<Self> {
numeric::pow(lhs, rhs)
}
fn float_powi(lhs: FloatTensor<Self>, rhs: IntTensor<Self>) -> FloatTensor<Self> {
ruprim::elementwise::binary::integer_power::tensor(lhs, rhs)
}
fn float_permute(tensor: FloatTensor<Self>, axes: &[usize]) -> FloatTensor<Self> {
permute(tensor, axes)
}
fn float_expand(tensor: FloatTensor<Self>, shape: Shape) -> FloatTensor<Self> {
expand(tensor, shape)
}
fn float_flip(tensor: FloatTensor<Self>, axes: &[usize]) -> FloatTensor<Self> {
let bool_dtype = get_device_settings::<Self>(&tensor.device).bool_dtype;
ruprim::indexing::flip(tensor, axes, bool_dtype.into())
}
fn float_cast(tensor: FloatTensor<Self>, dtype: FloatDType) -> FloatTensor<Self> {
ruprim::elementwise::cast::cast(tensor, dtype.into())
}
fn float_unfold(
tensor: FloatTensor<Self>,
dim: usize,
size: usize,
step: usize,
) -> FloatTensor<Self> {
unfold(tensor, dim, size, step)
}
fn float_is_nan(tensor: FloatTensor<Self>, out_dtype: BoolDType) -> BoolTensor<Self> {
ruprim::elementwise::comparison::is_nan(tensor, out_dtype.into())
}
fn float_is_inf(tensor: FloatTensor<Self>, out_dtype: BoolDType) -> BoolTensor<Self> {
ruprim::elementwise::comparison::is_inf(tensor, out_dtype.into())
}
fn float_grid_sample_2d(
tensor: FloatTensor<Self>,
grid: FloatTensor<Self>,
options: GridSampleOptions,
) -> FloatTensor<Self> {
rudnn::grid_sample::grid_sample(tensor, grid, options)
}
}