Skip to main content

sketch_spgemm/
error.rs

1use std::error::Error;
2use std::fmt;
3
4/// Identifies one input of a matrix product.
5#[derive(Clone, Copy, Debug, PartialEq, Eq)]
6pub enum MatrixOperand {
7    /// Left-hand matrix.
8    Left,
9    /// Right-hand matrix.
10    Right,
11}
12
13/// Arithmetic operation that overflowed in a checked kernel.
14#[derive(Clone, Copy, Debug, PartialEq, Eq)]
15pub enum ArithmeticOperation {
16    /// Multiplication of two matrix entries.
17    Multiply,
18    /// Addition into an output accumulator.
19    Add,
20}
21
22impl fmt::Display for ArithmeticOperation {
23    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
24        match self {
25            Self::Multiply => f.write_str("multiplication"),
26            Self::Add => f.write_str("addition"),
27        }
28    }
29}
30
31impl fmt::Display for MatrixOperand {
32    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
33        match self {
34            Self::Left => f.write_str("left"),
35            Self::Right => f.write_str("right"),
36        }
37    }
38}
39
40/// Errors reported by fallible multiplication and interoperability APIs.
41#[derive(Clone, Debug, PartialEq, Eq)]
42#[non_exhaustive]
43pub enum SpGemmError {
44    /// The inner dimensions of the two matrices do not agree.
45    DimensionMismatch {
46        /// Shape of the left-hand matrix.
47        left: (usize, usize),
48        /// Shape of the right-hand matrix.
49        right: (usize, usize),
50    },
51    /// A zero-copy `sprs` adapter received column-compressed storage.
52    NonCsrStorage {
53        /// Operand using unsupported storage.
54        operand: MatrixOperand,
55    },
56    /// A `usize` index cannot be represented by an ecosystem index type.
57    IndexOverflow {
58        /// Value that could not be converted.
59        value: usize,
60        /// Destination index type or buffer.
61        target: &'static str,
62    },
63    /// An ecosystem-native output rejected generated CSR buffers.
64    InvalidOutputStructure(String),
65    /// Scalar arithmetic overflowed while accumulating an output entry.
66    ArithmeticOverflow {
67        /// Operation that could not be represented by the scalar type.
68        operation: ArithmeticOperation,
69        /// Output row being computed.
70        row: usize,
71        /// Inner-dimension index of the candidate product.
72        inner: usize,
73        /// Output column being computed.
74        column: usize,
75    },
76    /// A configured output-size budget would be exceeded.
77    OutputNnzLimitExceeded {
78        /// Maximum number of entries allowed in the output.
79        limit: usize,
80        /// Number of entries required after completing the current row.
81        attempted: usize,
82    },
83}
84
85impl fmt::Display for SpGemmError {
86    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
87        match self {
88            Self::DimensionMismatch { left, right } => write!(
89                f,
90                "incompatible matrix dimensions: left is {}x{}, right is {}x{}",
91                left.0, left.1, right.0, right.1
92            ),
93            Self::NonCsrStorage { operand } => {
94                write!(f, "{operand} sprs operand must use CSR storage")
95            }
96            Self::IndexOverflow { value, target } => {
97                write!(f, "index {value} cannot be represented by {target}")
98            }
99            Self::InvalidOutputStructure(reason) => {
100                write!(f, "generated CSR output is invalid: {reason}")
101            }
102            Self::ArithmeticOverflow {
103                operation,
104                row,
105                inner,
106                column,
107            } => write!(
108                f,
109                "arithmetic overflow during {operation} at output ({row}, {column}) through inner index {inner}"
110            ),
111            Self::OutputNnzLimitExceeded { limit, attempted } => write!(
112                f,
113                "sparse product requires at least {attempted} output entries, exceeding limit {limit}"
114            ),
115        }
116    }
117}
118
119impl Error for SpGemmError {}