use vyre_foundation::ir::{AtomicOp, Expr, MemoryOrdering, Node, Program};
use vyre_foundation::visit::node_map::any_descendant;
use vyre_primitives::fixpoint::persistent_fixpoint::{
cpu_ref, persistent_fixpoint, persistent_fixpoint_grid, OP_ID_GRID,
PERSISTENT_FIXPOINT_WORKGROUP_SIZE,
};
use vyre_reference::value::Value;
fn lane() -> Expr {
Expr::InvocationId { axis: 0 }
}
fn wave_nodes(program: &Program) -> &[Node] {
match program.entry() {
[Node::Region {
generator, body, ..
}] => {
assert_eq!(
generator.as_str(),
OP_ID_GRID,
"the grid builder must attribute its region to its own op id"
);
body
}
other => panic!("expected exactly one generator Region at the entry, got {other:?}"),
}
}
fn count_nodes<P>(program: &Program, mut pred: P) -> usize
where
P: FnMut(&Node) -> bool,
{
let mut total = 0usize;
for node in program.entry() {
let _ = any_descendant(node, &mut |candidate| {
if pred(candidate) {
total += 1;
}
false
});
}
total
}
fn atomic_or_indices(program: &Program, buffer: &str) -> Vec<u32> {
let mut indices = Vec::new();
for node in program.entry() {
let _ = any_descendant(node, &mut |candidate| {
if let Node::Let {
value:
Expr::Atomic {
op: AtomicOp::Or,
buffer: target,
index,
..
},
..
} = candidate
{
if target.as_str() == buffer {
match index.as_ref() {
Expr::LitU32(word) => indices.push(*word),
other => panic!("flag index must be a literal word, got {other:?}"),
}
}
}
false
});
}
indices
}
fn identity_body(words: u32) -> Vec<Node> {
vec![Node::if_then(
Expr::lt(lane(), Expr::u32(words)),
vec![Node::store("next", lane(), Expr::load("current", lane()))],
)]
}
fn or_const_body(words: u32, mask: u32) -> Vec<Node> {
vec![Node::if_then(
Expr::lt(lane(), Expr::u32(words)),
vec![Node::store(
"next",
lane(),
Expr::bitor(Expr::load("current", lane()), Expr::u32(mask)),
)],
)]
}
fn shift_or_body(words: u32) -> Vec<Node> {
vec![Node::if_then(
Expr::lt(lane(), Expr::u32(words)),
vec![Node::store(
"next",
lane(),
Expr::bitor(
Expr::load("current", lane()),
Expr::shl(Expr::load("current", lane()), Expr::u32(8)),
),
)],
)]
}
fn carry_body(words: u32) -> Vec<Node> {
vec![Node::if_then(
Expr::lt(lane(), Expr::u32(words)),
vec![
Node::if_then(
Expr::eq(lane(), Expr::u32(0)),
vec![Node::store("next", lane(), Expr::load("current", lane()))],
),
Node::if_then(
Expr::gt(lane(), Expr::u32(0)),
vec![Node::store(
"next",
lane(),
Expr::bitor(
Expr::load("current", lane()),
Expr::load("current", Expr::sub(lane(), Expr::u32(1))),
),
)],
),
],
)]
}
fn pack(words: &[u32]) -> Value {
Value::from(vyre_primitives::wire::pack_u32_slice(words))
}
fn unpack(value: &Value) -> Vec<u32> {
value
.to_bytes()
.chunks_exact(4)
.map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect()
}
struct Run {
state: Vec<u32>,
changed: Vec<u32>,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
enum Order {
Forward,
Reversed,
}
fn eval(program: &Program, inputs: &[Value], order: Order) -> Vec<Value> {
match order {
Order::Forward => vyre_reference::reference_eval(program, inputs),
Order::Reversed => vyre_reference::reference_eval_lane_reversed(program, inputs),
}
.expect("reference evaluation must succeed")
}
fn run_grid_ordered(body: Vec<Node>, seed: &[u32], max_iterations: u32, order: Order) -> Run {
let words = u32::try_from(seed.len()).expect("seed length must fit u32");
let program =
persistent_fixpoint_grid(body, "current", "next", "changed", words, max_iterations);
let outputs = eval(
&program,
&[
pack(seed),
pack(&vec![0u32; seed.len()]),
pack(&vec![0u32; max_iterations.max(1) as usize]),
],
order,
);
Run {
state: unpack(&outputs[0]),
changed: unpack(&outputs[2]),
}
}
fn run_workgroup_ordered(body: Vec<Node>, seed: &[u32], max_iterations: u32, order: Order) -> Run {
let words = u32::try_from(seed.len()).expect("seed length must fit u32");
let program = persistent_fixpoint(body, "current", "next", "changed", words, max_iterations);
let outputs = eval(
&program,
&[pack(seed), pack(&vec![0u32; seed.len()]), pack(&[0u32])],
order,
);
Run {
state: unpack(&outputs[0]),
changed: unpack(&outputs[2]),
}
}
fn run_grid(body: Vec<Node>, seed: &[u32], max_iterations: u32) -> Run {
run_grid_ordered(body, seed, max_iterations, Order::Forward)
}
fn run_workgroup(body: Vec<Node>, seed: &[u32], max_iterations: u32) -> Run {
run_workgroup_ordered(body, seed, max_iterations, Order::Forward)
}
fn passes_from_flags(changed: &[u32], max_iterations: u32) -> u32 {
changed
.iter()
.position(|word| *word == 0)
.map_or(max_iterations, |first_zero| {
u32::try_from(first_zero).expect("flag index must fit u32") + 1
})
}
mod host_orchestration;
mod parity_and_races;
mod structure;