mod cuda;
pub mod cuda_event;
pub mod cuda_stream;
mod ops;
mod shape;
mod nn_ops;
pub use cuda::*;
pub use cuda_event::{CudaEvent, CudaEventFlags};
pub use cuda_stream::{CudaStream, StreamGuard};
pub use nn_ops::RnnParams;
use std::ffi::{c_void, CStr};
use std::fmt;
use std::ptr;
use std::sync::atomic::{AtomicU64, Ordering};
use flodl_sys::{self as ffi, FlodlTensor};
pub(super) static LIVE_TENSOR_COUNT: AtomicU64 = AtomicU64::new(0);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(i32)]
pub enum DType {
Float16 = ffi::FLODL_FLOAT16,
BFloat16 = ffi::FLODL_BFLOAT16,
Float32 = ffi::FLODL_FLOAT32,
Float64 = ffi::FLODL_FLOAT64,
Int32 = ffi::FLODL_INT32,
Int64 = ffi::FLODL_INT64,
}
impl DType {
fn from_raw(v: i32) -> Self {
match v {
ffi::FLODL_FLOAT16 => DType::Float16,
ffi::FLODL_BFLOAT16 => DType::BFloat16,
ffi::FLODL_FLOAT32 => DType::Float32,
ffi::FLODL_FLOAT64 => DType::Float64,
ffi::FLODL_INT32 => DType::Int32,
ffi::FLODL_INT64 => DType::Int64,
other => panic!(
"flodl: unknown dtype code {other} from the C++ shim — the \
Rust and flodl-sys dtype tables disagree; rebuild flodl-sys"
),
}
}
pub fn element_size(self) -> usize {
match self {
DType::Float16 | DType::BFloat16 => 2,
DType::Float32 | DType::Int32 => 4,
DType::Float64 | DType::Int64 => 8,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Device {
CPU,
CUDA(u8),
}
impl Device {
pub(crate) fn to_ffi(self) -> (i32, i32) {
match self {
Device::CPU => (ffi::FLODL_CPU, 0),
Device::CUDA(idx) => (ffi::FLODL_CUDA, idx as i32),
}
}
pub(crate) fn from_ffi(device_type: i32, device_index: i32) -> Self {
match device_type {
ffi::FLODL_CUDA => Device::CUDA(device_index as u8),
_ => Device::CPU,
}
}
pub fn is_cuda(&self) -> bool {
matches!(self, Device::CUDA(_))
}
pub fn index(&self) -> u8 {
match self {
Device::CPU => 0,
Device::CUDA(idx) => *idx,
}
}
}
impl fmt::Display for Device {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Device::CPU => write!(f, "cpu"),
Device::CUDA(0) => write!(f, "cuda"),
Device::CUDA(idx) => write!(f, "cuda:{}", idx),
}
}
}
#[derive(Debug, Clone)]
pub struct TensorError(String);
impl TensorError {
pub fn new(msg: &str) -> Self {
TensorError(msg.to_string())
}
pub fn is_cuda_oom(&self) -> bool {
self.0.contains("out of memory")
}
}
impl fmt::Display for TensorError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.0)
}
}
impl std::error::Error for TensorError {}
pub type Result<T> = std::result::Result<T, TensorError>;
pub(crate) fn check_err(err: *mut std::ffi::c_char) -> Result<()> {
if err.is_null() {
Ok(())
} else {
let msg = unsafe { CStr::from_ptr(err) }
.to_string_lossy()
.into_owned();
unsafe { ffi::flodl_free_string(err) };
Err(TensorError(msg))
}
}
macro_rules! ffi_call {
($ffi:ident $(, $arg:expr)* $(,)?) => {{
let mut handle: flodl_sys::FlodlTensor = ::std::ptr::null_mut();
let err = unsafe { flodl_sys::$ffi($($arg,)* &mut handle) };
$crate::tensor::check_err(err)?;
Ok($crate::tensor::Tensor::from_raw(handle))
}};
}
pub(crate) use ffi_call;
#[derive(Debug, Clone, Copy)]
pub struct TensorOptions {
pub dtype: DType,
pub device: Device,
}
impl Default for TensorOptions {
fn default() -> Self {
Self {
dtype: DType::Float32,
device: Device::CPU,
}
}
}
pub struct Tensor {
pub(crate) handle: FlodlTensor,
}
fn typed_bytes<T>(data: &[T]) -> &[u8] {
unsafe { std::slice::from_raw_parts(data.as_ptr() as *const u8, std::mem::size_of_val(data)) }
}
unsafe impl Send for Tensor {}
unsafe impl Sync for Tensor {}
impl Drop for Tensor {
fn drop(&mut self) {
if !self.handle.is_null() {
LIVE_TENSOR_COUNT.fetch_sub(1, Ordering::Relaxed);
unsafe { ffi::flodl_free_tensor(self.handle) };
}
}
}
impl Clone for Tensor {
fn clone(&self) -> Self {
let mut handle: FlodlTensor = ptr::null_mut();
let err = unsafe { ffi::flodl_shallow_clone(self.handle, &mut handle) };
if !err.is_null() {
let msg = unsafe { CStr::from_ptr(err) }
.to_string_lossy()
.into_owned();
unsafe { ffi::flodl_free_string(err) };
panic!("tensor clone failed: {}", msg);
}
Self::from_raw(handle)
}
}
impl Tensor {
pub(crate) fn from_raw(handle: FlodlTensor) -> Self {
debug_assert!(!handle.is_null());
LIVE_TENSOR_COUNT.fetch_add(1, Ordering::Relaxed);
Self { handle }
}
pub(crate) fn raw(&self) -> FlodlTensor {
self.handle
}
pub fn zeros(shape: &[i64], opts: TensorOptions) -> Result<Self> {
let mut shape = shape.to_vec();
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_zeros(
shape.as_mut_ptr(),
shape.len() as i32,
opts.dtype as i32,
dt, di,
&mut handle,
)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn ones(shape: &[i64], opts: TensorOptions) -> Result<Self> {
let mut shape = shape.to_vec();
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_ones(
shape.as_mut_ptr(),
shape.len() as i32,
opts.dtype as i32,
dt, di,
&mut handle,
)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn from_f32(data: &[f32], shape: &[i64], device: Device) -> Result<Self> {
Self::from_blob_impl("Tensor::from_f32", typed_bytes(data), shape, DType::Float32, device)
}
pub fn from_f64(data: &[f64], shape: &[i64], device: Device) -> Result<Self> {
Self::from_blob_impl("Tensor::from_f64", typed_bytes(data), shape, DType::Float64, device)
}
pub fn from_i64(data: &[i64], shape: &[i64], device: Device) -> Result<Self> {
Self::from_blob_impl("Tensor::from_i64", typed_bytes(data), shape, DType::Int64, device)
}
pub fn from_blob(data: &[u8], shape: &[i64], dtype: DType, device: Device) -> Result<Self> {
Self::from_blob_impl("Tensor::from_blob", data, shape, dtype, device)
}
fn from_blob_impl(
ctx: &str,
data: &[u8],
shape: &[i64],
dtype: DType,
device: Device,
) -> Result<Self> {
let numel = shape
.iter()
.try_fold(1i64, |acc, &d| if d < 0 { None } else { acc.checked_mul(d) })
.ok_or_else(|| {
TensorError::new(&format!(
"{ctx}: invalid shape {shape:?} (negative or overflowing dimension)"
))
})?;
let expected = (numel as usize)
.checked_mul(dtype.element_size())
.ok_or_else(|| {
TensorError::new(&format!(
"{ctx}: invalid shape {shape:?} (byte size overflows usize)"
))
})?;
if data.len() != expected {
return Err(TensorError::new(&format!(
"{ctx}: data is {} bytes, expected {expected} \
(numel={numel} × {} bytes/elem for {dtype:?})",
data.len(),
dtype.element_size(),
)));
}
let mut shape = shape.to_vec();
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = device.to_ffi();
let err = unsafe {
ffi::flodl_from_blob(
data.as_ptr() as *mut c_void,
shape.as_mut_ptr(),
shape.len() as i32,
dtype as i32,
dt, di,
&mut handle,
)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn zeros_like(t: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_zeros_like, t.handle)
}
pub fn ones_like(t: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_ones_like, t.handle)
}
pub fn full_like(t: &Tensor, value: f64) -> Result<Tensor> {
ffi_call!(flodl_full_like, t.handle, value)
}
pub fn rand_like(t: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_rand_like, t.handle)
}
pub fn randn_like(t: &Tensor) -> Result<Tensor> {
ffi_call!(flodl_randn_like, t.handle)
}
pub fn rand(shape: &[i64], opts: TensorOptions) -> Result<Self> {
let mut shape = shape.to_vec();
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_rand(
shape.as_mut_ptr(), shape.len() as i32,
opts.dtype as i32, dt, di,
&mut handle,
)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn randn(shape: &[i64], opts: TensorOptions) -> Result<Self> {
let mut shape = shape.to_vec();
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_randn(
shape.as_mut_ptr(), shape.len() as i32,
opts.dtype as i32, dt, di,
&mut handle,
)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn linspace(start: f64, end: f64, steps: i64, opts: TensorOptions) -> Result<Self> {
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_linspace(start, end, steps, opts.dtype as i32, dt, di, &mut handle)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn arange(start: f64, end: f64, step: f64, opts: TensorOptions) -> Result<Self> {
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_arange(start, end, step, opts.dtype as i32, dt, di, &mut handle)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn eye(n: i64, opts: TensorOptions) -> Result<Self> {
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_eye(n, opts.dtype as i32, dt, di, &mut handle)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn full(shape: &[i64], value: f64, opts: TensorOptions) -> Result<Self> {
let mut shape = shape.to_vec();
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_full(
shape.as_mut_ptr(), shape.len() as i32, value,
opts.dtype as i32, dt, di, &mut handle,
)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn randperm(n: i64, opts: TensorOptions) -> Result<Self> {
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_randperm(n, opts.dtype as i32, dt, di, &mut handle)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn randint(low: i64, high: i64, shape: &[i64], opts: TensorOptions) -> Result<Self> {
let mut shape = shape.to_vec();
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_randint(
low, high,
shape.as_mut_ptr(), shape.len() as i32,
opts.dtype as i32, dt, di,
&mut handle,
)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn empty(shape: &[i64], opts: TensorOptions) -> Result<Self> {
let mut shape = shape.to_vec();
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = opts.device.to_ffi();
let err = unsafe {
ffi::flodl_empty(
shape.as_mut_ptr(), shape.len() as i32,
opts.dtype as i32, dt, di,
&mut handle,
)
};
check_err(err)?;
Ok(Self::from_raw(handle))
}
pub fn one_hot(&self, num_classes: i64) -> Result<Tensor> {
ffi_call!(flodl_one_hot, self.handle, num_classes)
}
pub fn bernoulli(&self) -> Result<Tensor> {
ffi_call!(flodl_bernoulli, self.handle)
}
pub fn ndim(&self) -> usize {
unsafe { ffi::flodl_ndim(self.handle) as usize }
}
pub fn shape(&self) -> Vec<i64> {
let n = self.ndim();
(0..n)
.map(|i| unsafe { ffi::flodl_shape(self.handle, i as i32) })
.collect()
}
pub fn numel(&self) -> i64 {
unsafe { ffi::flodl_numel(self.handle) }
}
pub fn nbytes(&self) -> usize {
self.numel() as usize * self.dtype().element_size()
}
pub fn storage_nbytes(&self) -> usize {
unsafe { ffi::flodl_storage_nbytes(self.handle) as usize }
}
pub fn dtype(&self) -> DType {
DType::from_raw(unsafe { ffi::flodl_dtype(self.handle) })
}
pub fn device(&self) -> Device {
let dt = unsafe { ffi::flodl_device_type(self.handle) };
let di = unsafe { ffi::flodl_device_index(self.handle) };
Device::from_ffi(dt, di)
}
pub fn to_f32_vec(&self) -> Result<Vec<f32>> {
if self.dtype() != DType::Float32 {
return self.to_dtype(DType::Float32)?.to_f32_vec();
}
let n = self.numel() as usize;
let mut buf = vec![0f32; n];
let bytes = (n * 4) as i64;
let err = unsafe {
ffi::flodl_copy_data(self.handle, buf.as_mut_ptr() as *mut c_void, bytes)
};
check_err(err)?;
Ok(buf)
}
pub fn to_blob(&self) -> Result<Vec<u8>> {
let bytes = self.numel() as usize * self.dtype().element_size();
let mut buf = vec![0u8; bytes];
let err = unsafe {
ffi::flodl_copy_data(self.handle, buf.as_mut_ptr() as *mut c_void, bytes as i64)
};
check_err(err)?;
Ok(buf)
}
pub fn to_f64_vec(&self) -> Result<Vec<f64>> {
if self.dtype() != DType::Float64 {
return self.to_dtype(DType::Float64)?.to_f64_vec();
}
let n = self.numel() as usize;
let mut buf = vec![0.0f64; n];
let bytes = (n * 8) as i64;
let err = unsafe {
ffi::flodl_copy_data(self.handle, buf.as_mut_ptr() as *mut c_void, bytes)
};
check_err(err)?;
Ok(buf)
}
pub fn to_i64_vec(&self) -> Result<Vec<i64>> {
if self.dtype() != DType::Int64 {
return self.to_dtype(DType::Int64)?.to_i64_vec();
}
let n = self.numel() as usize;
let mut buf = vec![0i64; n];
let bytes = (n * 8) as i64;
let err = unsafe {
ffi::flodl_copy_data(self.handle, buf.as_mut_ptr() as *mut c_void, bytes)
};
check_err(err)?;
Ok(buf)
}
pub fn item(&self) -> Result<f64> {
if self.numel() != 1 {
return Err(TensorError::new(&format!(
"item() requires exactly 1 element, got {} (shape {:?})",
self.numel(), self.shape()
)));
}
if self.dtype() != DType::Float64 {
return self.to_dtype(DType::Float64)?.item();
}
let mut buf = [0.0f64; 1];
let err = unsafe {
ffi::flodl_copy_data(self.handle, buf.as_mut_ptr() as *mut c_void, 8)
};
check_err(err)?;
Ok(buf[0])
}
pub fn to_device(&self, device: Device) -> Result<Tensor> {
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = device.to_ffi();
let err = unsafe { ffi::flodl_to_device(self.handle, dt, di, &mut handle) };
check_err(err)?;
Ok(Tensor::from_raw(handle))
}
pub fn to_device_of(&self, other: &Tensor) -> Result<Tensor> {
let target = other.device();
if self.device() == target {
return Ok(self.clone());
}
self.to_device(target)
}
pub fn to_device_async(&self, device: Device) -> Result<Tensor> {
let mut handle: FlodlTensor = ptr::null_mut();
let (dt, di) = device.to_ffi();
let err = unsafe { ffi::flodl_to_device_async(self.handle, dt, di, &mut handle) };
check_err(err)?;
Ok(Tensor::from_raw(handle))
}
pub fn record_stream(&self, stream: &crate::tensor::cuda_stream::CudaStream) -> Result<()> {
let err = unsafe {
ffi::flodl_tensor_record_stream(self.handle, stream.as_ptr())
};
check_err(err)
}
pub fn set_requires_grad(&self, requires_grad: bool) -> Result<Tensor> {
ffi_call!(flodl_set_requires_grad, self.handle, requires_grad as i32)
}
pub fn requires_grad(&self) -> bool {
unsafe { ffi::flodl_requires_grad(self.handle) != 0 }
}
pub fn backward(&self) -> Result<()> {
let err = unsafe { ffi::flodl_backward(self.handle) };
check_err(err)
}
pub fn grad(&self) -> Option<Tensor> {
let mut handle: FlodlTensor = ptr::null_mut();
let err = unsafe { ffi::flodl_grad(self.handle, &mut handle) };
if !err.is_null() {
let msg = unsafe { CStr::from_ptr(err) }.to_string_lossy().into_owned();
unsafe { ffi::flodl_free_string(err) };
panic!("Tensor::grad failed: {msg}");
}
if handle.is_null() {
None
} else {
Some(Tensor::from_raw(handle))
}
}
pub fn set_grad(&self, grad: &Tensor) -> Result<()> {
let err = unsafe { ffi::flodl_set_grad(self.handle, grad.handle) };
check_err(err)
}
pub fn zero_grad(&self) -> Result<()> {
let err = unsafe { ffi::flodl_zero_grad(self.handle) };
check_err(err)
}
pub fn zero_grad_set_to_none(&self) {
unsafe { ffi::flodl_zero_grad_set_to_none(self.handle) }
}
pub fn clip_grad_norm_fused(params: &[Tensor], max_norm: f64) -> Result<f64> {
if params.is_empty() {
return Ok(0.0);
}
let mut handles: Vec<FlodlTensor> = params.iter().map(|t| t.handle).collect();
let mut total_norm: f64 = 0.0;
let err = unsafe {
ffi::flodl_clip_grad_norm(
handles.as_mut_ptr(),
handles.len() as i32,
max_norm,
&mut total_norm,
)
};
check_err(err)?;
Ok(total_norm)
}
pub fn is_leaf(&self) -> bool {
unsafe { ffi::flodl_is_leaf(self.handle) != 0 }
}
pub fn ensure_grad_accumulator(&self) -> Result<Option<GradAccumulatorHandle>> {
let mut handle: *mut std::ffi::c_void = std::ptr::null_mut();
let err = unsafe { ffi::flodl_ensure_grad_accumulator(self.handle, &mut handle) };
check_err(err)?;
if handle.is_null() {
Ok(None)
} else {
Ok(Some(GradAccumulatorHandle { handle }))
}
}
pub fn autograd_node_count(&self) -> i64 {
unsafe { ffi::flodl_autograd_node_count(self.handle) }
}
pub fn detach(&self) -> Result<Tensor> {
ffi_call!(flodl_detach, self.handle)
}
pub fn detach_(&self) -> Result<()> {
let err = unsafe { ffi::flodl_detach_(self.handle) };
check_err(err)
}
pub fn copy(&self) -> Result<Tensor> {
ffi_call!(flodl_deep_clone, self.handle)
}
pub fn add_(&self, other: &Tensor) -> Result<()> {
let err = unsafe { ffi::flodl_add_(self.handle, other.handle) };
check_err(err)
}
pub fn sub_(&self, other: &Tensor) -> Result<()> {
let err = unsafe { ffi::flodl_sub_(self.handle, other.handle) };
check_err(err)
}
pub fn mul_scalar_(&self, scalar: f64) -> Result<()> {
let err = unsafe { ffi::flodl_mul_scalar_(self.handle, scalar) };
check_err(err)
}
pub fn add_scalar_(&self, scalar: f64) -> Result<()> {
let err = unsafe { ffi::flodl_add_scalar_(self.handle, scalar) };
check_err(err)
}
pub fn zero_(&self) -> Result<()> {
let err = unsafe { ffi::flodl_zero_(self.handle) };
check_err(err)
}
pub fn mul_(&self, other: &Tensor) -> Result<()> {
let err = unsafe { ffi::flodl_mul_(self.handle, other.handle) };
check_err(err)
}
pub fn div_scalar_(&self, scalar: f64) -> Result<()> {
let err = unsafe { ffi::flodl_div_scalar_(self.handle, scalar) };
check_err(err)
}
pub fn div_(&self, other: &Tensor) -> Result<()> {
let err = unsafe { ffi::flodl_div_(self.handle, other.handle) };
check_err(err)
}
pub fn fill_(&self, value: f64) -> Result<()> {
let err = unsafe { ffi::flodl_fill_(self.handle, value) };
check_err(err)
}
pub fn copy_(&self, src: &Tensor, non_blocking: bool) -> Result<()> {
let err = unsafe { ffi::flodl_copy_(self.handle, src.handle, non_blocking as i32) };
check_err(err)
}
#[allow(clippy::too_many_arguments)]
pub fn adam_step(
&self, grad: &Tensor, m: &Tensor, v: &Tensor,
lr: f64, beta1: f64, beta2: f64, eps: f64,
weight_decay: f64, step: i64,
) -> Result<()> {
let err = unsafe {
ffi::flodl_adam_step(
self.handle, grad.handle, m.handle, v.handle,
lr, beta1, beta2, eps, weight_decay, step,
)
};
check_err(err)
}
#[allow(clippy::too_many_arguments)]
pub fn adam_step_batched(
params: &[Tensor], grads: &[Tensor], ms: &[Tensor], vs: &[Tensor],
lrs: &mut [f64], beta1: f64, beta2: f64, eps: f64,
weight_decay: f64, step: i64,
) -> Result<()> {
let count = params.len() as i32;
let mut p_handles: Vec<FlodlTensor> = params.iter().map(|t| t.handle).collect();
let mut g_handles: Vec<FlodlTensor> = grads.iter().map(|t| t.handle).collect();
let mut m_handles: Vec<FlodlTensor> = ms.iter().map(|t| t.handle).collect();
let mut v_handles: Vec<FlodlTensor> = vs.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_adam_step_batched(
p_handles.as_mut_ptr(), g_handles.as_mut_ptr(),
m_handles.as_mut_ptr(), v_handles.as_mut_ptr(),
lrs.as_mut_ptr(), count,
beta1, beta2, eps, weight_decay, step,
)
};
check_err(err)
}
#[allow(clippy::too_many_arguments)]
pub fn fused_adam_(
params: &[Tensor], grads: &[Tensor], exp_avgs: &[Tensor], exp_avg_sqs: &[Tensor],
lr: f64, beta1: f64, beta2: f64, eps: f64,
weight_decay: f64, steps: &[i64],
grad_scale: Option<&Tensor>, found_inf: Option<&Tensor>,
) -> Result<()> {
if params.is_empty() { return Ok(()); }
if steps.len() != params.len() {
return Err(TensorError::new(&format!(
"fused_adam_: steps length {} does not match params length {}",
steps.len(), params.len()
)));
}
let count = params.len() as i32;
let mut p = Self::handles(params);
let mut g = Self::handles(grads);
let mut m = Self::handles(exp_avgs);
let mut v = Self::handles(exp_avg_sqs);
let gs = grad_scale.map_or(ptr::null_mut(), |t| t.handle);
let fi = found_inf.map_or(ptr::null_mut(), |t| t.handle);
let err = unsafe {
ffi::flodl_fused_adam_(
p.as_mut_ptr(), g.as_mut_ptr(), m.as_mut_ptr(), v.as_mut_ptr(),
count, lr, beta1, beta2, eps, weight_decay, steps.as_ptr(), gs, fi,
)
};
check_err(err)
}
#[allow(clippy::too_many_arguments)]
pub fn fused_adamw_(
params: &[Tensor], grads: &[Tensor], exp_avgs: &[Tensor], exp_avg_sqs: &[Tensor],
lr: f64, beta1: f64, beta2: f64, eps: f64,
weight_decay: f64, steps: &[i64],
grad_scale: Option<&Tensor>, found_inf: Option<&Tensor>,
) -> Result<()> {
if params.is_empty() { return Ok(()); }
if steps.len() != params.len() {
return Err(TensorError::new(&format!(
"fused_adamw_: steps length {} does not match params length {}",
steps.len(), params.len()
)));
}
let count = params.len() as i32;
let mut p = Self::handles(params);
let mut g = Self::handles(grads);
let mut m = Self::handles(exp_avgs);
let mut v = Self::handles(exp_avg_sqs);
let gs = grad_scale.map_or(ptr::null_mut(), |t| t.handle);
let fi = found_inf.map_or(ptr::null_mut(), |t| t.handle);
let err = unsafe {
ffi::flodl_fused_adamw_(
p.as_mut_ptr(), g.as_mut_ptr(), m.as_mut_ptr(), v.as_mut_ptr(),
count, lr, beta1, beta2, eps, weight_decay, steps.as_ptr(), gs, fi,
)
};
check_err(err)
}
fn handles(tensors: &[Tensor]) -> Vec<FlodlTensor> {
tensors.iter().map(|t| t.handle).collect()
}
pub fn foreach_add_scalar_(tensors: &[Tensor], scalar: f64) -> Result<()> {
if tensors.is_empty() { return Ok(()); }
let mut handles: Vec<FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_foreach_add_scalar_(handles.as_mut_ptr(), handles.len() as i32, scalar)
};
check_err(err)
}
pub fn foreach_mul_scalar_(tensors: &[Tensor], scalar: f64) -> Result<()> {
if tensors.is_empty() { return Ok(()); }
let mut handles: Vec<FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_foreach_mul_scalar_(handles.as_mut_ptr(), handles.len() as i32, scalar)
};
check_err(err)
}
pub fn foreach_zero_(tensors: &[Tensor]) -> Result<()> {
if tensors.is_empty() { return Ok(()); }
let mut handles: Vec<FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_foreach_zero_(handles.as_mut_ptr(), handles.len() as i32)
};
check_err(err)
}
pub fn foreach_add_list_(tensors1: &[Tensor], tensors2: &[Tensor], alpha: f64) -> Result<()> {
if tensors1.is_empty() { return Ok(()); }
if tensors1.len() != tensors2.len() {
return Err(TensorError::new(&format!(
"foreach_add_list_: list length mismatch ({} vs {})",
tensors1.len(), tensors2.len(),
)));
}
let mut h1: Vec<FlodlTensor> = tensors1.iter().map(|t| t.handle).collect();
let mut h2: Vec<FlodlTensor> = tensors2.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_foreach_add_list_(
h1.as_mut_ptr(), h2.as_mut_ptr(), h1.len() as i32, alpha,
)
};
check_err(err)
}
pub fn foreach_norm(tensors: &[Tensor], ord: f64) -> Result<Vec<Tensor>> {
if tensors.is_empty() { return Ok(vec![]); }
let mut handles: Vec<FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let mut results: Vec<FlodlTensor> = vec![ptr::null_mut(); tensors.len()];
let err = unsafe {
ffi::flodl_foreach_norm(
handles.as_mut_ptr(), handles.len() as i32, ord,
results.as_mut_ptr(),
)
};
check_err(err)?;
Ok(results.into_iter().map(Tensor::from_raw).collect())
}
pub fn foreach_lerp_scalar_(tensors1: &[Tensor], tensors2: &[Tensor], weight: f64) -> Result<()> {
if tensors1.is_empty() { return Ok(()); }
if tensors1.len() != tensors2.len() {
return Err(TensorError::new(&format!(
"foreach_lerp_scalar_: list length mismatch ({} vs {})",
tensors1.len(), tensors2.len(),
)));
}
let mut h1: Vec<FlodlTensor> = tensors1.iter().map(|t| t.handle).collect();
let mut h2: Vec<FlodlTensor> = tensors2.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_foreach_lerp_scalar_(
h1.as_mut_ptr(), h2.as_mut_ptr(), h1.len() as i32, weight,
)
};
check_err(err)
}
pub fn foreach_sqrt_(tensors: &[Tensor]) -> Result<()> {
if tensors.is_empty() { return Ok(()); }
let mut handles: Vec<FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_foreach_sqrt_(handles.as_mut_ptr(), handles.len() as i32)
};
check_err(err)
}
pub fn pin_memory(&self) -> Result<Tensor> {
ffi_call!(flodl_pin_memory, self.handle)
}
pub fn is_pinned(&self) -> bool {
unsafe { ffi::flodl_is_pinned(self.handle) != 0 }
}
pub fn to_channels_last(&self) -> Result<Tensor> {
ffi_call!(flodl_to_channels_last, self.handle)
}
pub fn is_channels_last(&self) -> bool {
unsafe { ffi::flodl_is_channels_last(self.handle) != 0 }
}
pub fn is_contiguous(&self) -> bool {
unsafe { ffi::flodl_is_contiguous(self.handle) != 0 }
}
}
pub struct GradAccumulatorHandle {
handle: *mut std::ffi::c_void,
}
unsafe impl Send for GradAccumulatorHandle {}
unsafe impl Sync for GradAccumulatorHandle {}
impl Drop for GradAccumulatorHandle {
fn drop(&mut self) {
if !self.handle.is_null() {
unsafe { ffi::flodl_grad_accumulator_delete(self.handle) };
self.handle = std::ptr::null_mut();
}
}
}
impl fmt::Debug for Tensor {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"Tensor({:?}, {:?}, {:?})",
self.shape(),
self.dtype(),
self.device()
)
}
}
#[cfg(test)]
pub fn test_device() -> Device {
use std::sync::Once;
static PRINT: Once = Once::new();
let dev = if cfg!(feature = "cuda") && cuda_available() { Device::CUDA(0) } else { Device::CPU };
PRINT.call_once(|| eprintln!("\n*** flodl test device: {} ***\n", dev));
dev
}
#[cfg(test)]
pub fn test_opts() -> TensorOptions {
TensorOptions { dtype: DType::Float32, device: test_device() }
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;