use oxiblas_core::scalar::{Field, Scalar};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CscError {
InvalidColPtrs {
expected: usize,
actual: usize,
},
LengthMismatch {
values_len: usize,
row_indices_len: usize,
},
InvalidRowIndex {
index: usize,
nrows: usize,
},
InvalidColPtrOrder,
DuplicateEntry {
row: usize,
col: usize,
},
}
impl core::fmt::Display for CscError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::InvalidColPtrs { expected, actual } => {
write!(
f,
"Invalid col_ptrs length: expected {expected}, got {actual}"
)
}
Self::LengthMismatch {
values_len,
row_indices_len,
} => {
write!(
f,
"Length mismatch: values={values_len}, row_indices={row_indices_len}"
)
}
Self::InvalidRowIndex { index, nrows } => {
write!(f, "Row index {index} out of bounds for {nrows} rows")
}
Self::InvalidColPtrOrder => {
write!(f, "Column pointers must be monotonically increasing")
}
Self::DuplicateEntry { row, col } => {
write!(f, "Duplicate entry at ({row}, {col})")
}
}
}
}
impl std::error::Error for CscError {}
#[derive(Debug, Clone)]
pub struct CscMatrix<T: Scalar> {
nrows: usize,
ncols: usize,
col_ptrs: Vec<usize>,
row_indices: Vec<usize>,
values: Vec<T>,
}
impl<T: Scalar + Clone> CscMatrix<T> {
pub fn new(
nrows: usize,
ncols: usize,
col_ptrs: Vec<usize>,
row_indices: Vec<usize>,
values: Vec<T>,
) -> Result<Self, CscError> {
if col_ptrs.len() != ncols + 1 {
return Err(CscError::InvalidColPtrs {
expected: ncols + 1,
actual: col_ptrs.len(),
});
}
if values.len() != row_indices.len() {
return Err(CscError::LengthMismatch {
values_len: values.len(),
row_indices_len: row_indices.len(),
});
}
for i in 1..col_ptrs.len() {
if col_ptrs[i] < col_ptrs[i - 1] {
return Err(CscError::InvalidColPtrOrder);
}
}
let nnz = values.len();
if col_ptrs[ncols] != nnz {
return Err(CscError::InvalidColPtrs {
expected: nnz,
actual: col_ptrs[ncols],
});
}
for &row in &row_indices {
if row >= nrows {
return Err(CscError::InvalidRowIndex { index: row, nrows });
}
}
Ok(Self {
nrows,
ncols,
col_ptrs,
row_indices,
values,
})
}
#[inline]
pub unsafe fn new_unchecked(
nrows: usize,
ncols: usize,
col_ptrs: Vec<usize>,
row_indices: Vec<usize>,
values: Vec<T>,
) -> Self {
Self {
nrows,
ncols,
col_ptrs,
row_indices,
values,
}
}
pub fn zeros(nrows: usize, ncols: usize) -> Self {
Self {
nrows,
ncols,
col_ptrs: vec![0; ncols + 1],
row_indices: Vec::new(),
values: Vec::new(),
}
}
pub fn eye(n: usize) -> Self
where
T: Field,
{
let mut col_ptrs = Vec::with_capacity(n + 1);
let mut row_indices = Vec::with_capacity(n);
let mut values = Vec::with_capacity(n);
for i in 0..n {
col_ptrs.push(i);
row_indices.push(i);
values.push(T::one());
}
col_ptrs.push(n);
Self {
nrows: n,
ncols: n,
col_ptrs,
row_indices,
values,
}
}
#[inline]
pub fn nrows(&self) -> usize {
self.nrows
}
#[inline]
pub fn ncols(&self) -> usize {
self.ncols
}
#[inline]
pub fn shape(&self) -> (usize, usize) {
(self.nrows, self.ncols)
}
#[inline]
pub fn nnz(&self) -> usize {
self.values.len()
}
#[inline]
pub fn density(&self) -> f64 {
if self.nrows == 0 || self.ncols == 0 {
0.0
} else {
self.nnz() as f64 / (self.nrows * self.ncols) as f64
}
}
#[inline]
pub fn col_ptrs(&self) -> &[usize] {
&self.col_ptrs
}
#[inline]
pub fn row_indices(&self) -> &[usize] {
&self.row_indices
}
#[inline]
pub fn values(&self) -> &[T] {
&self.values
}
#[inline]
pub fn values_mut(&mut self) -> &mut [T] {
&mut self.values
}
pub fn get(&self, row: usize, col: usize) -> Option<&T> {
if row >= self.nrows || col >= self.ncols {
return None;
}
let start = self.col_ptrs[col];
let end = self.col_ptrs[col + 1];
for i in start..end {
if self.row_indices[i] == row {
return Some(&self.values[i]);
}
}
None
}
pub fn get_or_zero(&self, row: usize, col: usize) -> T
where
T: Field,
{
self.get(row, col).cloned().unwrap_or_else(T::zero)
}
pub fn col_iter(&self, col: usize) -> impl Iterator<Item = (usize, &T)> {
let start = self.col_ptrs[col];
let end = self.col_ptrs[col + 1];
self.row_indices[start..end]
.iter()
.zip(self.values[start..end].iter())
.map(|(&row, val)| (row, val))
}
pub fn iter(&self) -> impl Iterator<Item = (usize, usize, &T)> {
(0..self.ncols).flat_map(move |col| {
let start = self.col_ptrs[col];
let end = self.col_ptrs[col + 1];
self.row_indices[start..end]
.iter()
.zip(self.values[start..end].iter())
.map(move |(&row, val)| (row, col, val))
})
}
pub fn to_csr(&self) -> crate::csr::CsrMatrix<T> {
crate::convert::csc_to_csr(self)
}
pub fn to_dense(&self) -> oxiblas_matrix::Mat<T>
where
T: Field + bytemuck::Zeroable,
{
let mut dense = oxiblas_matrix::Mat::zeros(self.nrows, self.ncols);
for col in 0..self.ncols {
let start = self.col_ptrs[col];
let end = self.col_ptrs[col + 1];
for i in start..end {
dense[(self.row_indices[i], col)] = self.values[i].clone();
}
}
dense
}
pub fn from_dense(dense: &oxiblas_matrix::MatRef<'_, T>) -> Self
where
T: Field,
{
let (nrows, ncols) = dense.shape();
let mut col_ptrs = Vec::with_capacity(ncols + 1);
let mut row_indices = Vec::new();
let mut values = Vec::new();
let eps = <T as Scalar>::epsilon();
col_ptrs.push(0);
for j in 0..ncols {
for i in 0..nrows {
let val = dense[(i, j)].clone();
if Scalar::abs(val.clone()) > eps {
row_indices.push(i);
values.push(val);
}
}
col_ptrs.push(values.len());
}
Self {
nrows,
ncols,
col_ptrs,
row_indices,
values,
}
}
pub fn transpose(&self) -> Self {
let csr = self.to_csr();
csr.to_csc()
}
pub fn scale(&mut self, alpha: T) {
for val in &mut self.values {
*val = val.clone() * alpha.clone();
}
}
pub fn scaled(&self, alpha: T) -> Self {
let mut result = self.clone();
result.scale(alpha);
result
}
#[inline]
pub fn col_nnz(&self, col: usize) -> usize {
self.col_ptrs[col + 1] - self.col_ptrs[col]
}
pub fn is_structurally_symmetric(&self) -> bool {
if self.nrows != self.ncols {
return false;
}
for col in 0..self.ncols {
let start = self.col_ptrs[col];
let end = self.col_ptrs[col + 1];
for i in start..end {
let row = self.row_indices[i];
if self.get(col, row).is_none() {
return false;
}
}
}
true
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_csc_new() {
let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
let row_indices = vec![0, 2, 1, 0, 2];
let col_ptrs = vec![0, 2, 3, 5];
let csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values).unwrap();
assert_eq!(csc.nrows(), 3);
assert_eq!(csc.ncols(), 3);
assert_eq!(csc.nnz(), 5);
}
#[test]
fn test_csc_get() {
let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
let row_indices = vec![0, 2, 1, 0, 2];
let col_ptrs = vec![0, 2, 3, 5];
let csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values).unwrap();
assert_eq!(csc.get(0, 0), Some(&1.0));
assert_eq!(csc.get(2, 0), Some(&2.0));
assert_eq!(csc.get(1, 1), Some(&3.0));
assert_eq!(csc.get(0, 2), Some(&4.0));
assert_eq!(csc.get(2, 2), Some(&5.0));
assert_eq!(csc.get(1, 0), None);
assert_eq!(csc.get(0, 1), None);
}
#[test]
fn test_csc_zeros() {
let csc: CscMatrix<f64> = CscMatrix::zeros(5, 3);
assert_eq!(csc.nrows(), 5);
assert_eq!(csc.ncols(), 3);
assert_eq!(csc.nnz(), 0);
}
#[test]
fn test_csc_eye() {
let csc: CscMatrix<f64> = CscMatrix::eye(4);
assert_eq!(csc.nrows(), 4);
assert_eq!(csc.ncols(), 4);
assert_eq!(csc.nnz(), 4);
for i in 0..4 {
assert_eq!(csc.get(i, i), Some(&1.0));
}
}
#[test]
fn test_csc_density() {
let values = vec![1.0f64, 2.0, 3.0];
let row_indices = vec![0, 1, 2];
let col_ptrs = vec![0, 1, 2, 3];
let csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values).unwrap();
let density = csc.density();
assert!((density - 3.0 / 9.0).abs() < 1e-10);
}
#[test]
fn test_csc_col_iter() {
let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
let row_indices = vec![0, 2, 1, 0, 2];
let col_ptrs = vec![0, 2, 3, 5];
let csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values).unwrap();
let col0: Vec<_> = csc.col_iter(0).collect();
assert_eq!(col0, vec![(0, &1.0), (2, &2.0)]);
let col1: Vec<_> = csc.col_iter(1).collect();
assert_eq!(col1, vec![(1, &3.0)]);
}
#[test]
fn test_csc_scale() {
let values = vec![1.0f64, 2.0, 3.0];
let row_indices = vec![0, 1, 2];
let col_ptrs = vec![0, 1, 2, 3];
let mut csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values).unwrap();
csc.scale(2.0);
assert_eq!(csc.values(), &[2.0, 4.0, 6.0]);
}
#[test]
fn test_csc_invalid_col_ptrs() {
let values = vec![1.0f64, 2.0];
let row_indices = vec![0, 1];
let col_ptrs = vec![0, 1];
let result = CscMatrix::new(2, 2, col_ptrs, row_indices, values);
assert!(matches!(result, Err(CscError::InvalidColPtrs { .. })));
}
#[test]
fn test_csc_invalid_row_index() {
let values = vec![1.0f64];
let row_indices = vec![5]; let col_ptrs = vec![0, 1];
let result = CscMatrix::new(3, 1, col_ptrs, row_indices, values);
assert!(matches!(result, Err(CscError::InvalidRowIndex { .. })));
}
#[test]
fn test_csc_col_nnz() {
let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
let row_indices = vec![0, 2, 1, 0, 2];
let col_ptrs = vec![0, 2, 3, 5];
let csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values).unwrap();
assert_eq!(csc.col_nnz(0), 2);
assert_eq!(csc.col_nnz(1), 1);
assert_eq!(csc.col_nnz(2), 2);
}
#[test]
fn test_csc_structurally_symmetric() {
let values = vec![1.0f64, 2.0, 2.0, 3.0];
let row_indices = vec![0, 1, 0, 1];
let col_ptrs = vec![0, 2, 4];
let csc = CscMatrix::new(2, 2, col_ptrs, row_indices, values).unwrap();
assert!(csc.is_structurally_symmetric());
}
#[test]
fn test_csc_structurally_symmetric_block_diagonal_non_dense() {
let values = vec![1.0f64, 3.0, 2.0, 4.0, 5.0];
let row_indices = vec![0, 1, 0, 1, 2];
let col_ptrs = vec![0, 2, 4, 5];
let csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values).unwrap();
assert!(csc.is_structurally_symmetric());
}
#[test]
fn test_csc_not_structurally_symmetric_upper_triangular() {
let values = vec![1.0f64, 2.0, 3.0, 4.0];
let row_indices = vec![0, 0, 1, 2];
let col_ptrs = vec![0, 1, 3, 4];
let csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values).unwrap();
assert!(!csc.is_structurally_symmetric());
}
#[test]
fn test_csc_no_index_operator_use_get_instead() {
let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
let row_indices = vec![0, 2, 1, 0, 2];
let col_ptrs = vec![0, 2, 3, 5];
let csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values).expect("valid CSC");
assert_eq!(csc.get(0, 1), None);
assert_eq!(csc.get_or_zero(0, 1), 0.0);
assert_eq!(csc.get(99, 99), None);
assert_eq!(csc.get_or_zero(99, 99), 0.0);
assert_eq!(csc.get_or_zero(0, 0), 1.0);
assert_eq!(csc.get_or_zero(2, 2), 5.0);
}
}