use std::sync::Arc;
pub struct GpuBuffer {
data: cudarc::driver::CudaSlice<f32>,
len: usize,
cached_ptr: cudarc::driver::sys::CUdeviceptr,
}
impl GpuBuffer {
pub fn zeros(stream: &Arc<cudarc::driver::CudaStream>, len: usize) -> Result<Self, String> {
let data = stream
.alloc_zeros::<f32>(len)
.map_err(|e| format!("GPU alloc_zeros({}) failed: {:?}", len, e))?;
let cached_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _guard) = data.device_ptr(stream);
ptr
};
Ok(Self {
data,
len,
cached_ptr,
})
}
pub fn from_cpu(stream: &Arc<cudarc::driver::CudaStream>, src: &[f32]) -> Result<Self, String> {
let data = stream
.clone_htod(src)
.map_err(|e| format!("GPU upload({} floats) failed: {:?}", src.len(), e))?;
let cached_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _guard) = data.device_ptr(stream);
ptr
};
Ok(Self {
len: src.len(),
data,
cached_ptr,
})
}
pub fn to_cpu(&self, stream: &Arc<cudarc::driver::CudaStream>) -> Result<Vec<f32>, String> {
stream
.clone_dtoh(&self.data)
.map_err(|e| format!("GPU download({} floats) failed: {:?}", self.len, e))
}
pub fn upload(
&mut self,
stream: &Arc<cudarc::driver::CudaStream>,
src: &[f32],
) -> Result<(), String> {
assert_eq!(
src.len(),
self.len,
"upload size mismatch: src={} gpu={}",
src.len(),
self.len
);
stream
.memcpy_htod(src, &mut self.data)
.map_err(|e| format!("GPU op failed: {:?}", e))
}
pub fn download(
&self,
stream: &Arc<cudarc::driver::CudaStream>,
dst: &mut [f32],
) -> Result<(), String> {
assert_eq!(
dst.len(),
self.len,
"download size mismatch: dst={} gpu={}",
dst.len(),
self.len
);
stream
.memcpy_dtoh(&self.data, dst)
.map_err(|e| format!("GPU op failed: {:?}", e))
}
pub fn zero(&mut self, stream: &Arc<cudarc::driver::CudaStream>) -> Result<(), String> {
stream
.memset_zeros(&mut self.data)
.map_err(|e| format!("GPU op failed: {:?}", e))
}
pub fn copy_from(
&mut self,
src: &GpuBuffer,
stream: &Arc<cudarc::driver::CudaStream>,
) -> Result<(), String> {
assert_eq!(
self.len, src.len,
"D2D copy size mismatch: dst={} src={}",
self.len, src.len
);
stream
.memcpy_dtod(&src.data, &mut self.data)
.map_err(|e| format!("GPU op failed: {:?}", e))
}
pub fn copy_from_raw(
&mut self,
src: &GpuBuffer,
stream: &Arc<cudarc::driver::CudaStream>,
) -> Result<(), String> {
assert_eq!(
self.len, src.len,
"D2D copy size mismatch: dst={} src={}",
self.len, src.len
);
if self.len > 0 {
let byte_count = self.len * std::mem::size_of::<f32>();
let result = unsafe {
cudarc::driver::sys::cuMemcpyDtoDAsync_v2(
self.cached_ptr,
src.cached_ptr,
byte_count,
stream.cu_stream(),
)
};
if result != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!(
"D2D copy_raw({} floats) failed: {:?}",
self.len, result
));
}
}
Ok(())
}
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn inner(&self) -> &cudarc::driver::CudaSlice<f32> {
&self.data
}
pub fn inner_mut(&mut self) -> &mut cudarc::driver::CudaSlice<f32> {
&mut self.data
}
pub fn raw_ptr(
&self,
_stream: &std::sync::Arc<cudarc::driver::CudaStream>,
) -> cudarc::driver::sys::CUdeviceptr {
self.cached_ptr
}
pub fn size_bytes(&self) -> usize {
self.len * std::mem::size_of::<f32>()
}
pub fn cached_ptr(&self) -> cudarc::driver::sys::CUdeviceptr {
self.cached_ptr
}
pub fn raw_ptr_at(
&self,
_stream: &std::sync::Arc<cudarc::driver::CudaStream>,
offset: usize,
) -> cudarc::driver::sys::CUdeviceptr {
assert!(
offset < self.len,
"raw_ptr_at offset {} >= len {}",
offset,
self.len
);
let byte_off = (offset * std::mem::size_of::<f32>()) as u64;
self.cached_ptr + byte_off
}
pub fn inner_at(&self, offset: usize) -> cudarc::driver::sys::CUdeviceptr {
assert!(
offset < self.len,
"inner_at offset {} >= len {}",
offset,
self.len
);
self.cached_ptr + (offset * std::mem::size_of::<f32>()) as u64
}
pub fn inner_mut_at(&mut self, offset: usize) -> cudarc::driver::sys::CUdeviceptr {
assert!(
offset < self.len,
"inner_mut_at offset {} >= len {}",
offset,
self.len
);
self.cached_ptr + (offset * std::mem::size_of::<f32>()) as u64
}
}
pub struct GradSlice {
ptr: cudarc::driver::sys::CUdeviceptr,
len: usize,
}
impl GradSlice {
pub fn ptr(&self) -> cudarc::driver::sys::CUdeviceptr {
self.ptr
}
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn raw_ptr(
&self,
_stream: &std::sync::Arc<cudarc::driver::CudaStream>,
) -> cudarc::driver::sys::CUdeviceptr {
self.ptr
}
pub fn inner(&self) -> &cudarc::driver::sys::CUdeviceptr {
&self.ptr
}
pub fn from_offset(base: cudarc::driver::sys::CUdeviceptr, offset: usize, len: usize) -> Self {
Self {
ptr: base + (offset * std::mem::size_of::<f32>()) as u64,
len,
}
}
pub fn size_bytes(&self) -> usize {
self.len * std::mem::size_of::<f32>()
}
pub fn to_cpu(&self) -> Result<Vec<f32>, String> {
let mut dst = vec![0.0f32; self.len];
if self.len > 0 {
let byte_count = self.len * std::mem::size_of::<f32>();
let result = unsafe {
cudarc::driver::sys::cuMemcpyDtoH_v2(
dst.as_mut_ptr() as *mut std::ffi::c_void,
self.ptr,
byte_count,
)
};
if result != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!(
"GradSlice::to_cpu({} floats) failed: {:?}",
self.len, result
));
}
}
Ok(dst)
}
pub fn upload_from_cpu(&self, src: &[f32]) -> Result<(), String> {
assert_eq!(
src.len(),
self.len,
"GradSlice upload size mismatch: src={} slice={}",
src.len(),
self.len
);
if self.len > 0 {
let byte_count = self.len * std::mem::size_of::<f32>();
let result = unsafe {
cudarc::driver::sys::cuMemcpyHtoD_v2(
self.ptr,
src.as_ptr() as *const std::ffi::c_void,
byte_count,
)
};
if result != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!(
"GradSlice::upload_from_cpu({} floats) failed: {:?}",
self.len, result
));
}
}
Ok(())
}
}
pub type WeightSlice = GradSlice;
impl std::fmt::Debug for GpuBuffer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"GpuBuffer({} floats, {} KB)",
self.len,
self.size_bytes() / 1024
)
}
}
#[cfg(test)]
mod tests {
}