1use std::error::Error;
4use std::fmt::{Display, Formatter};
5
6
7pub mod conversion;
8pub mod host;
9mod formats;
10#[cfg(feature = "serde")]
11mod serde;
12mod symbolic;
13pub use formats::{
14 BlockDirection, BsrMatrix, CooMatrix, CooMatrixOwned, CscMatrix, CscMatrixOwned, EllMatrix,
15};
16
17
18
19
20#[cfg(feature = "tensor")]
21pub mod tensor;
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
24pub enum IndexBase {
25 Zero,
26 One,
27}
28
29impl IndexBase {
30 const fn value(self) -> u32 {
31 match self {
32 Self::Zero => 0,
33 Self::One => 1,
34 }
35 }
36}
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
39pub enum Operation {
40 None,
41 Transpose,
42 ConjugateTranspose,
43}
44
45#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
46pub enum DenseOrder {
47 RowMajor,
48 ColumnMajor,
49}
50
51#[derive(Debug, Clone, Copy)]
52pub struct DenseMatrix<'a> {
53 values: &'a [f32],
54 rows: usize,
55 columns: usize,
56 order: DenseOrder,
57}
58
59#[derive(Debug, Clone, PartialEq)]
60pub struct DenseMatrixOwned {
61 values: Vec<f32>,
62 rows: usize,
63 columns: usize,
64 order: DenseOrder,
65}
66
67impl DenseMatrixOwned {
68 pub fn new(
69 values: Vec<f32>,
70 rows: usize,
71 columns: usize,
72 order: DenseOrder,
73 ) -> Result<Self, SparseError> {
74 DenseMatrix::new(&values, rows, columns, order)?;
75 Ok(Self {
76 values,
77 rows,
78 columns,
79 order,
80 })
81 }
82
83 pub fn as_ref(&self) -> DenseMatrix<'_> {
84 DenseMatrix {
85 values: &self.values,
86 rows: self.rows,
87 columns: self.columns,
88 order: self.order,
89 }
90 }
91}
92
93#[derive(Debug, Clone, Copy)]
94pub struct SparseVector<'a> {
95 indices: &'a [u32],
96 values: &'a [f32],
97 size: usize,
98 index_base: IndexBase,
99}
100
101#[derive(Debug, Clone, PartialEq)]
102pub struct SparseVectorOwned {
103 indices: Vec<u32>,
104 values: Vec<f32>,
105 size: usize,
106 index_base: IndexBase,
107}
108
109impl SparseVectorOwned {
110 pub fn new(
111 size: usize,
112 indices: Vec<u32>,
113 values: Vec<f32>,
114 index_base: IndexBase,
115 ) -> Result<Self, SparseError> {
116 SparseVector::new(size, &indices, &values, index_base)?;
117 Ok(Self { indices, values, size, index_base })
118 }
119
120 pub fn as_ref(&self) -> SparseVector<'_> {
121 SparseVector {
122 indices: &self.indices,
123 values: &self.values,
124 size: self.size,
125 index_base: self.index_base,
126 }
127 }
128}
129
130impl<'a> SparseVector<'a> {
131 pub fn new(
132 size: usize,
133 indices: &'a [u32],
134 values: &'a [f32],
135 index_base: IndexBase,
136 ) -> Result<Self, SparseError> {
137 dimension(size, "sparse vector size")?;
138 if indices.len() != values.len() {
139 return Err(SparseError::BufferLength {
140 name: "sparse vector indices",
141 expected: values.len(),
142 actual: indices.len(),
143 });
144 }
145 let base = index_base.value();
146 let limit = base
147 .checked_add(dimension(size, "sparse vector index range")?)
148 .ok_or(SparseError::SizeOverflow("sparse vector index range"))?;
149 if indices.iter().any(|&index| index < base || index >= limit) {
150 return Err(SparseError::InvalidSparseIndex("sparse vector"));
151 }
152 Ok(Self {
153 indices,
154 values,
155 size,
156 index_base,
157 })
158 }
159
160 pub const fn size(&self) -> usize {
161 self.size
162 }
163
164 pub const fn nnz(&self) -> usize {
165 self.values.len()
166 }
167
168 pub const fn indices(&self) -> &'a [u32] {
169 self.indices
170 }
171
172 pub const fn values(&self) -> &'a [f32] {
173 self.values
174 }
175
176 pub const fn index_base(&self) -> IndexBase {
177 self.index_base
178 }
179}
180
181impl<'a> DenseMatrix<'a> {
182 pub fn new(
183 values: &'a [f32],
184 rows: usize,
185 columns: usize,
186 order: DenseOrder,
187 ) -> Result<Self, SparseError> {
188 let expected = rows
189 .checked_mul(columns)
190 .ok_or(SparseError::SizeOverflow("dense matrix"))?;
191 if values.len() != expected {
192 return Err(SparseError::BufferLength {
193 name: "dense matrix",
194 expected,
195 actual: values.len(),
196 });
197 }
198 Ok(Self {
199 values,
200 rows,
201 columns,
202 order,
203 })
204 }
205
206 pub const fn values(&self) -> &'a [f32] {
207 self.values
208 }
209
210 pub const fn rows(&self) -> usize {
211 self.rows
212 }
213
214 pub const fn columns(&self) -> usize {
215 self.columns
216 }
217
218 pub const fn order(&self) -> DenseOrder {
219 self.order
220 }
221
222 fn physical_strides(self) -> Result<(u32, u32), SparseError> {
223 let rows = dimension(self.rows, "dense rows")?;
224 let columns = dimension(self.columns, "dense columns")?;
225 Ok(match self.order {
226 DenseOrder::RowMajor => (columns, 1),
227 DenseOrder::ColumnMajor => (1, rows),
228 })
229 }
230
231 fn operation_shape(self, operation: Operation) -> (usize, usize) {
232 match operation {
233 Operation::None => (self.rows, self.columns),
234 Operation::Transpose | Operation::ConjugateTranspose => (self.columns, self.rows),
235 }
236 }
237
238 fn operation_strides(self, operation: Operation) -> Result<(u32, u32), SparseError> {
239 let (row, column) = self.physical_strides()?;
240 Ok(match operation {
241 Operation::None => (row, column),
242 Operation::Transpose | Operation::ConjugateTranspose => (column, row),
243 })
244 }
245}
246
247#[derive(Debug, Clone, Copy)]
248pub struct CsrMatrix<'a> {
249 values: &'a [f32],
250 row_offsets: &'a [u32],
251 column_indices: &'a [u32],
252 rows: usize,
253 columns: usize,
254 index_base: IndexBase,
255}
256
257impl<'a> CsrMatrix<'a> {
258 pub fn new(
259 rows: usize,
260 columns: usize,
261 row_offsets: &'a [u32],
262 column_indices: &'a [u32],
263 values: &'a [f32],
264 index_base: IndexBase,
265 ) -> Result<Self, SparseError> {
266 let matrix = Self {
267 values,
268 row_offsets,
269 column_indices,
270 rows,
271 columns,
272 index_base,
273 };
274 matrix.validate()?;
275 Ok(matrix)
276 }
277
278 pub const fn rows(&self) -> usize {
279 self.rows
280 }
281
282 pub const fn columns(&self) -> usize {
283 self.columns
284 }
285
286 pub const fn nnz(&self) -> usize {
287 self.values.len()
288 }
289
290 pub const fn index_base(&self) -> IndexBase {
291 self.index_base
292 }
293
294 pub const fn row_offsets(&self) -> &'a [u32] {
295 self.row_offsets
296 }
297
298 pub const fn column_indices(&self) -> &'a [u32] {
299 self.column_indices
300 }
301
302 pub const fn values(&self) -> &'a [f32] {
303 self.values
304 }
305
306 fn validate(self) -> Result<(), SparseError> {
307 dimension(self.rows, "CSR rows")?;
308 dimension(self.columns, "CSR columns")?;
309 let expected_offsets = self
310 .rows
311 .checked_add(1)
312 .ok_or(SparseError::SizeOverflow("CSR row offsets"))?;
313 if self.row_offsets.len() != expected_offsets {
314 return Err(SparseError::BufferLength {
315 name: "CSR row offsets",
316 expected: expected_offsets,
317 actual: self.row_offsets.len(),
318 });
319 }
320 if self.column_indices.len() != self.values.len() {
321 return Err(SparseError::BufferLength {
322 name: "CSR column indices",
323 expected: self.values.len(),
324 actual: self.column_indices.len(),
325 });
326 }
327 let base = self.index_base.value();
328 if self.row_offsets.first().copied() != Some(base) {
329 return Err(SparseError::InvalidRowOffsets(
330 "first row offset does not equal the index base",
331 ));
332 }
333 for offsets in self.row_offsets.windows(2) {
334 if offsets[0] > offsets[1] {
335 return Err(SparseError::InvalidRowOffsets(
336 "row offsets are not nondecreasing",
337 ));
338 }
339 }
340 let terminal = self.row_offsets.last().copied().unwrap_or(base);
341 let encoded_nnz = terminal
342 .checked_sub(base)
343 .ok_or(SparseError::InvalidRowOffsets(
344 "terminal row offset is below the index base",
345 ))?;
346 if encoded_nnz as usize != self.values.len() {
347 return Err(SparseError::InvalidRowOffsets(
348 "terminal row offset does not match nnz",
349 ));
350 }
351 let column_limit = base
352 .checked_add(dimension(self.columns, "CSR columns")?)
353 .ok_or(SparseError::SizeOverflow("CSR column index range"))?;
354 if self
355 .column_indices
356 .iter()
357 .any(|&column| column < base || column >= column_limit)
358 {
359 return Err(SparseError::InvalidColumnIndex);
360 }
361 Ok(())
362 }
363
364 pub fn transpose(self) -> Result<CsrMatrixOwned, SparseError> {
365 self.transpose_impl(false).map(|(matrix, _)| matrix)
366 }
367
368 pub fn transpose_with_permutation(self) -> Result<(CsrMatrixOwned, Vec<u32>), SparseError> {
369 self.transpose_impl(true)
370 }
371
372 fn transpose_impl(self, capture_permutation: bool) -> Result<(CsrMatrixOwned, Vec<u32>), SparseError> {
373 let base = self.index_base.value();
374 let mut row_offsets = vec![0_u32; self.columns + 1];
375 for &encoded_column in self.column_indices {
376 let column = (encoded_column - base) as usize;
377 row_offsets[column + 1] = row_offsets[column + 1]
378 .checked_add(1)
379 .ok_or(SparseError::SizeOverflow("transposed CSR row counts"))?;
380 }
381 for row in 0..self.columns {
382 row_offsets[row + 1] = row_offsets[row + 1]
383 .checked_add(row_offsets[row])
384 .ok_or(SparseError::SizeOverflow("transposed CSR row offsets"))?;
385 }
386 let mut positions = row_offsets[..self.columns].to_vec();
387 let mut column_indices = vec![0_u32; self.nnz()];
388 let mut values = vec![0.0_f32; self.nnz()];
389 let mut permutation = if capture_permutation { vec![0u32; self.nnz()] } else { Vec::new() };
390 for row in 0..self.rows {
391 let start = (self.row_offsets[row] - base) as usize;
392 let end = (self.row_offsets[row + 1] - base) as usize;
393 for entry in start..end {
394 let column = (self.column_indices[entry] - base) as usize;
395 let destination = positions[column] as usize;
396 column_indices[destination] = dimension(row, "transposed CSR column")? + base;
397 values[destination] = self.values[entry];
398 if capture_permutation {
399 permutation[destination] = dimension(entry, "transposed CSR permutation")?;
400 }
401 positions[column] += 1;
402 }
403 }
404 for offset in &mut row_offsets {
405 *offset = offset
406 .checked_add(base)
407 .ok_or(SparseError::SizeOverflow("transposed CSR index base"))?;
408 }
409 let matrix = CsrMatrixOwned::new(
410 self.columns,
411 self.rows,
412 row_offsets,
413 column_indices,
414 values,
415 self.index_base,
416 )?;
417 Ok((matrix, permutation))
418 }
419}
420
421#[derive(Debug, Clone, PartialEq)]
422pub struct CsrMatrixOwned {
423 values: Vec<f32>,
424 row_offsets: Vec<u32>,
425 column_indices: Vec<u32>,
426 rows: usize,
427 columns: usize,
428 index_base: IndexBase,
429}
430
431impl CsrMatrixOwned {
432 pub fn new(
433 rows: usize,
434 columns: usize,
435 row_offsets: Vec<u32>,
436 column_indices: Vec<u32>,
437 values: Vec<f32>,
438 index_base: IndexBase,
439 ) -> Result<Self, SparseError> {
440 CsrMatrix::new(
441 rows,
442 columns,
443 &row_offsets,
444 &column_indices,
445 &values,
446 index_base,
447 )?;
448 Ok(Self {
449 values,
450 row_offsets,
451 column_indices,
452 rows,
453 columns,
454 index_base,
455 })
456 }
457
458 pub fn as_ref(&self) -> CsrMatrix<'_> {
459 CsrMatrix {
460 values: &self.values,
461 row_offsets: &self.row_offsets,
462 column_indices: &self.column_indices,
463 rows: self.rows,
464 columns: self.columns,
465 index_base: self.index_base,
466 }
467 }
468}
469
470#[derive(Debug, Clone, Copy, PartialEq, Eq)]
471pub enum SparseAlgorithm {
472 RowSplitWavefront32,
473}
474
475#[derive(Debug, Clone, Copy, PartialEq, Eq)]
476pub struct SparsePlan {
477 pub algorithm: SparseAlgorithm,
478 pub output_elements: u32,
479 pub block_threads: u32,
480 pub grid_blocks: u32,
481}
482
483#[derive(Debug)]
484pub enum SparseError {
485 BufferLength {
486 name: &'static str,
487 expected: usize,
488 actual: usize,
489 },
490 DimensionMismatch(&'static str),
491 DimensionTooLarge(&'static str),
492 InvalidColumnIndex,
493 InvalidSparseIndex(&'static str),
494 InvalidSparseOffsets {
495 format: &'static str,
496 message: &'static str,
497 },
498 InvalidBlockDimension,
499 InvalidRowOffsets(&'static str),
500 SizeOverflow(&'static str),
501
502 #[cfg(feature = "tensor")]
503 TensorExecution(ruda_core::tensor::execution::ExecutionError),
504 #[cfg(feature = "tensor")]
505 TensorData(ruda_core::tensor::data::DataError),
506 Device(&'static str),
507}
508
509impl Display for SparseError {
510 fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
511 match self {
512 Self::BufferLength {
513 name,
514 expected,
515 actual,
516 } => write!(formatter, "{name} needs {expected} elements, got {actual}"),
517 Self::DimensionMismatch(message) => formatter.write_str(message),
518 Self::DimensionTooLarge(name) => write!(formatter, "{name} does not fit u32"),
519 Self::InvalidColumnIndex => formatter.write_str("CSR column index is out of bounds"),
520 Self::InvalidSparseIndex(name) => write!(formatter, "{name} index is out of bounds"),
521 Self::InvalidSparseOffsets { format, message } => {
522 write!(formatter, "invalid {format} offsets: {message}")
523 }
524 Self::InvalidBlockDimension => {
525 formatter.write_str("BSR block dimension must be greater than zero")
526 }
527 Self::InvalidRowOffsets(message) => {
528 write!(formatter, "invalid CSR row offsets: {message}")
529 }
530 Self::SizeOverflow(name) => write!(formatter, "{name} size overflows"),
531
532 #[cfg(feature = "tensor")]
533 Self::TensorExecution(error) => Display::fmt(error, formatter),
534 #[cfg(feature = "tensor")]
535 Self::TensorData(error) => Display::fmt(error, formatter),
536 Self::Device(message) => formatter.write_str(message),
537 }
538 }
539}
540
541impl Error for SparseError {
542 fn source(&self) -> Option<&(dyn Error + 'static)> {
543 match self {
544
545 #[cfg(feature = "tensor")]
546 Self::TensorExecution(error) => Some(error),
547 #[cfg(feature = "tensor")]
548 Self::TensorData(error) => Some(error),
549 _ => None,
550 }
551 }
552}
553
554
555
556fn dimension(value: usize, name: &'static str) -> Result<u32, SparseError> {
557 u32::try_from(value).map_err(|_| SparseError::DimensionTooLarge(name))
558}
559
560fn dense_strides(
561 rows: usize,
562 columns: usize,
563 order: DenseOrder,
564 name: &'static str,
565) -> Result<(u32, u32), SparseError> {
566 let rows = dimension(rows, name)?;
567 let columns = dimension(columns, name)?;
568 Ok(match order {
569 DenseOrder::RowMajor => (columns, 1),
570 DenseOrder::ColumnMajor => (1, rows),
571 })
572}