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(
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 edge_count = edge_targets.len() as u32;
let shape = ProgramGraphShape::new(node_count, edge_count.max(1));
let program = persistent_bfs(shape, "frontier_in", "frontier_out", allow_mask, max_iters);
let words = bitset_words(node_count) as usize;
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();
dispatch_with_grid_sync_split(
&CpuRefBackend,
&program,
&borrowed,
&DispatchConfig::default(),
)
.expect("Fix: persistent_bfs grid-sync split dispatch must succeed on a valid graph.")
} else {
reference_eval(
&program,
&inputs
.iter()
.map(|bytes| vyre_reference::value::Value::from(bytes.as_slice()))
.collect::<Vec<_>>(),
)
.expect("Fix: persistent_bfs reference dispatch must succeed on a valid graph.")
.into_iter()
.map(|value| value.to_bytes())
.collect()
};
let mut frontier_out = read_named_output(&program, &outputs, "frontier_out");
frontier_out.truncate(words);
let changed = read_named_output(&program, &outputs, "changed")[0];
let converged = read_named_output(&program, &outputs, "converged")[0];
(frontier_out, changed, converged)
}
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_device(
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)
}
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)
}
#[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_device(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_device(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);
}
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)
}
#[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 edge_count = targets.len() as u32;
let shape = ProgramGraphShape::new(node_count, edge_count.max(1));
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_device_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 edge_count = edge_targets.len() as u32;
let shape = ProgramGraphShape::new(node_count, edge_count.max(1));
let program = persistent_bfs_batch(
shape,
"frontier_in",
"frontier_out",
"changed",
"converged",
query_count,
allow_mask,
max_iters,
);
let words = bitset_words(node_count) as usize;
let total_words = words * query_count.max(1) as usize;
let inputs = build_inputs(
&program,
edge_offsets,
edge_targets,
edge_kind_mask,
frontier_in,
);
let grid = persistent_bfs_batch_dispatch_grid(node_count, query_count);
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 = Some(grid);
dispatch_with_grid_sync_split(&CpuRefBackend, &program, &borrowed, &config).expect(
"Fix: persistent_bfs_batch grid-sync split dispatch must succeed on a valid graph.",
)
} else {
reference_eval_with_grid(
&program,
&inputs
.iter()
.map(|bytes| vyre_reference::value::Value::from(bytes.as_slice()))
.collect::<Vec<_>>(),
grid,
)
.expect("Fix: persistent_bfs_batch reference dispatch must succeed on a valid graph.")
.into_iter()
.map(|value| value.to_bytes())
.collect()
};
let mut frontier_out = read_named_output(&program, &outputs, "frontier_out");
frontier_out.truncate(total_words);
let mut changed = read_named_output(&program, &outputs, "changed");
changed.truncate(query_count as usize);
let mut converged = read_named_output(&program, &outputs, "converged");
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 query_count = seeds.len() as u32;
let mut frontier_in = 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)"
);
frontier_in.extend_from_slice(seed);
}
let (device_frontier, device_changed, device_converged) = run_device_batch(
node_count,
query_count,
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 words = bitset_words(node_count) as usize;
let leaf = node_count - 1;
let mut leaf_seed = vec![0u32; words];
leaf_seed[(leaf / 32) as usize] = 1 << (leaf % 32);
let seeds = vec![root_seed, leaf_seed];
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)]);
}