Skip to main content

vyre_primitives/math/
sparse_selector.rs

1//! Sparse-kernel selector evidence for graph, flow, and math workloads.
2
3use std::error::Error;
4use std::fmt::{Display, Formatter};
5
6/// Sparse selector evidence schema version.
7pub const SPARSE_KERNEL_SELECTOR_SCHEMA_VERSION: u32 = 1;
8
9/// Sparse workload class.
10#[derive(Debug, Clone, Copy, Eq, PartialEq)]
11pub enum SparseKernelWorkloadClass {
12    /// Sparse matrix-vector shape.
13    SpmvLike,
14    /// Sparse matrix-matrix shape.
15    SpmmLike,
16    /// Masked sparse update shape.
17    MaskedUpdate,
18    /// Frontier expansion shape.
19    FrontierExpansion,
20}
21
22impl SparseKernelWorkloadClass {
23    /// Stable evidence label.
24    #[must_use]
25    pub const fn as_str(self) -> &'static str {
26        match self {
27            Self::SpmvLike => "spmv-like",
28            Self::SpmmLike => "spmm-like",
29            Self::MaskedUpdate => "masked-update",
30            Self::FrontierExpansion => "frontier-expansion",
31        }
32    }
33}
34
35/// Selected sparse execution path.
36#[derive(Debug, Clone, Copy, Eq, PartialEq)]
37pub enum SparseKernelSelectedPath {
38    /// cuSPARSE SpMV-style library baseline.
39    CusparseSpmv,
40    /// cuSPARSE SpMM-style library baseline.
41    CusparseSpmm,
42    /// cuSPARSE masked-update-style library baseline.
43    CusparseMaskedUpdate,
44    /// Native frontier expansion path.
45    FrontierExpansion,
46}
47
48impl SparseKernelSelectedPath {
49    /// Stable evidence label.
50    #[must_use]
51    pub const fn as_str(self) -> &'static str {
52        match self {
53            Self::CusparseSpmv => "cusparse-spmv",
54            Self::CusparseSpmm => "cusparse-spmm",
55            Self::CusparseMaskedUpdate => "cusparse-masked-update",
56            Self::FrontierExpansion => "frontier-expansion",
57        }
58    }
59}
60
61/// Selector request for one sparse workload.
62#[derive(Debug, Clone, Eq, PartialEq)]
63pub struct SparseKernelSelectorRequest {
64    /// Workload class.
65    pub workload_class: SparseKernelWorkloadClass,
66    /// Matrix rows.
67    pub rows: u32,
68    /// Matrix columns.
69    pub cols: u32,
70    /// Non-zero entries.
71    pub nnz: u32,
72    /// Dense RHS columns, `1` for SpMV-like workloads.
73    pub rhs_cols: u32,
74    /// Mask non-zero entries for masked update workloads.
75    pub mask_nnz: u32,
76    /// Active frontier entries for frontier expansion workloads.
77    pub frontier_nnz: u32,
78    /// Comparator baseline id.
79    pub baseline_id: String,
80    /// Result digest supplied by the caller's benchmark/evidence producer.
81    pub result_digest: [u8; 32],
82    /// Host/device transfer bytes for this selector decision.
83    pub transfer_bytes: u64,
84}
85
86/// Selector evidence emitted for one sparse workload.
87#[derive(Debug, Clone, Eq, PartialEq)]
88pub struct SparseKernelSelectorEvidence {
89    /// Evidence schema version.
90    pub schema_version: u32,
91    /// Workload class label.
92    pub workload_class: &'static str,
93    /// Selected execution path.
94    pub selected_path: &'static str,
95    /// Comparator baseline id.
96    pub baseline_id: String,
97    /// Matrix rows.
98    pub rows: u32,
99    /// Matrix columns.
100    pub cols: u32,
101    /// Non-zero entries.
102    pub nnz: u32,
103    /// Dense RHS columns.
104    pub rhs_cols: u32,
105    /// Mask non-zero entries.
106    pub mask_nnz: u32,
107    /// Active frontier entries.
108    pub frontier_nnz: u32,
109    /// Sparse density in basis points.
110    pub density_bps: u32,
111    /// Result digest supplied by the caller.
112    pub result_digest: [u8; 32],
113    /// Host/device transfer bytes.
114    pub transfer_bytes: u64,
115}
116
117/// Sparse selector validation failure.
118#[derive(Debug, Clone, Eq, PartialEq)]
119pub enum SparseKernelSelectorError {
120    /// Matrix dimensions are zero.
121    InvalidShape,
122    /// Non-zero count is invalid for the matrix shape.
123    InvalidNnz,
124    /// RHS column count is invalid for the workload class.
125    InvalidRhsColumns {
126        /// Required diagnostic.
127        reason: &'static str,
128    },
129    /// Mask count is required but missing.
130    MissingMask,
131    /// Frontier count is required but missing.
132    MissingFrontier,
133    /// Baseline id is blank.
134    MissingBaselineId,
135    /// Result digest is all zeros.
136    MissingResultDigest,
137    /// Transfer-byte accounting is missing.
138    MissingTransferBytes,
139}
140
141impl Display for SparseKernelSelectorError {
142    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
143        match self {
144            Self::InvalidShape => write!(
145                f,
146                "sparse kernel selector received zero rows or columns. Fix: record a non-empty sparse workload shape."
147            ),
148            Self::InvalidNnz => write!(
149                f,
150                "sparse kernel selector received invalid nnz for the matrix shape. Fix: bound nnz to rows*cols and keep sparse workloads non-empty."
151            ),
152            Self::InvalidRhsColumns { reason } => {
153                write!(f, "sparse kernel selector RHS columns are invalid: {reason}.")
154            }
155            Self::MissingMask => write!(
156                f,
157                "sparse kernel selector masked-update workload has no mask tuples. Fix: record mask_nnz before selecting a masked path."
158            ),
159            Self::MissingFrontier => write!(
160                f,
161                "sparse kernel selector frontier-expansion workload has no frontier tuples. Fix: record frontier_nnz before selecting a frontier path."
162            ),
163            Self::MissingBaselineId => write!(
164                f,
165                "sparse kernel selector baseline id is blank. Fix: record the cuSPARSE/frontier comparator id."
166            ),
167            Self::MissingResultDigest => write!(
168                f,
169                "sparse kernel selector result digest is missing. Fix: attach the benchmark output digest."
170            ),
171            Self::MissingTransferBytes => write!(
172                f,
173                "sparse kernel selector transfer bytes are zero. Fix: account host/device bytes separately from logical work."
174            ),
175        }
176    }
177}
178
179impl Error for SparseKernelSelectorError {}
180
181/// Select and validate one sparse kernel evidence row.
182///
183/// # Errors
184///
185/// Returns [`SparseKernelSelectorError`] when shape, baseline, digest,
186/// transfer-byte accounting, or workload-specific fields are incomplete.
187pub fn select_sparse_kernel(
188    request: SparseKernelSelectorRequest,
189) -> Result<SparseKernelSelectorEvidence, SparseKernelSelectorError> {
190    validate_sparse_request(&request)?;
191    let selected = match request.workload_class {
192        SparseKernelWorkloadClass::SpmvLike => SparseKernelSelectedPath::CusparseSpmv,
193        SparseKernelWorkloadClass::SpmmLike => SparseKernelSelectedPath::CusparseSpmm,
194        SparseKernelWorkloadClass::MaskedUpdate => SparseKernelSelectedPath::CusparseMaskedUpdate,
195        SparseKernelWorkloadClass::FrontierExpansion => SparseKernelSelectedPath::FrontierExpansion,
196    };
197    Ok(SparseKernelSelectorEvidence {
198        schema_version: SPARSE_KERNEL_SELECTOR_SCHEMA_VERSION,
199        workload_class: request.workload_class.as_str(),
200        selected_path: selected.as_str(),
201        baseline_id: request.baseline_id,
202        rows: request.rows,
203        cols: request.cols,
204        nnz: request.nnz,
205        rhs_cols: request.rhs_cols,
206        mask_nnz: request.mask_nnz,
207        frontier_nnz: request.frontier_nnz,
208        density_bps: density_bps(request.rows, request.cols, request.nnz),
209        result_digest: request.result_digest,
210        transfer_bytes: request.transfer_bytes,
211    })
212}
213
214fn validate_sparse_request(
215    request: &SparseKernelSelectorRequest,
216) -> Result<(), SparseKernelSelectorError> {
217    if request.rows == 0 || request.cols == 0 {
218        return Err(SparseKernelSelectorError::InvalidShape);
219    }
220    let Some(cells) = request.rows.checked_mul(request.cols) else {
221        return Err(SparseKernelSelectorError::InvalidNnz);
222    };
223    if request.nnz == 0 || request.nnz > cells {
224        return Err(SparseKernelSelectorError::InvalidNnz);
225    }
226    match request.workload_class {
227        SparseKernelWorkloadClass::SpmvLike => {
228            if request.rhs_cols != 1 {
229                return Err(SparseKernelSelectorError::InvalidRhsColumns {
230                    reason: "SpMV-like workloads require rhs_cols == 1",
231                });
232            }
233        }
234        SparseKernelWorkloadClass::SpmmLike => {
235            if request.rhs_cols <= 1 {
236                return Err(SparseKernelSelectorError::InvalidRhsColumns {
237                    reason: "SpMM-like workloads require rhs_cols > 1",
238                });
239            }
240        }
241        SparseKernelWorkloadClass::MaskedUpdate => {
242            if request.mask_nnz == 0 {
243                return Err(SparseKernelSelectorError::MissingMask);
244            }
245        }
246        SparseKernelWorkloadClass::FrontierExpansion => {
247            if request.frontier_nnz == 0 {
248                return Err(SparseKernelSelectorError::MissingFrontier);
249            }
250        }
251    }
252    if request.baseline_id.trim().is_empty() {
253        return Err(SparseKernelSelectorError::MissingBaselineId);
254    }
255    if request.result_digest == [0; 32] {
256        return Err(SparseKernelSelectorError::MissingResultDigest);
257    }
258    if request.transfer_bytes == 0 {
259        return Err(SparseKernelSelectorError::MissingTransferBytes);
260    }
261    Ok(())
262}
263
264fn density_bps(rows: u32, cols: u32, nnz: u32) -> u32 {
265    let cells = u64::from(rows).saturating_mul(u64::from(cols)).max(1);
266    ((u64::from(nnz) * 10_000) / cells).min(10_000) as u32
267}
268
269#[cfg(test)]
270mod tests {
271    use super::*;
272
273    fn digest(byte: u8) -> [u8; 32] {
274        [byte; 32]
275    }
276
277    fn request(workload_class: SparseKernelWorkloadClass) -> SparseKernelSelectorRequest {
278        SparseKernelSelectorRequest {
279            workload_class,
280            rows: 128,
281            cols: 256,
282            nnz: 512,
283            rhs_cols: 1,
284            mask_nnz: 0,
285            frontier_nnz: 0,
286            baseline_id: "cusparse-12.5".to_string(),
287            result_digest: digest(7),
288            transfer_bytes: 4096,
289        }
290    }
291
292    #[test]
293    fn sparse_selector_classifies_spmv_spmm_masked_and_frontier_workloads() {
294        let spmv = select_sparse_kernel(request(SparseKernelWorkloadClass::SpmvLike)).unwrap();
295        assert_eq!(spmv.selected_path, "cusparse-spmv");
296        assert_eq!(spmv.baseline_id, "cusparse-12.5");
297        assert_eq!(spmv.result_digest, digest(7));
298        assert_eq!(spmv.transfer_bytes, 4096);
299
300        let mut spmm_request = request(SparseKernelWorkloadClass::SpmmLike);
301        spmm_request.rhs_cols = 8;
302        let spmm = select_sparse_kernel(spmm_request).unwrap();
303        assert_eq!(spmm.selected_path, "cusparse-spmm");
304
305        let mut masked_request = request(SparseKernelWorkloadClass::MaskedUpdate);
306        masked_request.mask_nnz = 64;
307        let masked = select_sparse_kernel(masked_request).unwrap();
308        assert_eq!(masked.selected_path, "cusparse-masked-update");
309        assert_eq!(masked.mask_nnz, 64);
310
311        let mut frontier_request = request(SparseKernelWorkloadClass::FrontierExpansion);
312        frontier_request.frontier_nnz = 32;
313        let frontier = select_sparse_kernel(frontier_request).unwrap();
314        assert_eq!(frontier.selected_path, "frontier-expansion");
315        assert_eq!(frontier.frontier_nnz, 32);
316    }
317
318    #[test]
319    fn sparse_selector_rejects_missing_result_digest() {
320        let mut request = request(SparseKernelWorkloadClass::SpmvLike);
321        request.result_digest = [0; 32];
322
323        let error = select_sparse_kernel(request).unwrap_err();
324
325        assert_eq!(error, SparseKernelSelectorError::MissingResultDigest);
326    }
327
328    #[test]
329    fn sparse_selector_rejects_missing_transfer_bytes() {
330        let mut request = request(SparseKernelWorkloadClass::SpmvLike);
331        request.transfer_bytes = 0;
332
333        let error = select_sparse_kernel(request).unwrap_err();
334
335        assert_eq!(error, SparseKernelSelectorError::MissingTransferBytes);
336    }
337}