Skip to main content

ztensor_compat/
csr.rs

1//! Assembling a `zt.sparse_csr/1` object.
2//!
3//! Outside the core crate, and outside the reader, because a layout profile is
4//! registry vocabulary: L2 is open by design, so the one layout this
5//! implementation happens to understand cannot be a method on [`Tensor`] or a
6//! module of a crate whose thesis is that layouts are not its business. It
7//! lives here, beside the foreign-format projections, and a profile added
8//! downstream lives in its own crate with exactly this shape.
9
10use ztensor::{DType, Error, Result, Rule, Tensor};
11
12/// An assembled CSR object with its data-level rules checked.
13#[derive(Debug, Clone)]
14pub struct Csr {
15    pub rows: u64,
16    pub cols: u64,
17    /// Decoded value bytes: `nnz` elements of `dtype`/`logical`.
18    pub values: Vec<u8>,
19    pub dtype: DType,
20    pub logical: Option<String>,
21    /// Column index per value, widened to u64.
22    pub indices: Vec<u64>,
23    /// Row pointers, `rows + 1` entries.
24    pub indptr: Vec<u64>,
25}
26
27/// Reads and assembles a `zt.sparse_csr/1` tensor, enforcing the profile's
28/// data-level MUSTs: `indptr[0] == 0`, non-decreasing, `indptr[rows] == nnz`,
29/// per-row strictly increasing indices, and every index `< cols`.
30pub 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}