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 {
#[doc(hidden)]
pub fn from_raw(ptr: cudarc::driver::sys::CUdeviceptr, len: usize) -> Self {
Self { ptr, len }
}
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 {
cu_ctx_sync("GradSlice::to_cpu")?;
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 {
cu_ctx_sync("GradSlice::upload_from_cpu (pre)")?;
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
));
}
cu_ctx_sync("GradSlice::upload_from_cpu (post)")?;
}
Ok(())
}
}
fn cu_ctx_sync(what: &str) -> Result<(), String> {
let r = unsafe { cudarc::driver::sys::cuCtxSynchronize() };
if r != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!("{what}: cuCtxSynchronize failed: {r:?}"));
}
Ok(())
}
pub type WeightSlice = GradSlice;
use super::dtype::WeightDtype;
pub struct GpuByteBuffer {
data: cudarc::driver::CudaSlice<u8>,
len_bytes: usize,
cached_ptr: cudarc::driver::sys::CUdeviceptr,
}
impl GpuByteBuffer {
pub fn zeros(
stream: &Arc<cudarc::driver::CudaStream>,
len_bytes: usize,
) -> Result<Self, String> {
let data = stream
.alloc_zeros::<u8>(len_bytes)
.map_err(|e| format!("GPU alloc_zeros({len_bytes} bytes) failed: {e:?}"))?;
let cached_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _g) = data.device_ptr(stream);
ptr
};
Ok(Self {
data,
len_bytes,
cached_ptr,
})
}
pub fn cached_ptr(&self) -> cudarc::driver::sys::CUdeviceptr {
self.cached_ptr
}
pub fn len_bytes(&self) -> usize {
self.len_bytes
}
pub fn inner(&self) -> &cudarc::driver::CudaSlice<u8> {
&self.data
}
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 download_f64(
&self,
stream: &Arc<cudarc::driver::CudaStream>,
dst: &mut [f64],
) -> Result<(), String> {
assert_eq!(
std::mem::size_of_val(dst),
self.len_bytes,
"download_f64 size mismatch: dst={} f64s, gpu={} bytes",
dst.len(),
self.len_bytes
);
let bytes: &mut [u8] = bytemuck::cast_slice_mut(dst);
stream
.memcpy_dtoh(&self.data, bytes)
.map_err(|e| format!("GPU op failed: {:?}", e))
}
}
pub struct DtypedBuf {
inner: GpuByteBuffer,
n_elems: usize,
dtype: WeightDtype,
}
impl DtypedBuf {
pub fn zeros(
stream: &Arc<cudarc::driver::CudaStream>,
n_elems: usize,
dtype: WeightDtype,
) -> Result<Self, String> {
let inner = GpuByteBuffer::zeros(stream, n_elems * dtype.size_bytes())?;
Ok(Self {
inner,
n_elems,
dtype,
})
}
pub fn cached_ptr(&self) -> cudarc::driver::sys::CUdeviceptr {
self.inner.cached_ptr()
}
pub fn len_elems(&self) -> usize {
self.n_elems
}
pub fn dtype(&self) -> WeightDtype {
self.dtype
}
pub fn size_bytes(&self) -> usize {
self.n_elems * self.dtype.size_bytes()
}
pub fn zero(&mut self, stream: &Arc<cudarc::driver::CudaStream>) -> Result<(), String> {
self.inner.zero(stream)
}
pub fn upload_f32(
&self,
stream: &Arc<cudarc::driver::CudaStream>,
src: &[f32],
) -> Result<(), String> {
assert_eq!(src.len(), self.n_elems, "DtypedBuf upload size mismatch");
let ptr = self.inner.cached_ptr();
match self.dtype {
WeightDtype::F32 => {
let bytes: &[u8] = bytemuck::cast_slice(src);
cu_memcpy_htod_raw(stream, ptr, bytes)
}
WeightDtype::Bf16 => {
let buf: Vec<half::bf16> = src.iter().map(|&v| half::bf16::from_f32(v)).collect();
let bytes: &[u8] = bytemuck::cast_slice(&buf);
cu_memcpy_htod_raw(stream, ptr, bytes)
}
WeightDtype::F16 => {
let buf: Vec<half::f16> = src.iter().map(|&v| half::f16::from_f32(v)).collect();
let bytes: &[u8] = bytemuck::cast_slice(&buf);
cu_memcpy_htod_raw(stream, ptr, bytes)
}
}
}
pub fn download_f32(
&self,
stream: &Arc<cudarc::driver::CudaStream>,
dst: &mut [f32],
) -> Result<(), String> {
assert_eq!(dst.len(), self.n_elems, "DtypedBuf download size mismatch");
let ptr = self.inner.cached_ptr();
match self.dtype {
WeightDtype::F32 => {
let bytes: &mut [u8] = bytemuck::cast_slice_mut(dst);
cu_memcpy_dtoh_raw(stream, ptr, bytes)
}
WeightDtype::Bf16 => {
let mut buf = vec![half::bf16::ZERO; self.n_elems];
let bytes: &mut [u8] = bytemuck::cast_slice_mut(&mut buf);
cu_memcpy_dtoh_raw(stream, ptr, bytes)?;
for (d, &v) in dst.iter_mut().zip(&buf) {
*d = v.to_f32();
}
Ok(())
}
WeightDtype::F16 => {
let mut buf = vec![half::f16::ZERO; self.n_elems];
let bytes: &mut [u8] = bytemuck::cast_slice_mut(&mut buf);
cu_memcpy_dtoh_raw(stream, ptr, bytes)?;
for (d, &v) in dst.iter_mut().zip(&buf) {
*d = v.to_f32();
}
Ok(())
}
}
}
}
pub(crate) fn cu_memcpy_htod_raw(
stream: &Arc<cudarc::driver::CudaStream>,
dst: cudarc::driver::sys::CUdeviceptr,
bytes: &[u8],
) -> Result<(), String> {
if bytes.is_empty() {
return Ok(());
}
let r = unsafe {
cudarc::driver::sys::cuMemcpyHtoDAsync_v2(
dst,
bytes.as_ptr() as *const std::ffi::c_void,
bytes.len(),
stream.cu_stream(),
)
};
if r != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!("cuMemcpyHtoDAsync: {r:?}"));
}
stream
.synchronize()
.map_err(|e| format!("cuMemcpyHtoDAsync sync: {e:?}"))
}
pub(crate) fn cu_memcpy_dtoh_raw(
stream: &Arc<cudarc::driver::CudaStream>,
src: cudarc::driver::sys::CUdeviceptr,
bytes: &mut [u8],
) -> Result<(), String> {
if bytes.is_empty() {
return Ok(());
}
let r = unsafe {
cudarc::driver::sys::cuMemcpyDtoHAsync_v2(
bytes.as_mut_ptr() as *mut std::ffi::c_void,
src,
bytes.len(),
stream.cu_stream(),
)
};
if r != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(format!("cuMemcpyDtoHAsync: {r:?}"));
}
stream
.synchronize()
.map_err(|e| format!("cuMemcpyDtoHAsync sync: {e:?}"))
}
#[derive(Clone, Copy)]
pub struct WeightSliceDyn {
ptr: cudarc::driver::sys::CUdeviceptr,
len_elems: usize,
dtype: WeightDtype,
}
impl WeightSliceDyn {
pub fn from_byte_offset(
base: cudarc::driver::sys::CUdeviceptr,
byte_offset: usize,
len_elems: usize,
dtype: WeightDtype,
) -> Self {
Self {
ptr: base + byte_offset as u64,
len_elems,
dtype,
}
}
pub fn ptr(&self) -> cudarc::driver::sys::CUdeviceptr {
self.ptr
}
pub fn len_elems(&self) -> usize {
self.len_elems
}
pub fn dtype(&self) -> WeightDtype {
self.dtype
}
pub fn size_bytes(&self) -> usize {
self.len_elems * self.dtype.size_bytes()
}
pub fn download_to_f32(
&self,
stream: &Arc<cudarc::driver::CudaStream>,
dst: &mut [f32],
) -> Result<(), String> {
assert_eq!(dst.len(), self.len_elems, "size mismatch");
if self.len_elems == 0 {
return Ok(());
}
match self.dtype {
WeightDtype::F32 => {
let bytes: &mut [u8] = bytemuck::cast_slice_mut(dst);
cu_memcpy_dtoh_raw(stream, self.ptr, bytes)
}
WeightDtype::Bf16 => {
let mut buf = vec![half::bf16::ZERO; self.len_elems];
let bytes: &mut [u8] = bytemuck::cast_slice_mut(&mut buf);
cu_memcpy_dtoh_raw(stream, self.ptr, bytes)?;
for (d, &v) in dst.iter_mut().zip(&buf) {
*d = v.to_f32();
}
Ok(())
}
WeightDtype::F16 => {
let mut buf = vec![half::f16::ZERO; self.len_elems];
let bytes: &mut [u8] = bytemuck::cast_slice_mut(&mut buf);
cu_memcpy_dtoh_raw(stream, self.ptr, bytes)?;
for (d, &v) in dst.iter_mut().zip(&buf) {
*d = v.to_f32();
}
Ok(())
}
}
}
pub fn upload_from_cpu_f32(
&self,
stream: &Arc<cudarc::driver::CudaStream>,
src: &[f32],
) -> Result<(), String> {
assert_eq!(src.len(), self.len_elems, "size mismatch");
if self.len_elems == 0 {
return Ok(());
}
match self.dtype {
WeightDtype::F32 => self.upload_raw_bytes(stream, bytemuck::cast_slice(src)),
WeightDtype::Bf16 => {
let buf: Vec<half::bf16> = src.iter().map(|&v| half::bf16::from_f32(v)).collect();
self.upload_raw_bytes(stream, bytemuck::cast_slice(&buf))
}
WeightDtype::F16 => {
let buf: Vec<half::f16> = src.iter().map(|&v| half::f16::from_f32(v)).collect();
self.upload_raw_bytes(stream, bytemuck::cast_slice(&buf))
}
}
}
pub fn upload_raw_bytes(
&self,
stream: &Arc<cudarc::driver::CudaStream>,
bytes: &[u8],
) -> Result<(), String> {
assert_eq!(bytes.len(), self.size_bytes(), "byte size mismatch");
cu_memcpy_htod_raw(stream, self.ptr, bytes)
}
}
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 {
}