use super::Error;
use crate::{
DlpackElement, ManagedTensorBase,
allocation::{dynamic, fixed},
ffi::DLDevice,
metadata::{Copied, Dynamic, Fixed},
};
use cudarc::driver::{CudaSlice, CudaStream, DevicePtr};
use std::{os::raw::c_void, sync::Arc};
impl<T: DlpackElement, M: ManagedTensorBase> TryFrom<Box<CudaSlice<T>>>
for fixed::Initialized<M, 1>
{
type Error = Error;
fn try_from(slice: Box<CudaSlice<T>>) -> Result<Self, Self::Error> {
let len = i64::try_from(slice.len()).map_err(|source| Error::LengthOverflow {
len: slice.len(),
source,
})?;
let device_id =
i32::try_from(slice.ordinal()).map_err(|source| Error::DeviceIdOverflow {
ordinal: slice.ordinal(),
source,
})?;
let data_ptr = device_ptr_of(&slice);
let prepared = Fixed::new(Copied([len]), Copied([1])).prepare::<M>()?;
let mut initialized = prepared.initialize(slice);
initialized.set_device(DLDevice::cuda(device_id));
initialized.set_data(data_ptr);
initialized.set_dtype(T::DTYPE);
Ok(initialized)
}
}
#[allow(clippy::type_complexity)]
pub fn from_cuda_slice<T: DlpackElement, M: ManagedTensorBase>(
slice: Box<CudaSlice<T>>,
shape: &[i64],
strides: &[i64],
) -> Result<(dynamic::Initialized<M>, Arc<CudaStream>), Error> {
let device_id = i32::try_from(slice.ordinal()).map_err(|source| Error::DeviceIdOverflow {
ordinal: slice.ordinal(),
source,
})?;
let stream = slice.stream().clone();
let data_ptr = device_ptr_of(&slice);
let prepared = Dynamic::new(Copied(shape), Copied(strides)).prepare::<M>()?;
let mut initialized = prepared
.initialize(slice)
.map_err(crate::metadata::Error::from)?;
initialized.set_device(DLDevice::cuda(device_id));
initialized.set_dtype(T::DTYPE);
initialized.set_data(data_ptr);
Ok((initialized, stream))
}
fn device_ptr_of<T>(slice: &CudaSlice<T>) -> *mut c_void {
let stream = slice.stream().clone();
let (cu_ptr, sync) = slice.device_ptr(&stream);
drop(sync);
cu_ptr as usize as *mut c_void
}