Skip to main content

vyre_primitives/graph/
csr_queue_strided.rs

1//! Row-strided sparse CSR expansion for skewed active queues.
2//!
3//! `csr_frontier_queue::csr_queue_forward_traverse` maps one lane to one queued
4//! source row. That is the right shape for tiny rows, but power-law graphs can
5//! put thousands of edges behind one active source and leave the rest of the GPU
6//! idle. This primitive keeps the same queue ABI and assigns a fixed lane team
7//! to each queued source:
8//!
9//! ```text
10//! queue index = global_lane / 32
11//! edge lane   = global_lane % 32
12//! for e = row_start + edge_lane; e < row_end; e += 32:
13//!     emit edge target
14//! ```
15//!
16//! It is intentionally a separate Program builder so callers can keep the
17//! one-lane-per-source kernel for low-degree graphs and select this path only
18//! when row skew is large enough to amortize the extra lanes.
19
20use 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
33/// Canonical op id for row-strided queue-driven CSR expansion.
34pub const CSR_QUEUE_STRIDED_FORWARD_OP_ID: &str =
35    "vyre-primitives::graph::csr_queue_strided_forward_traverse";
36
37/// Fixed lane team assigned to each queued source row.
38pub const CSR_QUEUE_STRIDED_FORWARD_LANES_PER_SOURCE: u32 = 32;
39
40/// Workgroup shape for row-strided queue-driven CSR expansion.
41pub const CSR_QUEUE_STRIDED_FORWARD_WORKGROUP_SIZE: [u32; 3] = [256, 1, 1];
42
43/// Dispatch grid that launches one 32-lane team for every queue slot.
44#[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/// Positional inputs for [`csr_queue_strided_forward_traverse`].
52#[derive(Clone, Copy, Debug)]
53pub struct CsrQueueStridedForwardParams<'a> {
54    /// Compacted queue of active source nodes.
55    pub active_queue: &'a str,
56    /// Single-element resident length of `active_queue`.
57    pub queue_len: &'a str,
58    /// CSR row pointers, `node_count + 1` entries.
59    pub edge_offsets: &'a str,
60    /// CSR edge destinations.
61    pub edge_targets: &'a str,
62    /// Per-edge kind bits tested against `allow_mask`.
63    pub edge_kind_mask: &'a str,
64    /// Packed bitset the reached destinations are ORed into.
65    pub frontier_out: &'a str,
66    /// Node count the CSR row pointers and destination bounds are sized by.
67    pub node_count: u32,
68    /// Logical edge count the edge-slot bound check uses.
69    pub edge_count: u32,
70    /// Static capacity of `active_queue`.
71    pub queue_capacity: u32,
72    /// Edge kinds this traversal is allowed to follow.
73    pub allow_mask: u32,
74}
75
76/// Build a GPU program that expands queued CSR source rows with a fixed lane
77/// team per row.
78#[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/// Build a GPU program that expands queued CSR source rows with a fixed lane
107/// team per row.
108#[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/// CPU reference for the row-strided queue traversal.
157#[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/// Fallible CPU reference for the row-strided queue traversal.
184#[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/// Fallible CPU reference into caller-owned storage.
207#[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]),            // active_queue
251                to_bytes(&[2]),               // queue_len
252                to_bytes(&[0, 3, 3, 4, 4]),   // edge_offsets
253                to_bytes(&[1, 2, 3, 0]),      // edge_targets
254                to_bytes(&[1, 2, 1, 1]),      // edge_kind_mask
255                to_bytes(&[0]),               // frontier_out
256            ]]
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}