use std::ptr;
use flodl_sys::{self as ffi, FlodlTensor};
use super::{Tensor, check_err, Result, ffi_call};
impl Tensor {
pub fn add(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_add, self.handle, other.handle)
}
pub fn sub(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_sub, self.handle, other.handle)
}
pub fn mul(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_mul, self.handle, other.handle)
}
pub fn matmul(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_matmul, self.handle, other.handle)
}
pub fn mul_scalar(&self, scalar: f64) -> Result<Tensor> {
ffi_call!(flodl_mul_scalar, self.handle, scalar)
}
pub fn div(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_div, self.handle, other.handle)
}
pub fn neg(&self) -> Result<Tensor> {
ffi_call!(flodl_neg, self.handle)
}
pub fn add_scalar(&self, scalar: f64) -> Result<Tensor> {
ffi_call!(flodl_add_scalar, self.handle, scalar)
}
pub fn div_scalar(&self, scalar: f64) -> Result<Tensor> {
ffi_call!(flodl_div_scalar, self.handle, scalar)
}
pub fn exp(&self) -> Result<Tensor> {
ffi_call!(flodl_exp, self.handle)
}
pub fn log(&self) -> Result<Tensor> {
ffi_call!(flodl_log, self.handle)
}
pub fn sqrt(&self) -> Result<Tensor> {
ffi_call!(flodl_sqrt, self.handle)
}
pub fn abs(&self) -> Result<Tensor> {
ffi_call!(flodl_abs, self.handle)
}
pub fn triu(&self, diagonal: i64) -> Result<Tensor> {
ffi_call!(flodl_triu, self.handle, diagonal)
}
pub fn tril(&self, diagonal: i64) -> Result<Tensor> {
ffi_call!(flodl_tril, self.handle, diagonal)
}
pub fn pow_scalar(&self, exponent: f64) -> Result<Tensor> {
ffi_call!(flodl_pow_scalar, self.handle, exponent)
}
pub fn clamp(&self, min: f64, max: f64) -> Result<Tensor> {
ffi_call!(flodl_clamp, self.handle, min, max)
}
pub fn clamp_min(&self, min: f64) -> Result<Tensor> {
ffi_call!(flodl_clamp_min, self.handle, min)
}
pub fn clamp_max(&self, max: f64) -> Result<Tensor> {
ffi_call!(flodl_clamp_max, self.handle, max)
}
pub fn log1p(&self) -> Result<Tensor> {
ffi_call!(flodl_log1p, self.handle)
}
pub fn expm1(&self) -> Result<Tensor> {
ffi_call!(flodl_expm1, self.handle)
}
pub fn log2(&self) -> Result<Tensor> {
ffi_call!(flodl_log2, self.handle)
}
pub fn log10(&self) -> Result<Tensor> {
ffi_call!(flodl_log10, self.handle)
}
pub fn sin(&self) -> Result<Tensor> {
ffi_call!(flodl_sin, self.handle)
}
pub fn cos(&self) -> Result<Tensor> {
ffi_call!(flodl_cos, self.handle)
}
pub fn tan(&self) -> Result<Tensor> {
ffi_call!(flodl_tan, self.handle)
}
pub fn asin(&self) -> Result<Tensor> {
ffi_call!(flodl_asin, self.handle)
}
pub fn acos(&self) -> Result<Tensor> {
ffi_call!(flodl_acos, self.handle)
}
pub fn atan(&self) -> Result<Tensor> {
ffi_call!(flodl_atan, self.handle)
}
pub fn sign(&self) -> Result<Tensor> {
ffi_call!(flodl_sign, self.handle)
}
pub fn floor(&self) -> Result<Tensor> {
ffi_call!(flodl_floor, self.handle)
}
pub fn ceil(&self) -> Result<Tensor> {
ffi_call!(flodl_ceil, self.handle)
}
pub fn round(&self) -> Result<Tensor> {
ffi_call!(flodl_round, self.handle)
}
pub fn reciprocal(&self) -> Result<Tensor> {
ffi_call!(flodl_reciprocal, self.handle)
}
pub fn erf(&self) -> Result<Tensor> {
ffi_call!(flodl_erf, self.handle)
}
pub fn erfc(&self) -> Result<Tensor> {
ffi_call!(flodl_erfc, self.handle)
}
pub fn trunc(&self) -> Result<Tensor> {
ffi_call!(flodl_trunc, self.handle)
}
pub fn frac(&self) -> Result<Tensor> {
ffi_call!(flodl_frac, self.handle)
}
pub fn fmod(&self, divisor: f64) -> Result<Tensor> {
ffi_call!(flodl_fmod_scalar, self.handle, divisor)
}
pub fn fmod_tensor(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_fmod_tensor, self.handle, other.handle)
}
pub fn remainder(&self, divisor: f64) -> Result<Tensor> {
ffi_call!(flodl_remainder_scalar, self.handle, divisor)
}
pub fn remainder_tensor(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_remainder_tensor, self.handle, other.handle)
}
pub fn lerp(&self, end: &Tensor, weight: f64) -> Result<Tensor> {
ffi_call!(flodl_lerp, self.handle, end.handle, weight)
}
pub fn lerp_tensor(&self, end: &Tensor, weight: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_lerp_tensor, self.handle, end.handle, weight.handle)
}
pub fn isclose(&self, other: &Tensor, rtol: f64, atol: f64) -> Result<Tensor> {
ffi_call!(flodl_isclose, self.handle, other.handle, rtol, atol)
}
pub fn addmm(&self, mat1: &Tensor, mat2: &Tensor, beta: f64, alpha: f64) -> Result<Tensor> {
ffi_call!(flodl_addmm, self.handle, mat1.handle, mat2.handle, beta, alpha)
}
pub fn addcmul(&self, tensor1: &Tensor, tensor2: &Tensor, value: f64) -> Result<Tensor> {
ffi_call!(flodl_addcmul, self.handle, tensor1.handle, tensor2.handle, value)
}
pub fn addcdiv(&self, tensor1: &Tensor, tensor2: &Tensor, value: f64) -> Result<Tensor> {
ffi_call!(flodl_addcdiv, self.handle, tensor1.handle, tensor2.handle, value)
}
pub fn selu(&self) -> Result<Tensor> {
ffi_call!(flodl_selu, self.handle)
}
pub fn hardswish(&self) -> Result<Tensor> {
ffi_call!(flodl_hardswish, self.handle)
}
pub fn hardsigmoid(&self) -> Result<Tensor> {
ffi_call!(flodl_hardsigmoid, self.handle)
}
pub fn prelu(&self, weight: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_prelu, self.handle, weight.handle)
}
pub fn relu(&self) -> Result<Tensor> {
ffi_call!(flodl_relu, self.handle)
}
pub fn sigmoid(&self) -> Result<Tensor> {
ffi_call!(flodl_sigmoid, self.handle)
}
pub fn tanh(&self) -> Result<Tensor> {
ffi_call!(flodl_tanh_op, self.handle)
}
pub fn softmax(&self, dim: i32) -> Result<Tensor> {
ffi_call!(flodl_softmax, self.handle, dim)
}
pub fn log_softmax(&self, dim: i32) -> Result<Tensor> {
ffi_call!(flodl_log_softmax, self.handle, dim)
}
pub fn gelu(&self) -> Result<Tensor> {
ffi_call!(flodl_gelu, self.handle)
}
pub fn gelu_tanh(&self) -> Result<Tensor> {
ffi_call!(flodl_gelu_tanh, self.handle)
}
pub fn silu(&self) -> Result<Tensor> {
ffi_call!(flodl_silu, self.handle)
}
pub fn leaky_relu(&self, negative_slope: f64) -> Result<Tensor> {
ffi_call!(flodl_leaky_relu, self.handle, negative_slope)
}
pub fn elu(&self, alpha: f64) -> Result<Tensor> {
ffi_call!(flodl_elu, self.handle, alpha)
}
pub fn softplus(&self, beta: f64, threshold: f64) -> Result<Tensor> {
ffi_call!(flodl_softplus, self.handle, beta, threshold)
}
pub fn mish(&self) -> Result<Tensor> {
ffi_call!(flodl_mish, self.handle)
}
pub fn sum(&self) -> Result<Tensor> {
ffi_call!(flodl_sum, self.handle)
}
pub fn mean(&self) -> Result<Tensor> {
ffi_call!(flodl_mean, self.handle)
}
pub fn sum_dim(&self, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_sum_dim, self.handle, dim, keepdim as i32)
}
pub fn mean_dim(&self, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_mean_dim, self.handle, dim, keepdim as i32)
}
pub fn prod(&self) -> Result<Tensor> {
ffi_call!(flodl_prod, self.handle)
}
pub fn prod_dim(&self, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_prod_dim, self.handle, dim, keepdim as i32)
}
pub fn cumsum(&self, dim: i32) -> Result<Tensor> {
ffi_call!(flodl_cumsum, self.handle, dim)
}
pub fn logsumexp(&self, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_logsumexp, self.handle, dim, keepdim as i32)
}
pub fn min(&self) -> Result<Tensor> {
ffi_call!(flodl_min, self.handle)
}
pub fn max(&self) -> Result<Tensor> {
ffi_call!(flodl_max, self.handle)
}
pub fn norm(&self) -> Result<Tensor> {
ffi_call!(flodl_norm, self.handle)
}
pub fn norm_p(&self, p: f64, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_norm_p_dim, self.handle, p, dim, keepdim as i32)
}
pub fn sum_dims(&self, dims: &[i32], keepdim: bool) -> Result<Tensor> {
let mut dims64: Vec<i64> = dims.iter().map(|&d| d as i64).collect();
ffi_call!(flodl_sum_dims, self.handle, dims64.as_mut_ptr(), dims.len() as i32, keepdim as i32)
}
pub fn cumprod(&self, dim: i32) -> Result<Tensor> {
ffi_call!(flodl_cumprod, self.handle, dim)
}
pub fn median(&self) -> Result<Tensor> {
ffi_call!(flodl_median, self.handle)
}
pub fn median_dim(&self, dim: i32, keepdim: bool) -> Result<(Tensor, Tensor)> {
let mut vals: FlodlTensor = ptr::null_mut();
let mut idxs: FlodlTensor = ptr::null_mut();
let err = unsafe { ffi::flodl_median_dim(self.handle, dim, keepdim as i32, &mut vals, &mut idxs) };
check_err(err)?;
Ok((Tensor::from_raw(vals), Tensor::from_raw(idxs)))
}
pub fn count_nonzero(&self) -> Result<Tensor> {
ffi_call!(flodl_count_nonzero, self.handle)
}
pub fn count_nonzero_dim(&self, dim: i32) -> Result<Tensor> {
ffi_call!(flodl_count_nonzero_dim, self.handle, dim)
}
pub fn nonzero(&self) -> Result<Tensor> {
ffi_call!(flodl_nonzero, self.handle)
}
pub fn unique(&self, sorted: bool, return_inverse: bool) -> Result<(Tensor, Tensor)> {
let mut output: FlodlTensor = ptr::null_mut();
let mut inverse: FlodlTensor = ptr::null_mut();
let err = unsafe {
ffi::flodl_unique(self.handle, sorted as i32, return_inverse as i32, &mut output, &mut inverse)
};
check_err(err)?;
let inv = if inverse.is_null() {
Tensor::from_i64(&[], &[0], super::Device::CPU)?
} else {
Tensor::from_raw(inverse)
};
Ok((Tensor::from_raw(output), inv))
}
pub fn unique_consecutive(&self, return_inverse: bool) -> Result<(Tensor, Tensor)> {
let mut output: FlodlTensor = ptr::null_mut();
let mut inverse: FlodlTensor = ptr::null_mut();
let err = unsafe {
ffi::flodl_unique_consecutive(self.handle, return_inverse as i32, &mut output, &mut inverse)
};
check_err(err)?;
let inv = if inverse.is_null() {
Tensor::from_i64(&[], &[0], super::Device::CPU)?
} else {
Tensor::from_raw(inverse)
};
Ok((Tensor::from_raw(output), inv))
}
pub fn searchsorted(&self, values: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_searchsorted, self.handle, values.handle)
}
pub fn min_dim(&self, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_min_dim, self.handle, dim, keepdim as i32)
}
pub fn max_dim(&self, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_max_dim, self.handle, dim, keepdim as i32)
}
pub fn argmax(&self, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_argmax, self.handle, dim, keepdim as i32)
}
pub fn argmin(&self, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_argmin, self.handle, dim, keepdim as i32)
}
pub fn var(&self) -> Result<Tensor> {
ffi_call!(flodl_var, self.handle)
}
#[allow(clippy::should_implement_trait)]
pub fn std(&self) -> Result<Tensor> {
ffi_call!(flodl_std_op, self.handle)
}
pub fn var_dim(&self, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_var_dim, self.handle, dim, keepdim as i32)
}
pub fn std_dim(&self, dim: i32, keepdim: bool) -> Result<Tensor> {
ffi_call!(flodl_std_dim, self.handle, dim, keepdim as i32)
}
pub fn gt_scalar(&self, scalar: f64) -> Result<Tensor> {
ffi_call!(flodl_gt_scalar, self.handle, scalar)
}
pub fn ge_scalar(&self, scalar: f64) -> Result<Tensor> {
ffi_call!(flodl_ge_scalar, self.handle, scalar)
}
pub fn le_scalar(&self, scalar: f64) -> Result<Tensor> {
ffi_call!(flodl_le_scalar, self.handle, scalar)
}
pub fn lt_scalar(&self, scalar: f64) -> Result<Tensor> {
ffi_call!(flodl_lt_scalar, self.handle, scalar)
}
pub fn eq_scalar(&self, scalar: f64) -> Result<Tensor> {
ffi_call!(flodl_eq_scalar, self.handle, scalar)
}
pub fn ne_scalar(&self, scalar: f64) -> Result<Tensor> {
ffi_call!(flodl_ne_scalar, self.handle, scalar)
}
pub fn isnan(&self) -> Result<Tensor> {
ffi_call!(flodl_isnan, self.handle)
}
pub fn isinf(&self) -> Result<Tensor> {
ffi_call!(flodl_isinf, self.handle)
}
pub fn logical_and(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_logical_and, self.handle, other.handle)
}
pub fn logical_or(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_logical_or, self.handle, other.handle)
}
pub fn logical_not(&self) -> Result<Tensor> {
ffi_call!(flodl_logical_not, self.handle)
}
pub fn any(&self) -> Result<Tensor> {
ffi_call!(flodl_any, self.handle)
}
pub fn all(&self) -> Result<Tensor> {
ffi_call!(flodl_all, self.handle)
}
pub fn atan2(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_atan2, self.handle, other.handle)
}
pub fn maximum(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_maximum, self.handle, other.handle)
}
pub fn minimum(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_minimum, self.handle, other.handle)
}
pub fn gt(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_gt_tensor, self.handle, other.handle)
}
pub fn lt(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_lt_tensor, self.handle, other.handle)
}
pub fn ge(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_ge_tensor, self.handle, other.handle)
}
pub fn le(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_le_tensor, self.handle, other.handle)
}
pub fn eq_tensor(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_eq_tensor, self.handle, other.handle)
}
pub fn ne_tensor(&self, other: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_ne_tensor, self.handle, other.handle)
}
pub fn masked_fill(&self, mask: &Tensor, value: f64) -> Result<Tensor> {
ffi_call!(flodl_masked_fill, self.handle, mask.handle, value)
}
pub fn where_cond(condition: &Tensor, x: &Tensor, y: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_where, condition.handle, x.handle, y.handle)
}
pub fn topk(&self, k: i64, dim: i32, largest: bool, sorted: bool) -> Result<(Tensor, Tensor)> {
let mut values: FlodlTensor = ptr::null_mut();
let mut indices: FlodlTensor = ptr::null_mut();
let err = unsafe {
ffi::flodl_topk(
self.handle, k, dim, largest as i32, sorted as i32,
&mut values, &mut indices,
)
};
check_err(err)?;
Ok((Tensor::from_raw(values), Tensor::from_raw(indices)))
}
pub fn sort(&self, dim: i32, descending: bool) -> Result<(Tensor, Tensor)> {
let mut values: FlodlTensor = ptr::null_mut();
let mut indices: FlodlTensor = ptr::null_mut();
let err = unsafe {
ffi::flodl_sort(self.handle, dim, descending as i32, &mut values, &mut indices)
};
check_err(err)?;
Ok((Tensor::from_raw(values), Tensor::from_raw(indices)))
}
pub fn argsort(&self, dim: i32, descending: bool) -> Result<Tensor> {
ffi_call!(flodl_argsort, self.handle, dim, descending as i32)
}
pub fn gather(&self, dim: i32, index: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_gather, self.handle, dim, index.handle)
}
pub fn scatter_add(&self, dim: i32, index: &Tensor, src: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_scatter_add, self.handle, dim, index.handle, src.handle)
}
pub fn scatter(&self, dim: i32, index: &Tensor, src: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_scatter, self.handle, dim, index.handle, src.handle)
}
pub fn index_select(&self, dim: i32, index: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_index_select, self.handle, dim, index.handle)
}
pub fn index_add(&self, dim: i32, index: &Tensor, src: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_index_add, self.handle, dim, index.handle, src.handle)
}
pub fn select_scatter(&self, src: &Tensor, dim: i32, index: i64) -> Result<Tensor> {
ffi_call!(flodl_select_scatter, self.handle, src.handle, dim, index)
}
pub fn normalize(&self, p: f64, dim: i32) -> Result<Tensor> {
ffi_call!(flodl_normalize, self.handle, p, dim)
}
pub fn multinomial(&self, num_samples: i64, replacement: bool) -> Result<Tensor> {
let mut handle: FlodlTensor = ptr::null_mut();
let err = unsafe {
ffi::flodl_multinomial(
self.handle, num_samples, replacement as i32, &mut handle,
)
};
check_err(err)?;
Ok(Tensor::from_raw(handle))
}
pub fn cdist(&self, other: &Tensor) -> Result<Tensor> {
self.cdist_p(other, 2.0)
}
pub fn cdist_p(&self, other: &Tensor, p: f64) -> Result<Tensor> {
ffi_call!(flodl_cdist, self.handle, other.handle, p)
}
pub fn cosine_similarity(&self, other: &Tensor, dim: i64, eps: f64) -> Result<Tensor> {
ffi_call!(flodl_cosine_similarity, self.handle, other.handle, dim, eps)
}
pub fn to_dtype(&self, dtype: super::DType) -> Result<Tensor> {
ffi_call!(flodl_to_dtype, self.handle, dtype as i32)
}
pub fn all_finite(&self) -> Result<bool> {
let mut result: i32 = 0;
let err = unsafe { ffi::flodl_all_finite(self.handle, &mut result) };
check_err(err)?;
Ok(result != 0)
}
}
#[cfg(test)]
#[path = "ops_tests.rs"]
mod tests;