1use ztensor::{DType, Error, Result, Rule, Tensor};
11
12#[derive(Debug, Clone)]
14pub struct Csr {
15 pub rows: u64,
16 pub cols: u64,
17 pub values: Vec<u8>,
19 pub dtype: DType,
20 pub logical: Option<String>,
21 pub indices: Vec<u64>,
23 pub indptr: Vec<u64>,
25}
26
27pub fn read(tensor: &Tensor<'_>) -> Result<Csr> {
31 if tensor.layout() != "zt.sparse_csr/1" {
32 return Err(Error::Unsupported(format!(
33 "{:?} has layout {:?}, not zt.sparse_csr/1",
34 tensor.name(),
35 tensor.layout()
36 )));
37 }
38 let [rows, cols] = tensor.shape()[..] else {
39 return Err(Error::reject(
40 Rule::LayoutRule,
41 format!("{:?}: sparse_csr requires rank-2 shape", tensor.name()),
42 ));
43 };
44
45 let idx_part = tensor.part("indices")?;
46 let idx_dtype = idx_part.dtype();
47 let values_part = tensor.part("values")?;
48 let (dtype, logical) = (
49 values_part.dtype(),
50 values_part.logical().map(str::to_string),
51 );
52
53 let indices = widen(&idx_part.bytes()?, idx_dtype);
54 let indptr = widen(&tensor.part("indptr")?.bytes()?, idx_dtype);
55 let values = values_part.bytes()?.into_owned();
56 let nnz = indices.len() as u64;
57
58 let name = tensor.name();
59 let bad = |detail: String| Err(Error::reject(Rule::LayoutData, detail));
60 if indptr.first() != Some(&0) {
61 return bad(format!("{name:?}: indptr must start at 0"));
62 }
63 if indptr.windows(2).any(|w| w[0] > w[1]) {
64 return bad(format!("{name:?}: indptr must be non-decreasing"));
65 }
66 if indptr.last() != Some(&nnz) {
67 return bad(format!("{name:?}: indptr must end at nnz ({nnz})"));
68 }
69 for r in 0..rows as usize {
70 let row = &indices[indptr[r] as usize..indptr[r + 1] as usize];
71 if row.windows(2).any(|w| w[0] >= w[1]) {
72 return bad(format!("{name:?}: row {r} indices not strictly increasing"));
73 }
74 if row.last().is_some_and(|&c| c >= cols) {
75 return bad(format!("{name:?}: row {r} has an index >= cols ({cols})"));
76 }
77 }
78
79 Ok(Csr {
80 rows,
81 cols,
82 values,
83 dtype,
84 logical,
85 indices,
86 indptr,
87 })
88}
89
90fn widen(bytes: &[u8], dtype: DType) -> Vec<u64> {
91 match dtype {
92 DType::U32 => bytes
93 .chunks_exact(4)
94 .map(|c| u32::from_le_bytes(c.try_into().unwrap()) as u64)
95 .collect(),
96 _ => bytes
97 .chunks_exact(8)
98 .map(|c| u64::from_le_bytes(c.try_into().unwrap()))
99 .collect(),
100 }
101}