1use std::error::Error;
4use std::fmt::{Display, Formatter};
5
6pub const SPARSE_KERNEL_SELECTOR_SCHEMA_VERSION: u32 = 1;
8
9#[derive(Debug, Clone, Copy, Eq, PartialEq)]
11pub enum SparseKernelWorkloadClass {
12 SpmvLike,
14 SpmmLike,
16 MaskedUpdate,
18 FrontierExpansion,
20}
21
22impl SparseKernelWorkloadClass {
23 #[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#[derive(Debug, Clone, Copy, Eq, PartialEq)]
37pub enum SparseKernelSelectedPath {
38 CusparseSpmv,
40 CusparseSpmm,
42 CusparseMaskedUpdate,
44 FrontierExpansion,
46}
47
48impl SparseKernelSelectedPath {
49 #[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#[derive(Debug, Clone, Eq, PartialEq)]
63pub struct SparseKernelSelectorRequest {
64 pub workload_class: SparseKernelWorkloadClass,
66 pub rows: u32,
68 pub cols: u32,
70 pub nnz: u32,
72 pub rhs_cols: u32,
74 pub mask_nnz: u32,
76 pub frontier_nnz: u32,
78 pub baseline_id: String,
80 pub result_digest: [u8; 32],
82 pub transfer_bytes: u64,
84}
85
86#[derive(Debug, Clone, Eq, PartialEq)]
88pub struct SparseKernelSelectorEvidence {
89 pub schema_version: u32,
91 pub workload_class: &'static str,
93 pub selected_path: &'static str,
95 pub baseline_id: String,
97 pub rows: u32,
99 pub cols: u32,
101 pub nnz: u32,
103 pub rhs_cols: u32,
105 pub mask_nnz: u32,
107 pub frontier_nnz: u32,
109 pub density_bps: u32,
111 pub result_digest: [u8; 32],
113 pub transfer_bytes: u64,
115}
116
117#[derive(Debug, Clone, Eq, PartialEq)]
119pub enum SparseKernelSelectorError {
120 InvalidShape,
122 InvalidNnz,
124 InvalidRhsColumns {
126 reason: &'static str,
128 },
129 MissingMask,
131 MissingFrontier,
133 MissingBaselineId,
135 MissingResultDigest,
137 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
181pub 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}