use std::{fmt, num::NonZeroUsize, sync::Arc};
use crate::device::{Device, DeviceBuffer, OperationError};
use super::DenseMatrix;
pub struct SparseMatrix<D: Device> {
pub buf: D::BufferI32,
pub nnz: usize,
pub single_size: usize,
pub batch_size: Option<NonZeroUsize>,
}
impl<D: Device> fmt::Debug for SparseMatrix<D> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}x{:?}xf32", self.single_size, self.batch_size)
}
}
impl<D: Device> SparseMatrix<D> {
pub fn zeroed(
device: Arc<D>,
single_size: usize,
nnz: usize,
batch_size: Option<usize>,
) -> Result<Self, D::DeviceError> {
let buf = D::BufferI32::new(device, nnz * batch_size.unwrap_or(1))?;
Ok(Self { buf, single_size, nnz, batch_size: batch_size.map(|b| b.try_into().unwrap()) })
}
pub fn allocated_size(&self) -> usize {
self.buf.size()
}
pub fn set_batch_size(&mut self, batch_size: Option<usize>) -> Result<(), D::DeviceError> {
let new_size = self.nnz * batch_size.unwrap_or(1);
if new_size > self.allocated_size() {
self.buf = D::BufferI32::new(self.buf.device(), new_size)?;
} else if batch_size != self.batch_size() {
self.buf.set_zero()?;
}
self.batch_size = batch_size.map(|x| NonZeroUsize::new(x).unwrap());
Ok(())
}
pub fn swap_with(&mut self, other: &mut Self) -> Result<(), D::DeviceError> {
if self.single_size != other.single_size || self.nnz != other.nnz {
return Err(D::DeviceError::default());
}
std::mem::swap(self, other);
Ok(())
}
pub fn copy_from(&mut self, other: &Self) -> Result<(), D::DeviceError> {
if self.single_size != other.single_size || self.nnz != other.nnz {
return Err(D::DeviceError::default());
}
self.set_batch_size(other.batch_size())?;
self.buf.load_from_device(&other.buf, other.nnz * other.batch_size().unwrap_or(1))
}
pub fn single_size(&self) -> usize {
self.single_size
}
pub fn batch_size(&self) -> Option<usize> {
self.batch_size.map(NonZeroUsize::get)
}
pub fn size(&self) -> usize {
self.single_size * self.batch_size().unwrap_or(1)
}
pub unsafe fn load_from_slice(
&mut self,
nnz: usize,
batch_size: Option<usize>,
buf: &[i32],
) -> Result<(), D::DeviceError> {
if self.nnz != nnz || nnz * batch_size.unwrap_or(1) != buf.len() {
return Err(D::DeviceError::default());
}
self.set_batch_size(batch_size)?;
self.buf.load_from_slice(buf)
}
pub unsafe fn load_non_blocking_from_host(
&mut self,
nnz: usize,
batch_size: Option<usize>,
buf: &[i32],
) -> Result<(), D::DeviceError> {
if self.nnz != nnz || nnz * batch_size.unwrap_or(1) != buf.len() {
return Err(D::DeviceError::default());
}
self.set_batch_size(batch_size)?;
unsafe { self.buf.load_non_blocking_from_host(buf) }
}
pub fn copy_into_dense(&self, dst: &mut DenseMatrix<D>) -> Result<(), OperationError<D::DeviceError>> {
let batch_size = self.batch_size();
let size = self.single_size();
if size != dst.single_size() {
return Err(OperationError::InvalidTensorFormat);
}
dst.set_batch_size(batch_size)?;
D::sparse_to_dense(batch_size.unwrap_or(1), size, self.nnz, &self.buf, &mut dst.buf)
}
}