use crate::auto::{AutoSpGemmConfig, AutoSpGemmStats};
use crate::error::{MatrixOperand, SpGemmError};
use crate::matrix::{CsrInput, Scalar};
use sprs::{CsMatI, CsMatViewI, SpIndex};
use std::slice;
#[derive(Clone, Debug)]
pub struct SprsCsrView<'a, N = Scalar, I: SpIndex = usize, Iptr: SpIndex = I> {
matrix: CsMatViewI<'a, N, I, Iptr>,
}
impl<'a, N, I, Iptr> SprsCsrView<'a, N, I, Iptr>
where
I: SpIndex,
Iptr: SpIndex,
{
pub fn try_new(
matrix: CsMatViewI<'a, N, I, Iptr>,
operand: MatrixOperand,
) -> Result<Self, SpGemmError> {
if !matrix.is_csr() {
return Err(SpGemmError::NonCsrStorage { operand });
}
Ok(Self { matrix })
}
pub fn as_inner(&self) -> &CsMatViewI<'a, N, I, Iptr> {
&self.matrix
}
}
#[derive(Clone, Debug)]
pub struct SprsRowIter<'a, N, I: SpIndex> {
columns: slice::Iter<'a, I>,
values: slice::Iter<'a, N>,
}
impl<N, I> Iterator for SprsRowIter<'_, N, I>
where
N: Copy + Default + PartialEq,
I: SpIndex,
{
type Item = (usize, N);
fn next(&mut self) -> Option<Self::Item> {
loop {
let &column = self.columns.next()?;
let &value = self.values.next()?;
if value != N::default() {
return Some((column.index(), value));
}
}
}
}
impl<N, I, Iptr> CsrInput for SprsCsrView<'_, N, I, Iptr>
where
N: Copy + Default + PartialEq,
I: SpIndex,
Iptr: SpIndex,
{
type Scalar = N;
type RowIter<'a>
= SprsRowIter<'a, N, I>
where
Self: 'a;
fn rows(&self) -> usize {
self.matrix.rows()
}
fn cols(&self) -> usize {
self.matrix.cols()
}
fn nnz(&self) -> usize {
self.matrix
.data()
.iter()
.filter(|&&value| value != N::default())
.count()
}
fn row(&self, row: usize) -> Self::RowIter<'_> {
let range = self.matrix.indptr().outer_inds_sz(row);
SprsRowIter {
columns: self.matrix.indices()[range.clone()].iter(),
values: self.matrix.data()[range].iter(),
}
}
}
pub fn auto_spgemm<I, Iptr>(
left: CsMatViewI<'_, Scalar, I, Iptr>,
right: CsMatViewI<'_, Scalar, I, Iptr>,
config: AutoSpGemmConfig,
) -> Result<(CsMatI<Scalar, I, Iptr>, AutoSpGemmStats), SpGemmError>
where
I: SpIndex,
Iptr: SpIndex,
{
let left = SprsCsrView::try_new(left, MatrixOperand::Left)?;
let right = SprsCsrView::try_new(right, MatrixOperand::Right)?;
let (product, stats) = crate::try_auto_spgemm(&left, &right, config)?;
let row_ptr = product
.row_ptr
.into_iter()
.map(|value| {
Iptr::try_from_usize(value).ok_or(SpGemmError::IndexOverflow {
value,
target: "sprs row-pointer index type",
})
})
.collect::<Result<Vec<_>, _>>()?;
let col_idx = product
.col_idx
.into_iter()
.map(|value| {
I::try_from_usize(value).ok_or(SpGemmError::IndexOverflow {
value,
target: "sprs column-index type",
})
})
.collect::<Result<Vec<_>, _>>()?;
let output = CsMatI::try_new(
(product.rows, product.cols),
row_ptr,
col_idx,
product.values,
)
.map_err(|(_, _, _, error)| SpGemmError::InvalidOutputStructure(error.to_string()))?;
Ok((output, stats))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::matrix::CsrInput;
use sprs::CsMatI;
#[test]
fn multiplies_borrowed_csr_and_preserves_index_types() {
let left = CsMatI::<i64, u16>::new((2, 3), vec![0, 2, 3], vec![0, 2, 1], vec![2, 3, 4]);
let right =
CsMatI::<i64, u16>::new((3, 2), vec![0, 1, 2, 4], vec![0, 1, 0, 1], vec![5, 6, 7, 8]);
let (product, _) = auto_spgemm(left.view(), right.view(), AutoSpGemmConfig::default())
.expect("compatible CSR inputs");
assert_eq!(product.to_dense().as_slice().unwrap(), &[31, 24, 0, 24]);
let _: &CsMatI<i64, u16> = &product;
}
#[test]
fn rejects_csc_and_filters_explicit_zeros() {
let csc = CsMatI::<i64, usize>::new_csc((2, 2), vec![0, 1, 1], vec![0], vec![1]);
let error = SprsCsrView::try_new(csc.view(), MatrixOperand::Left).unwrap_err();
assert_eq!(
error,
SpGemmError::NonCsrStorage {
operand: MatrixOperand::Left
}
);
let csr = CsMatI::<i64, usize>::new((1, 2), vec![0, 2], vec![0, 1], vec![0, 9]);
let view = SprsCsrView::try_new(csr.view(), MatrixOperand::Left).unwrap();
assert_eq!(view.nnz(), 1);
assert_eq!(view.row(0).collect::<Vec<_>>(), vec![(1, 9)]);
}
#[test]
fn generic_view_supports_exact_i32_kernel() {
let left = CsMatI::<i32, u16>::new((1, 2), vec![0, 2], vec![0, 1], vec![2, 3]);
let right = CsMatI::<i32, u16>::new((2, 1), vec![0, 1, 2], vec![0, 0], vec![5, 7]);
let left = SprsCsrView::try_new(left.view(), MatrixOperand::Left).unwrap();
let right = SprsCsrView::try_new(right.view(), MatrixOperand::Right).unwrap();
let (product, _) = crate::try_spgemm_hash(&left, &right).unwrap();
assert_eq!(product.values, vec![31_i32]);
}
}