use vyre_foundation::ir::{BufferDecl, Node, Program};
use vyre_primitives::graph::level_wave::{level_wave_program, level_wave_program_with_buffers};
#[must_use]
pub fn build_callee_before_caller_program(
step_body: Vec<Node>,
depth_buf: &str,
max_depth: u32,
function_count: u32,
) -> Program {
use crate::observability::{bump, level_wave_pass_calls};
bump(&level_wave_pass_calls);
level_wave_program(step_body, depth_buf, max_depth, function_count)
}
#[must_use]
pub fn build_callee_before_caller_program_with_buffers(
step_body: Vec<Node>,
depth_buf: &str,
extra_buffers: Vec<BufferDecl>,
max_depth: u32,
function_count: u32,
) -> Program {
use crate::observability::{bump, level_wave_pass_calls};
bump(&level_wave_pass_calls);
level_wave_program_with_buffers(
step_body,
depth_buf,
extra_buffers,
max_depth,
function_count,
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn builds_nonempty_program() {
let body = vec![Node::barrier()];
let program = build_callee_before_caller_program(body, "depths", 4, 16);
assert_ne!(program.entry().len(), 0);
}
#[test]
fn zero_depth_still_builds() {
let body = vec![Node::barrier()];
let program = build_callee_before_caller_program(body, "depths", 0, 1);
assert_eq!(program.workgroup_size(), [256, 1, 1]);
assert!(!program.buffers().is_empty());
}
#[test]
fn callee_before_caller_commits_children_before_parents() {
use vyre_foundation::ir::{BufferAccess, BufferDecl, DataType, Expr};
use vyre_reference::reference_eval;
use vyre_reference::value::Value;
let t = Expr::InvocationId { axis: 0 };
let step_body = vec![
Node::let_bind("c", Expr::load("callee", t.clone())),
Node::store(
"out",
t.clone(),
Expr::add(Expr::u32(1), Expr::load("out", Expr::var("c"))),
),
];
let extra_buffers = vec![
BufferDecl::storage("callee", 1, BufferAccess::ReadOnly, DataType::U32).with_count(4),
BufferDecl::storage("out", 2, BufferAccess::ReadWrite, DataType::U32).with_count(4),
];
let program = build_callee_before_caller_program_with_buffers(
step_body,
"depths",
extra_buffers,
4, 4, );
let pack = |data: &[u32]| Value::from(vyre_primitives::wire::pack_u32_slice(data));
let inputs = vec![
pack(&[0, 1, 2, 3]), pack(&[0, 0, 1, 2]), pack(&[0, 0, 0, 0]), ];
let results = reference_eval(&program, &inputs).expect("Fix: level-wave pass eval failed");
let out: Vec<u32> = results[0]
.to_bytes()
.chunks_exact(4)
.map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
assert_eq!(
out,
vec![1, 2, 3, 4],
"each caller must read its callee's committed value: fn0=1, fn1=1+1, fn2=1+2, fn3=1+3"
);
}
}