1use vyre_foundation::ir::{DataType, Program};
21
22#[cfg(test)]
23use crate::bitset::bitset_words;
24#[cfg(any(test, feature = "cpu-parity"))]
25use crate::graph::csr_frontier_queue::{
26 try_csr_queue_forward_traverse_cpu, try_csr_queue_forward_traverse_cpu_into,
27};
28use crate::graph::csr_frontier_step::{
29 csr_queue_step_program, CsrQueueEmit, CsrQueueInputs, CsrQueueLanes, CsrQueueRowPlan,
30 CsrQueueStepSpec,
31};
32
33pub const CSR_QUEUE_STRIDED_FORWARD_OP_ID: &str =
35 "vyre-primitives::graph::csr_queue_strided_forward_traverse";
36
37pub const CSR_QUEUE_STRIDED_FORWARD_LANES_PER_SOURCE: u32 = 32;
39
40pub const CSR_QUEUE_STRIDED_FORWARD_WORKGROUP_SIZE: [u32; 3] = [256, 1, 1];
42
43#[must_use]
45pub const fn csr_queue_strided_forward_dispatch_grid(queue_capacity: u32) -> [u32; 3] {
46 let total_lanes = queue_capacity.saturating_mul(CSR_QUEUE_STRIDED_FORWARD_LANES_PER_SOURCE);
47 let blocks = total_lanes.div_ceil(CSR_QUEUE_STRIDED_FORWARD_WORKGROUP_SIZE[0]);
48 [if blocks == 0 { 1 } else { blocks }, 1, 1]
49}
50
51#[derive(Clone, Copy, Debug)]
53pub struct CsrQueueStridedForwardParams<'a> {
54 pub active_queue: &'a str,
56 pub queue_len: &'a str,
58 pub edge_offsets: &'a str,
60 pub edge_targets: &'a str,
62 pub edge_kind_mask: &'a str,
64 pub frontier_out: &'a str,
66 pub node_count: u32,
68 pub edge_count: u32,
70 pub queue_capacity: u32,
72 pub allow_mask: u32,
74}
75
76#[must_use]
79#[allow(clippy::too_many_arguments)]
80pub fn csr_queue_strided_forward_traverse(
81 active_queue: &str,
82 queue_len: &str,
83 edge_offsets: &str,
84 edge_targets: &str,
85 edge_kind_mask: &str,
86 frontier_out: &str,
87 node_count: u32,
88 edge_count: u32,
89 queue_capacity: u32,
90 allow_mask: u32,
91) -> Program {
92 csr_queue_strided_forward_traverse_with(CsrQueueStridedForwardParams {
93 active_queue,
94 queue_len,
95 edge_offsets,
96 edge_targets,
97 edge_kind_mask,
98 frontier_out,
99 node_count,
100 edge_count,
101 queue_capacity,
102 allow_mask,
103 })
104}
105
106#[must_use]
109pub fn csr_queue_strided_forward_traverse_with(
110 params: CsrQueueStridedForwardParams<'_>,
111) -> Program {
112 let CsrQueueStridedForwardParams {
113 active_queue,
114 queue_len,
115 edge_offsets,
116 edge_targets,
117 edge_kind_mask,
118 frontier_out,
119 node_count,
120 edge_count,
121 queue_capacity,
122 allow_mask,
123 } = params;
124 if node_count == 0 || queue_capacity == 0 {
125 return crate::invalid_output_program(CSR_QUEUE_STRIDED_FORWARD_OP_ID,
126 frontier_out,
127 DataType::U32,
128 format!(
129 "Fix: csr_queue_strided_forward_traverse requires node_count > 0 and queue_capacity > 0, got node_count={node_count} queue_capacity={queue_capacity}."
130 ),);
131 }
132 csr_queue_step_program(&CsrQueueStepSpec {
133 op_id: CSR_QUEUE_STRIDED_FORWARD_OP_ID,
134 builder_name: "csr_queue_strided_forward_traverse",
135 prefix: "qs",
136 workgroup_size: CSR_QUEUE_STRIDED_FORWARD_WORKGROUP_SIZE,
137 inputs: CsrQueueInputs {
138 active_queue,
139 queue_len,
140 edge_offsets,
141 edge_targets,
142 edge_kind_mask,
143 },
144 lanes: CsrQueueLanes::Team {
145 lanes: CSR_QUEUE_STRIDED_FORWARD_LANES_PER_SOURCE,
146 },
147 row_plan: CsrQueueRowPlan::ExpandAll,
148 emit: CsrQueueEmit::Frontier { frontier_out },
149 node_count,
150 edge_count,
151 queue_capacity,
152 allow_mask,
153 })
154}
155
156#[must_use]
158#[cfg(any(test, feature = "cpu-parity"))]
159#[allow(clippy::too_many_arguments)]
160pub fn csr_queue_strided_forward_traverse_cpu(
161 active_queue: &[u32],
162 queue_len: u32,
163 edge_offsets: &[u32],
164 edge_targets: &[u32],
165 edge_kind_mask: &[u32],
166 node_count: u32,
167 allow_mask: u32,
168) -> Vec<u32> {
169 try_csr_queue_strided_forward_traverse_cpu(
170 active_queue,
171 queue_len,
172 edge_offsets,
173 edge_targets,
174 edge_kind_mask,
175 node_count,
176 allow_mask,
177 )
178 .unwrap_or_else(|err| {
179 panic!("csr_queue_strided_forward_traverse CPU oracle received malformed input. {err}")
180 })
181}
182
183#[cfg(any(test, feature = "cpu-parity"))]
185#[allow(clippy::too_many_arguments)]
186pub fn try_csr_queue_strided_forward_traverse_cpu(
187 active_queue: &[u32],
188 queue_len: u32,
189 edge_offsets: &[u32],
190 edge_targets: &[u32],
191 edge_kind_mask: &[u32],
192 node_count: u32,
193 allow_mask: u32,
194) -> Result<Vec<u32>, String> {
195 try_csr_queue_forward_traverse_cpu(
196 active_queue,
197 queue_len,
198 edge_offsets,
199 edge_targets,
200 edge_kind_mask,
201 node_count,
202 allow_mask,
203 )
204}
205
206#[cfg(any(test, feature = "cpu-parity"))]
208#[allow(clippy::too_many_arguments)]
209pub fn try_csr_queue_strided_forward_traverse_cpu_into(
210 active_queue: &[u32],
211 queue_len: u32,
212 edge_offsets: &[u32],
213 edge_targets: &[u32],
214 edge_kind_mask: &[u32],
215 node_count: u32,
216 allow_mask: u32,
217 out: &mut Vec<u32>,
218) -> Result<(), String> {
219 try_csr_queue_forward_traverse_cpu_into(
220 active_queue,
221 queue_len,
222 edge_offsets,
223 edge_targets,
224 edge_kind_mask,
225 node_count,
226 allow_mask,
227 out,
228 )
229}
230
231#[cfg(feature = "inventory-registry")]
232inventory::submit! {
233 vyre_foundation::operation::OperationRegistration::primitive(
234 CSR_QUEUE_STRIDED_FORWARD_OP_ID,
235 || csr_queue_strided_forward_traverse(
236 "active_queue",
237 "queue_len",
238 "edge_offsets",
239 "edge_targets",
240 "edge_kind_mask",
241 "frontier_out",
242 4,
243 4,
244 2,
245 1,
246 ),
247 Some(|| {
248 let to_bytes = |w: &[u32]| crate::wire::pack_u32_slice(w);
249 vec![vec![
250 to_bytes(&[0, 3]), to_bytes(&[2]), to_bytes(&[0, 3, 3, 4, 4]), to_bytes(&[1, 2, 3, 0]), to_bytes(&[1, 2, 1, 1]), to_bytes(&[0]), ]]
257 }),
258 Some(|| {
259 let to_bytes = |w: &[u32]| crate::wire::pack_u32_slice(w);
260 vec![vec![to_bytes(&[0b1010])]]
261 }),
262 )
263}
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268
269 fn scalar_queue_forward(
270 active_queue: &[u32],
271 queue_len: u32,
272 edge_offsets: &[u32],
273 edge_targets: &[u32],
274 edge_kind_mask: &[u32],
275 node_count: u32,
276 allow_mask: u32,
277 ) -> Vec<u32> {
278 let mut out = vec![0u32; bitset_words(node_count) as usize];
279 let take = (queue_len as usize).min(active_queue.len());
280 for &src in &active_queue[..take] {
281 if src >= node_count {
282 continue;
283 }
284 for edge in edge_offsets[src as usize]..edge_offsets[src as usize + 1] {
285 let edge = edge as usize;
286 if edge_kind_mask[edge] & allow_mask == 0 {
287 continue;
288 }
289 let dst = edge_targets[edge];
290 out[dst as usize / 32] |= 1u32 << (dst % 32);
291 }
292 }
293 out
294 }
295
296 #[test]
297 fn dispatch_grid_assigns_32_lanes_per_queue_slot() {
298 assert_eq!(csr_queue_strided_forward_dispatch_grid(0), [1, 1, 1]);
299 assert_eq!(csr_queue_strided_forward_dispatch_grid(1), [1, 1, 1]);
300 assert_eq!(csr_queue_strided_forward_dispatch_grid(8), [1, 1, 1]);
301 assert_eq!(csr_queue_strided_forward_dispatch_grid(9), [2, 1, 1]);
302 assert_eq!(csr_queue_strided_forward_dispatch_grid(256), [32, 1, 1]);
303 }
304
305 #[test]
306 fn build_program_returns_well_formed_program() {
307 let program = csr_queue_strided_forward_traverse(
308 "queue", "len", "offsets", "targets", "kinds", "out", 64, 4096, 9, 0x55,
309 );
310 assert_eq!(
311 program.workgroup_size(),
312 CSR_QUEUE_STRIDED_FORWARD_WORKGROUP_SIZE
313 );
314 assert_eq!(program.buffers().len(), 6);
315 assert!(!program.stats().trap());
316 }
317
318 #[test]
319 fn generated_strided_cpu_matches_scalar_reference_on_skewed_rows() {
320 let mut seed = 0x51A7_7EED_u32;
321 for case in 0..4096u32 {
322 seed = seed.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
323 let node_count = 33 + (seed % 224);
324 let queue_capacity = 1 + (seed.rotate_left(5) % node_count);
325 let mut offsets = Vec::with_capacity(node_count as usize + 1);
326 let mut targets = Vec::new();
327 let mut masks = Vec::new();
328 offsets.push(0);
329 for src in 0..node_count {
330 seed ^= src.wrapping_mul(0x9E37_79B9).rotate_left((src & 15) + 1);
331 let degree = if src == case % node_count {
332 96 + (seed % 257)
333 } else {
334 seed % 5
335 };
336 for edge in 0..degree {
337 targets.push(src.wrapping_mul(17).wrapping_add(edge * 3 + seed) % node_count);
338 masks.push(if (edge ^ src ^ seed) & 3 == 0 { 2 } else { 1 });
339 }
340 offsets.push(targets.len() as u32);
341 }
342 let mut queue = Vec::with_capacity(queue_capacity as usize);
343 for slot in 0..queue_capacity {
344 queue.push(slot.wrapping_mul(7).wrapping_add(seed) % node_count);
345 }
346 let queue_len = queue_capacity.saturating_add(seed % 3);
347 let expected =
348 scalar_queue_forward(&queue, queue_len, &offsets, &targets, &masks, node_count, 1);
349
350 assert_eq!(
351 try_csr_queue_strided_forward_traverse_cpu(
352 &queue, queue_len, &offsets, &targets, &masks, node_count, 1,
353 ),
354 Ok(expected),
355 "generated skewed CSR queue case {case}"
356 );
357 }
358 }
359
360 #[test]
361 fn invalid_shape_returns_trap_program() {
362 let program = csr_queue_strided_forward_traverse(
363 "queue", "len", "offsets", "targets", "kinds", "out", 0, 0, 1, 1,
364 );
365
366 assert!(
367 program.stats().trap(),
368 "invalid node_count must compile to a trap program"
369 );
370 }
371
372 #[test]
373 fn offset_count_overflow_returns_trap_program_without_panic() {
374 let result = std::panic::catch_unwind(|| {
375 csr_queue_strided_forward_traverse(
376 "queue",
377 "len",
378 "offsets",
379 "targets",
380 "kinds",
381 "out",
382 u32::MAX,
383 0,
384 1,
385 1,
386 )
387 });
388
389 assert!(
390 result.is_ok(),
391 "CSR queue strided builder must reject offset-count overflow without panicking"
392 );
393 let program = result.unwrap();
394 assert!(program.stats().trap());
395 let entry = format!("{:?}", program.entry());
396 assert!(
397 entry.contains("node_count + 1 overflows u32"),
398 "Fix: trap must retain the CSR offset-count overflow diagnostic, got: {entry}"
399 );
400 }
401}