use super::*;
use crate::wire::pack_u32_slice;
use vyre_driver::backend::VyreBackend;
use vyre_driver::grid_sync::{
contains_grid_sync, dispatch_with_grid_sync_split, dispatch_with_grid_sync_split_via,
};
use vyre_driver::DispatchConfig;
use vyre_driver_reference::CpuRefBackend;
use vyre_foundation::ir::{BufferAccess, Program};
use vyre_reference::{output_index, reference_eval, reference_eval_with_grid};
fn build_inputs(
program: &Program,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
frontier_in: &[u32],
) -> Vec<Vec<u8>> {
let mut inputs: Vec<Vec<u8>> = Vec::new();
for buffer in program.buffers() {
if buffer.access() == BufferAccess::Workgroup {
continue;
}
let count = buffer.count() as usize;
let mut data = match buffer.name() {
"pg_edge_offsets" => edge_offsets.to_vec(),
"pg_edge_targets" => edge_targets.to_vec(),
"pg_edge_kind_mask" => edge_kind_mask.to_vec(),
"frontier_in" => frontier_in.to_vec(),
_ => Vec::new(),
};
data.resize(count, 0);
inputs.push(pack_u32_slice(&data));
}
inputs
}
fn read_named_output(program: &Program, outputs: &[Vec<u8>], name: &str) -> Vec<u32> {
let idx = output_index(program, name)
.unwrap_or_else(|| panic!("Fix: persistent_bfs must expose the `{name}` output buffer."));
outputs[idx]
.chunks_exact(4)
.map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect()
}
fn run_device(
program: &Program,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
frontier_in: &[u32],
grid: Option<[u32; 3]>,
reads: &[&str],
) -> Vec<Vec<u32>> {
let inputs = build_inputs(
program,
edge_offsets,
edge_targets,
edge_kind_mask,
frontier_in,
);
let outputs: Vec<Vec<u8>> = if contains_grid_sync(program) {
let borrowed: Vec<&[u8]> = inputs.iter().map(Vec::as_slice).collect();
let mut config = DispatchConfig::default();
config.dispatch_grid = grid;
dispatch_with_grid_sync_split(&CpuRefBackend, program, &borrowed, &config)
.expect("Fix: persistent_bfs grid-sync split dispatch must succeed on a valid graph.")
} else {
let values: Vec<vyre_reference::value::Value> = inputs
.iter()
.map(|bytes| vyre_reference::value::Value::from(bytes.as_slice()))
.collect();
match grid {
None => reference_eval(program, &values),
Some(grid) => reference_eval_with_grid(program, &values, grid),
}
.expect("Fix: persistent_bfs reference dispatch must succeed on a valid graph.")
.into_iter()
.map(|value| value.to_bytes())
.collect()
};
reads
.iter()
.map(|name| read_named_output(program, &outputs, name))
.collect()
}
fn reverse_chain(node_count: u32) -> (Vec<u32>, Vec<u32>, Vec<u32>, Vec<u32>) {
let mut offsets = vec![0u32];
for i in 0..node_count {
offsets.push(i);
}
let targets: Vec<u32> = (0..node_count.saturating_sub(1)).collect();
let masks = vec![1u32; targets.len()];
let words = bitset_words(node_count) as usize;
let mut seed = vec![0u32; words];
let top = node_count - 1;
seed[(top / 32) as usize] = 1 << (top % 32);
(offsets, targets, masks, seed)
}
fn grid_sync_two_level(fanout: u32) -> (u32, Vec<u32>, Vec<u32>, Vec<u32>, Vec<u32>) {
let node_count = fanout + 2; let leaf = node_count - 1;
let mut targets: Vec<u32> = (1..=fanout).collect();
targets.push(leaf); let mut offsets = vec![0u32, fanout];
for _ in 2..=node_count {
offsets.push(fanout + 1);
}
let masks = vec![1u32; targets.len()];
let words = bitset_words(node_count) as usize;
let mut seed = vec![0u32; words];
seed[0] = 1; (node_count, offsets, targets, masks, seed)
}
fn shape_for(node_count: u32, edge_targets: &[u32]) -> ProgramGraphShape {
ProgramGraphShape::new(node_count, (edge_targets.len() as u32).max(1))
}
fn pack_seeds(node_count: u32, seeds: &[Vec<u32>]) -> Vec<u32> {
let words = bitset_words(node_count) as usize;
let mut packed = Vec::with_capacity(words * seeds.len());
for seed in seeds {
assert_eq!(
seed.len(),
words,
"each batch seed must be exactly one frontier bitset ({words} words)"
);
packed.extend_from_slice(seed);
}
packed
}
fn run_converged(
node_count: u32,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
frontier_in: &[u32],
allow_mask: u32,
max_iters: u32,
) -> (Vec<u32>, u32, u32) {
let program = persistent_bfs(
shape_for(node_count, edge_targets),
"frontier_in",
"frontier_out",
allow_mask,
max_iters,
);
let reads = run_device(
&program,
edge_offsets,
edge_targets,
edge_kind_mask,
frontier_in,
None,
&["frontier_out", "changed", "converged"],
);
let mut frontier_out = reads[0].clone();
frontier_out.truncate(bitset_words(node_count) as usize);
(frontier_out, reads[1][0], reads[2][0])
}
fn assert_device_matches_oracle(
node_count: u32,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
frontier_in: &[u32],
allow_mask: u32,
max_iters: u32,
) -> (u32, bool) {
let (frontier, outcome) = try_cpu_ref_converged(
node_count,
edge_offsets,
edge_targets,
edge_kind_mask,
frontier_in,
allow_mask,
max_iters,
)
.expect("Fix: CPU oracle must accept a valid graph.");
let (device_frontier, device_changed, device_converged) = run_converged(
node_count,
edge_offsets,
edge_targets,
edge_kind_mask,
frontier_in,
allow_mask,
max_iters,
);
assert_eq!(
device_frontier, frontier,
"node_count={node_count} max_iters={max_iters}: device frontier must equal the CPU oracle."
);
assert_eq!(
device_changed, outcome.changed,
"node_count={node_count} max_iters={max_iters}: device changed flag must equal the CPU oracle."
);
assert_eq!(
device_converged,
u32::from(outcome.converged),
"node_count={node_count} max_iters={max_iters}: device converged word must equal the CPU oracle (converged={}).",
outcome.converged
);
(outcome.changed, outcome.converged)
}
#[test]
fn single_workgroup_converged_word_matches_oracle_below_and_above_diameter() {
let (offsets, targets, masks, seed) = reverse_chain(4);
let (changed, converged) =
assert_device_matches_oracle(4, &offsets, &targets, &masks, &seed, 0xFFFF_FFFF, 2);
assert_eq!((changed, converged), (1, false));
let (changed, converged) =
assert_device_matches_oracle(4, &offsets, &targets, &masks, &seed, 0xFFFF_FFFF, 8);
assert_eq!((changed, converged), (1, true));
}
#[test]
fn single_workgroup_converged_is_false_when_full_set_reached_on_last_allowed_step() {
let (offsets, targets, masks, seed) = reverse_chain(5);
let (device_frontier, device_changed, device_converged) =
run_converged(5, &offsets, &targets, &masks, &seed, 0xFFFF_FFFF, 4);
assert_eq!(device_frontier, vec![0b11111]);
assert_eq!(device_changed, 1);
assert_eq!(
device_converged, 0,
"Fix: a full set reached only on the last allowed step must report converged=0 on the device."
);
assert_device_matches_oracle(5, &offsets, &targets, &masks, &seed, 0xFFFF_FFFF, 4);
}
#[test]
fn single_workgroup_converged_is_true_when_seed_is_already_a_fixpoint() {
let offsets = [0u32, 1, 1];
let targets = [1u32];
let masks = [1u32];
let seed = [0b11u32];
let (device_frontier, device_changed, device_converged) =
run_converged(2, &offsets, &targets, &masks, &seed, 0xFFFF_FFFF, 4);
assert_eq!(device_frontier, vec![0b11]);
assert_eq!(device_changed, 0);
assert_eq!(device_converged, 1);
assert_device_matches_oracle(2, &offsets, &targets, &masks, &seed, 0xFFFF_FFFF, 4);
}
#[test]
fn grid_sync_seed_copy_preserves_every_frontier_word() {
let node_count = 8193;
let offsets = vec![0_u32; node_count as usize + 1];
let targets = [0_u32];
let masks = [1_u32];
let mut seed = (0..bitset_words(node_count))
.map(|word| 0x9e37_79b9_u32.wrapping_mul(word.wrapping_add(1)))
.collect::<Vec<_>>();
let trailing_bits = node_count % 32;
if trailing_bits != 0 {
let mask = (1_u32 << trailing_bits) - 1;
*seed
.last_mut()
.expect("Fix: a nonempty graph must have one frontier word") &= mask;
}
let (device_frontier, device_changed, device_converged) =
run_converged(node_count, &offsets, &targets, &masks, &seed, u32::MAX, 0);
assert_eq!(
device_frontier, seed,
"Fix: grid-sync persistent_bfs must copy every seed frontier word."
);
assert_eq!(device_changed, 0);
assert_eq!(device_converged, 0);
}
#[test]
fn grid_sync_converged_word_matches_oracle_across_the_budget_boundary() {
let (node_count, offsets, targets, masks, seed) = grid_sync_two_level(256);
assert!(node_count > 256, "must exercise the grid-sync path");
let (changed, converged) = assert_device_matches_oracle(
node_count,
&offsets,
&targets,
&masks,
&seed,
0xFFFF_FFFF,
1,
);
assert_eq!((changed, converged), (1, false));
let (changed, converged) = assert_device_matches_oracle(
node_count,
&offsets,
&targets,
&masks,
&seed,
0xFFFF_FFFF,
2,
);
assert_eq!((changed, converged), (1, false));
let (changed, converged) = assert_device_matches_oracle(
node_count,
&offsets,
&targets,
&masks,
&seed,
0xFFFF_FFFF,
3,
);
assert_eq!((changed, converged), (1, true));
}
#[test]
fn grid_sync_converged_word_matches_oracle_through_the_closure_split_entry() {
let (node_count, offsets, targets, masks, seed) = grid_sync_two_level(256);
assert!(node_count > 256, "must exercise the grid-sync path");
let shape = shape_for(node_count, &targets);
let dispatch = |program: &Program,
inputs: &[&[u8]],
grid: Option<[u32; 3]>,
outputs: &mut Vec<Vec<u8>>|
-> Result<(), String> {
let mut config = DispatchConfig::default();
config.grid_override = grid;
config.fixpoint_iterations = Some(1);
CpuRefBackend
.dispatch_borrowed_into(program, inputs, &config, outputs)
.map_err(|error| error.to_string())
};
for (max_iters, expect_converged) in [(1u32, 0u32), (2, 0), (3, 1)] {
let program = persistent_bfs(shape, "frontier_in", "frontier_out", 0xFFFF_FFFF, max_iters);
assert!(
contains_grid_sync(&program),
"258-node persistent_bfs must be a grid-sync program"
);
let inputs = build_inputs(&program, &offsets, &targets, &masks, &seed);
let borrowed: Vec<&[u8]> = inputs.iter().map(Vec::as_slice).collect();
let outputs = dispatch_with_grid_sync_split_via(
&program,
&borrowed,
&DispatchConfig::default(),
&dispatch,
)
.expect("Fix: persistent_bfs closure-split dispatch must succeed on a valid graph.");
let converged = read_named_output(&program, &outputs, "converged")[0];
let (_, oracle) = try_cpu_ref_converged(
node_count,
&offsets,
&targets,
&masks,
&seed,
0xFFFF_FFFF,
max_iters,
)
.expect("Fix: CPU oracle must accept a valid graph.");
assert_eq!(
converged,
expect_converged,
"max_iters={max_iters}: closure-split converged word must equal the oracle (converged={}).",
oracle.converged
);
assert_eq!(converged, u32::from(oracle.converged));
}
}
fn run_converged_batch(
node_count: u32,
query_count: u32,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
frontier_in: &[u32],
allow_mask: u32,
max_iters: u32,
) -> (Vec<u32>, Vec<u32>, Vec<u32>) {
let program = persistent_bfs_batch(
shape_for(node_count, edge_targets),
"frontier_in",
"frontier_out",
"changed",
"converged",
query_count,
allow_mask,
max_iters,
);
let reads = run_device(
&program,
edge_offsets,
edge_targets,
edge_kind_mask,
frontier_in,
Some(persistent_bfs_batch_dispatch_grid(node_count, query_count)),
&["frontier_out", "changed", "converged"],
);
let total_words = bitset_words(node_count) as usize * query_count.max(1) as usize;
let mut frontier_out = reads[0].clone();
frontier_out.truncate(total_words);
let mut changed = reads[1].clone();
changed.truncate(query_count as usize);
let mut converged = reads[2].clone();
converged.truncate(query_count as usize);
(frontier_out, changed, converged)
}
fn assert_batch_device_matches_oracle(
node_count: u32,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
seeds: &[Vec<u32>],
allow_mask: u32,
max_iters: u32,
) -> Vec<(u32, bool)> {
let words = bitset_words(node_count) as usize;
let frontier_in = pack_seeds(node_count, seeds);
let (device_frontier, device_changed, device_converged) = run_converged_batch(
node_count,
seeds.len() as u32,
edge_offsets,
edge_targets,
edge_kind_mask,
&frontier_in,
allow_mask,
max_iters,
);
let mut oracle_outcomes = Vec::with_capacity(seeds.len());
for (query, seed) in seeds.iter().enumerate() {
let (frontier, outcome) = try_cpu_ref_converged(
node_count,
edge_offsets,
edge_targets,
edge_kind_mask,
seed,
allow_mask,
max_iters,
)
.expect("Fix: CPU oracle must accept a valid graph.");
let start = query * words;
let end = start + words;
assert_eq!(
&device_frontier[start..end],
frontier.as_slice(),
"query {query}: batch device frontier must equal the CPU oracle."
);
assert_eq!(
device_changed[query], outcome.changed,
"query {query}: batch device changed flag must equal the CPU oracle."
);
assert_eq!(
device_converged[query],
u32::from(outcome.converged),
"query {query}: batch device converged word must equal the CPU oracle (converged={}).",
outcome.converged
);
oracle_outcomes.push((outcome.changed, outcome.converged));
}
oracle_outcomes
}
#[test]
fn batch_single_workgroup_converged_matches_oracle_across_queries_and_budget() {
let (offsets, targets, masks, _seed) = reverse_chain(4);
let seeds = vec![vec![0b1000u32], vec![0b0010u32], vec![0b1111u32]];
let outcomes =
assert_batch_device_matches_oracle(4, &offsets, &targets, &masks, &seeds, 0xFFFF_FFFF, 2);
assert_eq!(outcomes, vec![(1, false), (1, true), (0, true)]);
let outcomes =
assert_batch_device_matches_oracle(4, &offsets, &targets, &masks, &seeds, 0xFFFF_FFFF, 8);
assert_eq!(outcomes, vec![(1, true), (1, true), (0, true)]);
}
#[test]
fn batch_grid_sync_converged_matches_oracle_across_queries_and_budget() {
let (node_count, offsets, targets, masks, root_seed) = grid_sync_two_level(256);
assert!(node_count > 256, "must exercise the grid-sync batch path");
let seeds = vec![root_seed, leaf_seed(node_count)];
let outcomes = assert_batch_device_matches_oracle(
node_count,
&offsets,
&targets,
&masks,
&seeds,
0xFFFF_FFFF,
1,
);
assert_eq!(outcomes, vec![(1, false), (0, true)]);
let outcomes = assert_batch_device_matches_oracle(
node_count,
&offsets,
&targets,
&masks,
&seeds,
0xFFFF_FFFF,
3,
);
assert_eq!(outcomes, vec![(1, true), (0, true)]);
}
fn leaf_seed(node_count: u32) -> Vec<u32> {
let mut seed = vec![0u32; bitset_words(node_count) as usize];
let leaf = node_count - 1;
seed[(leaf / 32) as usize] = 1 << (leaf % 32);
seed
}
fn run_density(
node_count: u32,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
frontier_in: &[u32],
allow_mask: u32,
max_iters: u32,
) -> (Vec<u32>, Vec<u32>) {
let program = persistent_bfs_with_density(
shape_for(node_count, edge_targets),
"frontier_in",
"frontier_out",
DENSITY_ACTIVE_BUFFER,
allow_mask,
max_iters,
);
let reads = run_device(
&program,
edge_offsets,
edge_targets,
edge_kind_mask,
frontier_in,
None,
&["frontier_out", DENSITY_ACTIVE_BUFFER],
);
let mut frontier_out = reads[0].clone();
frontier_out.truncate(bitset_words(node_count) as usize);
let mut density = reads[1].clone();
density.truncate(max_iters as usize);
(frontier_out, density)
}
fn assert_device_density_matches_oracle(
node_count: u32,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
frontier_in: &[u32],
allow_mask: u32,
max_iters: u32,
) -> Vec<u32> {
let (frontier, _outcome, active) = try_cpu_ref_density(
node_count,
edge_offsets,
edge_targets,
edge_kind_mask,
frontier_in,
allow_mask,
max_iters,
)
.expect("Fix: CPU density oracle must accept a valid graph.");
assert_eq!(
active.len(),
max_iters as usize,
"oracle density array must have exactly max_iters entries"
);
let (device_frontier, device_density) = run_density(
node_count,
edge_offsets,
edge_targets,
edge_kind_mask,
frontier_in,
allow_mask,
max_iters,
);
assert_eq!(
device_frontier, frontier,
"node_count={node_count} max_iters={max_iters}: device frontier must equal the CPU oracle."
);
assert_eq!(
device_density, active,
"node_count={node_count} max_iters={max_iters}: device density array must equal the CPU oracle."
);
active
}
#[test]
fn single_workgroup_density_matches_oracle_growing_then_flat_after_convergence() {
let (offsets, targets, masks, seed) = reverse_chain(4);
let active =
assert_device_density_matches_oracle(4, &offsets, &targets, &masks, &seed, 0xFFFF_FFFF, 2);
assert_eq!(active, vec![2, 3]);
let active =
assert_device_density_matches_oracle(4, &offsets, &targets, &masks, &seed, 0xFFFF_FFFF, 8);
assert_eq!(active, vec![2, 3, 4, 4, 4, 4, 4, 4]);
}
#[test]
fn single_workgroup_density_repeats_seed_popcount_when_seed_is_already_a_fixpoint() {
let offsets = [0u32, 1, 1];
let targets = [1u32];
let masks = [1u32];
let seed = [0b11u32];
let active =
assert_device_density_matches_oracle(2, &offsets, &targets, &masks, &seed, 0xFFFF_FFFF, 4);
assert_eq!(active, vec![2, 2, 2, 2]);
}
#[test]
fn grid_sync_density_matches_oracle_across_the_budget_boundary() {
let (node_count, offsets, targets, masks, seed) = grid_sync_two_level(256);
assert!(node_count > 256, "must exercise the grid-sync density path");
let active = assert_device_density_matches_oracle(
node_count,
&offsets,
&targets,
&masks,
&seed,
0xFFFF_FFFF,
1,
);
assert_eq!(active, vec![257]);
let active = assert_device_density_matches_oracle(
node_count,
&offsets,
&targets,
&masks,
&seed,
0xFFFF_FFFF,
2,
);
assert_eq!(active, vec![257, 258]);
let active = assert_device_density_matches_oracle(
node_count,
&offsets,
&targets,
&masks,
&seed,
0xFFFF_FFFF,
3,
);
assert_eq!(active, vec![257, 258, 258]);
}
fn run_density_batch(
node_count: u32,
query_count: u32,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
frontier_in: &[u32],
allow_mask: u32,
max_iters: u32,
) -> Vec<u32> {
let program = persistent_bfs_batch_with_density(
shape_for(node_count, edge_targets),
"frontier_in",
"frontier_out",
"changed",
"converged",
DENSITY_ACTIVE_BUFFER,
query_count,
allow_mask,
max_iters,
);
let reads = run_device(
&program,
edge_offsets,
edge_targets,
edge_kind_mask,
frontier_in,
Some(persistent_bfs_batch_dispatch_grid(node_count, query_count)),
&[DENSITY_ACTIVE_BUFFER],
);
let mut density = reads[0].clone();
density.truncate((query_count * max_iters) as usize);
density
}
fn assert_batch_device_density_matches_oracle(
node_count: u32,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
seeds: &[Vec<u32>],
allow_mask: u32,
max_iters: u32,
) -> Vec<u32> {
let frontier_in = pack_seeds(node_count, seeds);
let device_density = run_density_batch(
node_count,
seeds.len() as u32,
edge_offsets,
edge_targets,
edge_kind_mask,
&frontier_in,
allow_mask,
max_iters,
);
let mut oracle = Vec::with_capacity(seeds.len() * max_iters as usize);
for (query, seed) in seeds.iter().enumerate() {
let (_frontier, _outcome, active) = try_cpu_ref_density(
node_count,
edge_offsets,
edge_targets,
edge_kind_mask,
seed,
allow_mask,
max_iters,
)
.expect("Fix: CPU density oracle must accept a valid graph.");
let start = query * max_iters as usize;
let end = start + max_iters as usize;
assert_eq!(
&device_density[start..end],
active.as_slice(),
"query {query}: batch device density must equal the CPU oracle."
);
oracle.extend_from_slice(&active);
}
oracle
}
#[test]
fn batch_single_workgroup_density_matches_oracle_across_queries() {
let (offsets, targets, masks, _seed) = reverse_chain(4);
let seeds = vec![vec![0b1000u32], vec![0b0010u32], vec![0b1111u32]];
let oracle = assert_batch_device_density_matches_oracle(
4,
&offsets,
&targets,
&masks,
&seeds,
0xFFFF_FFFF,
2,
);
assert_eq!(oracle, vec![2, 3, 2, 2, 4, 4]);
}
#[test]
fn batch_grid_sync_density_matches_oracle_across_queries() {
let (node_count, offsets, targets, masks, root_seed) = grid_sync_two_level(256);
assert!(
node_count > 256,
"must exercise the grid-sync batch density path"
);
let seeds = vec![root_seed, leaf_seed(node_count)];
let oracle = assert_batch_device_density_matches_oracle(
node_count,
&offsets,
&targets,
&masks,
&seeds,
0xFFFF_FFFF,
3,
);
assert_eq!(oracle, vec![257, 258, 258, 1, 1, 1]);
}