use super::*;
use crate::wire::pack_u32_slice;
use vyre_driver::grid_sync::{contains_grid_sync, dispatch_with_grid_sync_split};
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_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 edge_count = edge_targets.len() as u32;
let shape = ProgramGraphShape::new(node_count, edge_count.max(1));
let program = persistent_bfs_with_density(
shape,
"frontier_in",
"frontier_out",
DENSITY_ACTIVE_BUFFER,
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: density 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: density 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 mut density = read_named_output(&program, &outputs, DENSITY_ACTIVE_BUFFER);
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_device_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
}
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)
}
#[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_device_batch_density(
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 edge_count = edge_targets.len() as u32;
let shape = ProgramGraphShape::new(node_count, edge_count.max(1));
let program = persistent_bfs_batch_with_density(
shape,
"frontier_in",
"frontier_out",
"changed",
"converged",
DENSITY_ACTIVE_BUFFER,
query_count,
allow_mask,
max_iters,
);
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: density 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: density persistent_bfs_batch reference dispatch must succeed on a valid graph.",
)
.into_iter()
.map(|value| value.to_bytes())
.collect()
};
let mut density = read_named_output(&program, &outputs, DENSITY_ACTIVE_BUFFER);
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 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 one frontier bitset"
);
frontier_in.extend_from_slice(seed);
}
let device_density = run_device_batch_density(
node_count,
query_count,
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 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 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]);
}