1use crate::{CsrMatrixOwned, IndexBase, SparseError};
2
3#[derive(Debug, Clone, Copy)]
4pub struct CooMatrix<'a> {
5 rows: usize,
6 columns: usize,
7 row_indices: &'a [u32],
8 column_indices: &'a [u32],
9 values: &'a [f32],
10 index_base: IndexBase,
11}
12
13#[derive(Debug, Clone, PartialEq)]
14pub struct CooMatrixOwned {
15 rows: usize,
16 columns: usize,
17 row_indices: Vec<u32>,
18 column_indices: Vec<u32>,
19 values: Vec<f32>,
20 index_base: IndexBase,
21}
22
23impl CooMatrixOwned {
24 pub fn new(
25 rows: usize,
26 columns: usize,
27 row_indices: Vec<u32>,
28 column_indices: Vec<u32>,
29 values: Vec<f32>,
30 index_base: IndexBase,
31 ) -> Result<Self, SparseError> {
32 CooMatrix::new(
33 rows,
34 columns,
35 &row_indices,
36 &column_indices,
37 &values,
38 index_base,
39 )?;
40 Ok(Self {
41 rows,
42 columns,
43 row_indices,
44 column_indices,
45 values,
46 index_base,
47 })
48 }
49
50 pub fn as_ref(&self) -> CooMatrix<'_> {
51 CooMatrix {
52 rows: self.rows,
53 columns: self.columns,
54 row_indices: &self.row_indices,
55 column_indices: &self.column_indices,
56 values: &self.values,
57 index_base: self.index_base,
58 }
59 }
60}
61
62impl<'a> CooMatrix<'a> {
63 pub fn new(
64 rows: usize,
65 columns: usize,
66 row_indices: &'a [u32],
67 column_indices: &'a [u32],
68 values: &'a [f32],
69 index_base: IndexBase,
70 ) -> Result<Self, SparseError> {
71 validate_dimension(rows, "COO rows")?;
72 validate_dimension(columns, "COO columns")?;
73 check_length("COO row indices", values.len(), row_indices.len())?;
74 check_length("COO column indices", values.len(), column_indices.len())?;
75 let base = index_base.value();
76 let row_limit = index_limit(base, rows, "COO row range")?;
77 let column_limit = index_limit(base, columns, "COO column range")?;
78 if row_indices
79 .iter()
80 .any(|&row| row < base || row >= row_limit)
81 {
82 return Err(SparseError::InvalidSparseIndex("COO row"));
83 }
84 if column_indices
85 .iter()
86 .any(|&column| column < base || column >= column_limit)
87 {
88 return Err(SparseError::InvalidSparseIndex("COO column"));
89 }
90 Ok(Self {
91 rows,
92 columns,
93 row_indices,
94 column_indices,
95 values,
96 index_base,
97 })
98 }
99
100 pub const fn rows(&self) -> usize {
101 self.rows
102 }
103
104 pub const fn columns(&self) -> usize {
105 self.columns
106 }
107
108 pub const fn nnz(&self) -> usize {
109 self.values.len()
110 }
111
112 pub const fn row_indices(&self) -> &'a [u32] {
113 self.row_indices
114 }
115
116 pub const fn column_indices(&self) -> &'a [u32] {
117 self.column_indices
118 }
119
120 pub const fn values(&self) -> &'a [f32] {
121 self.values
122 }
123
124 pub const fn index_base(&self) -> IndexBase {
125 self.index_base
126 }
127
128 pub fn to_csr(self) -> Result<CsrMatrixOwned, SparseError> {
131 let base = self.index_base.value();
132 let mut row_offsets = vec![0_u32; self.rows + 1];
133 for &row in self.row_indices {
134 let row = (row - base) as usize;
135 row_offsets[row] = row_offsets[row]
136 .checked_add(1)
137 .ok_or(SparseError::SizeOverflow("COO row counts"))?;
138 }
139 prefix_offsets(&mut row_offsets, base, "COO row offsets")?;
140 let mut column_indices = vec![0_u32; self.nnz()];
141 let mut values = vec![0.0_f32; self.nnz()];
142 for entry in (0..self.nnz()).rev() {
143 let row = (self.row_indices[entry] - base) as usize;
144 row_offsets[row] -= 1;
145 let destination = (row_offsets[row] - base) as usize;
146 column_indices[destination] = self.column_indices[entry];
147 values[destination] = self.values[entry];
148 }
149 CsrMatrixOwned::new(
150 self.rows,
151 self.columns,
152 row_offsets,
153 column_indices,
154 values,
155 self.index_base,
156 )
157 }
158}
159
160#[derive(Debug, Clone, Copy)]
161pub struct CscMatrix<'a> {
162 rows: usize,
163 columns: usize,
164 column_offsets: &'a [u32],
165 row_indices: &'a [u32],
166 values: &'a [f32],
167 index_base: IndexBase,
168}
169
170#[derive(Debug, Clone, PartialEq)]
171pub struct CscMatrixOwned {
172 rows: usize,
173 columns: usize,
174 column_offsets: Vec<u32>,
175 row_indices: Vec<u32>,
176 values: Vec<f32>,
177 index_base: IndexBase,
178}
179
180impl CscMatrixOwned {
181 pub fn new(
182 rows: usize,
183 columns: usize,
184 column_offsets: Vec<u32>,
185 row_indices: Vec<u32>,
186 values: Vec<f32>,
187 index_base: IndexBase,
188 ) -> Result<Self, SparseError> {
189 CscMatrix::new(
190 rows,
191 columns,
192 &column_offsets,
193 &row_indices,
194 &values,
195 index_base,
196 )?;
197 Ok(Self {
198 rows,
199 columns,
200 column_offsets,
201 row_indices,
202 values,
203 index_base,
204 })
205 }
206
207 pub fn as_ref(&self) -> CscMatrix<'_> {
208 CscMatrix {
209 rows: self.rows,
210 columns: self.columns,
211 column_offsets: &self.column_offsets,
212 row_indices: &self.row_indices,
213 values: &self.values,
214 index_base: self.index_base,
215 }
216 }
217}
218
219impl<'a> CscMatrix<'a> {
220 pub fn new(
221 rows: usize,
222 columns: usize,
223 column_offsets: &'a [u32],
224 row_indices: &'a [u32],
225 values: &'a [f32],
226 index_base: IndexBase,
227 ) -> Result<Self, SparseError> {
228 validate_dimension(rows, "CSC rows")?;
229 validate_dimension(columns, "CSC columns")?;
230 let expected_offsets = columns
231 .checked_add(1)
232 .ok_or(SparseError::SizeOverflow("CSC column offsets"))?;
233 check_length("CSC column offsets", expected_offsets, column_offsets.len())?;
234 check_length("CSC row indices", values.len(), row_indices.len())?;
235 validate_offsets("CSC column", column_offsets, values.len(), index_base)?;
236 let base = index_base.value();
237 let row_limit = index_limit(base, rows, "CSC row range")?;
238 if row_indices
239 .iter()
240 .any(|&row| row < base || row >= row_limit)
241 {
242 return Err(SparseError::InvalidSparseIndex("CSC row"));
243 }
244 Ok(Self {
245 rows,
246 columns,
247 column_offsets,
248 row_indices,
249 values,
250 index_base,
251 })
252 }
253
254 pub const fn rows(&self) -> usize {
255 self.rows
256 }
257
258 pub const fn columns(&self) -> usize {
259 self.columns
260 }
261
262 pub const fn nnz(&self) -> usize {
263 self.values.len()
264 }
265
266 pub const fn column_offsets(&self) -> &'a [u32] {
267 self.column_offsets
268 }
269
270 pub const fn row_indices(&self) -> &'a [u32] {
271 self.row_indices
272 }
273
274 pub const fn values(&self) -> &'a [f32] {
275 self.values
276 }
277
278 pub const fn index_base(&self) -> IndexBase {
279 self.index_base
280 }
281
282 pub fn to_csr(self) -> Result<CsrMatrixOwned, SparseError> {
283 let base = self.index_base.value();
284 let mut row_offsets = vec![0_u32; self.rows + 1];
285 for &row in self.row_indices {
286 let row = (row - base) as usize;
287 row_offsets[row + 1] = row_offsets[row + 1]
288 .checked_add(1)
289 .ok_or(SparseError::SizeOverflow("CSC to CSR row counts"))?;
290 }
291 prefix_offsets(&mut row_offsets, base, "CSC to CSR row offsets")?;
292 let mut positions = row_offsets[..self.rows]
293 .iter()
294 .map(|offset| offset - base)
295 .collect::<Vec<_>>();
296 let mut column_indices = vec![0_u32; self.nnz()];
297 let mut values = vec![0.0_f32; self.nnz()];
298 for column in 0..self.columns {
299 let start = (self.column_offsets[column] - base) as usize;
300 let end = (self.column_offsets[column + 1] - base) as usize;
301 for entry in start..end {
302 let row = (self.row_indices[entry] - base) as usize;
303 let destination = positions[row] as usize;
304 column_indices[destination] = validate_dimension(column, "CSC column")? + base;
305 values[destination] = self.values[entry];
306 positions[row] += 1;
307 }
308 }
309 CsrMatrixOwned::new(
310 self.rows,
311 self.columns,
312 row_offsets,
313 column_indices,
314 values,
315 self.index_base,
316 )
317 }
318}
319
320#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
321pub enum BlockDirection {
322 RowMajor,
323 ColumnMajor,
324}
325
326#[derive(Debug, Clone, Copy)]
327pub struct BsrMatrix<'a> {
328 block_rows: usize,
329 block_columns: usize,
330 block_dimension: usize,
331 row_offsets: &'a [u32],
332 column_indices: &'a [u32],
333 values: &'a [f32],
334 direction: BlockDirection,
335 index_base: IndexBase,
336}
337
338impl<'a> BsrMatrix<'a> {
339 #[allow(clippy::too_many_arguments)]
340 pub fn new(
341 block_rows: usize,
342 block_columns: usize,
343 block_dimension: usize,
344 row_offsets: &'a [u32],
345 column_indices: &'a [u32],
346 values: &'a [f32],
347 direction: BlockDirection,
348 index_base: IndexBase,
349 ) -> Result<Self, SparseError> {
350 validate_dimension(block_rows, "BSR block rows")?;
351 validate_dimension(block_columns, "BSR block columns")?;
352 if block_dimension == 0 {
353 return Err(SparseError::InvalidBlockDimension);
354 }
355 validate_dimension(block_dimension, "BSR block dimension")?;
356 let expected_offsets = block_rows
357 .checked_add(1)
358 .ok_or(SparseError::SizeOverflow("BSR row offsets"))?;
359 check_length("BSR row offsets", expected_offsets, row_offsets.len())?;
360 validate_offsets("BSR row", row_offsets, column_indices.len(), index_base)?;
361 let block_elements = block_dimension
362 .checked_mul(block_dimension)
363 .ok_or(SparseError::SizeOverflow("BSR block"))?;
364 let expected_values = column_indices
365 .len()
366 .checked_mul(block_elements)
367 .ok_or(SparseError::SizeOverflow("BSR values"))?;
368 check_length("BSR values", expected_values, values.len())?;
369 let base = index_base.value();
370 let column_limit = index_limit(base, block_columns, "BSR block-column range")?;
371 if column_indices
372 .iter()
373 .any(|&column| column < base || column >= column_limit)
374 {
375 return Err(SparseError::InvalidSparseIndex("BSR block column"));
376 }
377 block_rows
378 .checked_mul(block_dimension)
379 .ok_or(SparseError::SizeOverflow("BSR rows"))?;
380 block_columns
381 .checked_mul(block_dimension)
382 .ok_or(SparseError::SizeOverflow("BSR columns"))?;
383 Ok(Self {
384 block_rows,
385 block_columns,
386 block_dimension,
387 row_offsets,
388 column_indices,
389 values,
390 direction,
391 index_base,
392 })
393 }
394
395 pub const fn block_rows(&self) -> usize {
396 self.block_rows
397 }
398
399 pub const fn block_columns(&self) -> usize {
400 self.block_columns
401 }
402
403 pub const fn block_dimension(&self) -> usize {
404 self.block_dimension
405 }
406
407 pub const fn nnzb(&self) -> usize {
408 self.column_indices.len()
409 }
410
411 pub const fn direction(&self) -> BlockDirection {
412 self.direction
413 }
414
415 pub const fn index_base(&self) -> IndexBase {
416 self.index_base
417 }
418
419 pub fn rows(&self) -> usize {
420 self.block_rows * self.block_dimension
421 }
422
423 pub fn columns(&self) -> usize {
424 self.block_columns * self.block_dimension
425 }
426
427 pub fn to_csr(self) -> Result<CsrMatrixOwned, SparseError> {
428 let base = self.index_base.value();
429 let rows = self.rows();
430 let columns = self.columns();
431 let entries_per_block_row = self.block_dimension;
432 let nnz = self
433 .nnzb()
434 .checked_mul(self.block_dimension)
435 .and_then(|value| value.checked_mul(self.block_dimension))
436 .ok_or(SparseError::SizeOverflow("BSR to CSR nnz"))?;
437 let mut row_offsets = vec![base; rows + 1];
438 let mut column_indices = Vec::with_capacity(nnz);
439 let mut values = Vec::with_capacity(nnz);
440 for block_row in 0..self.block_rows {
441 let block_start = (self.row_offsets[block_row] - base) as usize;
442 let block_end = (self.row_offsets[block_row + 1] - base) as usize;
443 for row_in_block in 0..entries_per_block_row {
444 for block_entry in block_start..block_end {
445 let block_column = (self.column_indices[block_entry] - base) as usize;
446 for column_in_block in 0..self.block_dimension {
447 let value_offset = match self.direction {
448 BlockDirection::RowMajor => {
449 row_in_block * self.block_dimension + column_in_block
450 }
451 BlockDirection::ColumnMajor => {
452 column_in_block * self.block_dimension + row_in_block
453 }
454 };
455 column_indices.push(
456 validate_dimension(
457 block_column * self.block_dimension + column_in_block,
458 "BSR to CSR column",
459 )? + base,
460 );
461 values.push(
462 self.values[block_entry * self.block_dimension * self.block_dimension
463 + value_offset],
464 );
465 }
466 }
467 let row = block_row * self.block_dimension + row_in_block;
468 row_offsets[row + 1] = validate_dimension(column_indices.len(), "BSR to CSR nnz")?
469 .checked_add(base)
470 .ok_or(SparseError::SizeOverflow("BSR to CSR row offsets"))?;
471 }
472 }
473 CsrMatrixOwned::new(
474 rows,
475 columns,
476 row_offsets,
477 column_indices,
478 values,
479 self.index_base,
480 )
481 }
482}
483
484#[derive(Debug, Clone, Copy)]
485pub struct EllMatrix<'a> {
486 rows: usize,
487 columns: usize,
488 width: usize,
489 column_indices: &'a [u32],
490 values: &'a [f32],
491 index_base: IndexBase,
492}
493
494impl<'a> EllMatrix<'a> {
495 pub fn new(
496 rows: usize,
497 columns: usize,
498 width: usize,
499 column_indices: &'a [u32],
500 values: &'a [f32],
501 index_base: IndexBase,
502 ) -> Result<Self, SparseError> {
503 validate_dimension(rows, "ELL rows")?;
504 validate_dimension(columns, "ELL columns")?;
505 validate_dimension(width, "ELL width")?;
506 let expected = rows
507 .checked_mul(width)
508 .ok_or(SparseError::SizeOverflow("ELL storage"))?;
509 check_length("ELL column indices", expected, column_indices.len())?;
510 check_length("ELL values", expected, values.len())?;
511 Ok(Self {
512 rows,
513 columns,
514 width,
515 column_indices,
516 values,
517 index_base,
518 })
519 }
520
521 pub const fn rows(&self) -> usize {
522 self.rows
523 }
524
525 pub const fn columns(&self) -> usize {
526 self.columns
527 }
528
529 pub const fn width(&self) -> usize {
530 self.width
531 }
532
533 pub const fn column_indices(&self) -> &'a [u32] {
534 self.column_indices
535 }
536
537 pub const fn values(&self) -> &'a [f32] {
538 self.values
539 }
540
541 pub const fn index_base(&self) -> IndexBase {
542 self.index_base
543 }
544
545 pub const fn padding_index() -> u32 {
546 u32::MAX
547 }
548
549 pub fn to_csr(self) -> Result<CsrMatrixOwned, SparseError> {
552 let base = self.index_base.value();
553 let column_limit = index_limit(base, self.columns, "ELL column range")?;
554 let mut row_offsets = Vec::with_capacity(self.rows + 1);
555 let mut column_indices = Vec::new();
556 let mut values = Vec::new();
557 row_offsets.push(base);
558 for row in 0..self.rows {
559 for slot in 0..self.width {
560 let entry = slot * self.rows + row;
561 let column = self.column_indices[entry];
562 if column >= base && column < column_limit {
563 column_indices.push(column);
564 values.push(self.values[entry]);
565 }
566 }
567 row_offsets.push(
568 validate_dimension(column_indices.len(), "ELL to CSR nnz")?
569 .checked_add(base)
570 .ok_or(SparseError::SizeOverflow("ELL to CSR row offsets"))?,
571 );
572 }
573 CsrMatrixOwned::new(
574 self.rows,
575 self.columns,
576 row_offsets,
577 column_indices,
578 values,
579 self.index_base,
580 )
581 }
582}
583
584fn validate_dimension(value: usize, name: &'static str) -> Result<u32, SparseError> {
585 u32::try_from(value).map_err(|_| SparseError::DimensionTooLarge(name))
586}
587
588fn check_length(name: &'static str, expected: usize, actual: usize) -> Result<(), SparseError> {
589 if expected != actual {
590 return Err(SparseError::BufferLength {
591 name,
592 expected,
593 actual,
594 });
595 }
596 Ok(())
597}
598
599fn index_limit(base: u32, dimension: usize, name: &'static str) -> Result<u32, SparseError> {
600 base.checked_add(validate_dimension(dimension, name)?)
601 .ok_or(SparseError::SizeOverflow(name))
602}
603
604fn validate_offsets(
605 format: &'static str,
606 offsets: &[u32],
607 entries: usize,
608 index_base: IndexBase,
609) -> Result<(), SparseError> {
610 let base = index_base.value();
611 if offsets.first().copied() != Some(base) {
612 return Err(SparseError::InvalidSparseOffsets {
613 format,
614 message: "first offset does not equal the index base",
615 });
616 }
617 if offsets.windows(2).any(|pair| pair[0] > pair[1]) {
618 return Err(SparseError::InvalidSparseOffsets {
619 format,
620 message: "offsets are not nondecreasing",
621 });
622 }
623 let terminal = offsets.last().copied().unwrap_or(base);
624 if terminal.checked_sub(base).map(|value| value as usize) != Some(entries) {
625 return Err(SparseError::InvalidSparseOffsets {
626 format,
627 message: "terminal offset does not match the entry count",
628 });
629 }
630 Ok(())
631}
632
633fn prefix_offsets(offsets: &mut [u32], base: u32, name: &'static str) -> Result<(), SparseError> {
634 for index in 0..offsets.len().saturating_sub(1) {
635 offsets[index + 1] = offsets[index + 1]
636 .checked_add(offsets[index])
637 .ok_or(SparseError::SizeOverflow(name))?;
638 }
639 for offset in offsets {
640 *offset = offset
641 .checked_add(base)
642 .ok_or(SparseError::SizeOverflow(name))?;
643 }
644 Ok(())
645}