Skip to main content

vyre_primitives/graph/exploded/
program_key.rs

1#[cfg(any(test, feature = "cpu-parity"))]
2use super::layout::IfdsCsrProgramCacheKey;
3#[cfg(any(test, feature = "cpu-parity"))]
4use super::validation::max_ifds_col_count;
5#[cfg(any(test, feature = "cpu-parity"))]
6use vyre_foundation::ir::{Expr, Node, Program};
7
8/// Recover the exploded IFDS program cache key baked into a generated CSR
9/// builder [`Program`].
10///
11/// Test and parity dispatchers use this to route GPU-shaped byte inputs through
12/// the CPU reference without re-deriving dimensions from padded buffers alone.
13#[cfg(any(test, feature = "cpu-parity"))]
14pub fn ifds_program_cache_key_from_program(
15    program: &Program,
16) -> Result<IfdsCsrProgramCacheKey, String> {
17    let intra_count = loop_upper_bound(program, "intra_i")
18        .ok_or_else(|| "Fix: exploded IFDS program missing intra_i loop bound.".to_string())?;
19    let inter_count = loop_upper_bound(program, "inter_i")
20        .ok_or_else(|| "Fix: exploded IFDS program missing inter_i loop bound.".to_string())?;
21    let gen_count = loop_upper_bound(program, "gen_i")
22        .ok_or_else(|| "Fix: exploded IFDS program missing gen_i loop bound.".to_string())?;
23    let kill_count = loop_upper_bound(program, "kill_i")
24        .ok_or_else(|| "Fix: exploded IFDS program missing kill_i loop bound.".to_string())?;
25    let facts_per_proc = loop_upper_bound(program, "fact")
26        .ok_or_else(|| "Fix: exploded IFDS program missing fact loop bound.".to_string())?;
27    let total_nodes = loop_upper_bound(program, "prefix_row")
28        .or_else(|| loop_upper_bound(program, "cursor_row"))
29        .ok_or_else(|| "Fix: exploded IFDS program missing total_nodes loop bound.".to_string())?;
30
31    let num_procs = upper_limit_for_var(program, "intra_p")
32        .or_else(|| upper_limit_for_var(program, "inter_sp"))
33        .ok_or_else(|| "Fix: exploded IFDS program missing num_procs bound.".to_string())?;
34    let blocks_per_proc = upper_limit_for_var(program, "intra_dst_b")
35        .or_else(|| upper_limit_for_var(program, "intra_src_b"))
36        .or_else(|| upper_limit_for_var(program, "inter_sb"))
37        .ok_or_else(|| "Fix: exploded IFDS program missing blocks_per_proc bound.".to_string())?;
38
39    let slots_per_proc = blocks_per_proc
40        .checked_mul(facts_per_proc)
41        .ok_or_else(|| "Fix: exploded IFDS blocks*facts overflowed u32.".to_string())?;
42    let expected_total = num_procs
43        .checked_mul(slots_per_proc)
44        .ok_or_else(|| "Fix: exploded IFDS procs*blocks*facts overflowed u32.".to_string())?;
45    if expected_total != total_nodes {
46        return Err(format!(
47            "Fix: exploded IFDS program shape mismatch: procs={num_procs} blocks={blocks_per_proc} facts={facts_per_proc} implies total_nodes={expected_total}, program loop bound={total_nodes}."
48        ));
49    }
50
51    let max_col_count = max_ifds_col_count(intra_count, inter_count, gen_count, facts_per_proc)
52        .ok_or_else(|| "Fix: exploded IFDS maximum column count overflowed u32.".to_string())?;
53
54    Ok(IfdsCsrProgramCacheKey {
55        num_procs,
56        blocks_per_proc,
57        facts_per_proc,
58        intra_count,
59        inter_count,
60        gen_count,
61        kill_count,
62        max_col_count,
63    })
64}
65
66#[cfg(any(test, feature = "cpu-parity"))]
67fn loop_upper_bound(program: &Program, var: &str) -> Option<u32> {
68    use vyre_foundation::transform::visit::walk_nodes;
69
70    let mut found: Option<u32> = None;
71    walk_nodes(program, |node| {
72        if let Node::Loop {
73            var: loop_var, to, ..
74        } = node
75        {
76            if loop_var.as_str() == var {
77                if let Expr::LitU32(limit) = to {
78                    found = Some(*limit);
79                }
80            }
81        }
82    });
83    found
84}
85
86#[cfg(any(test, feature = "cpu-parity"))]
87fn upper_limit_for_var(program: &Program, var: &str) -> Option<u32> {
88    use vyre_foundation::ir::BinOp;
89    use vyre_foundation::transform::visit::walk_exprs;
90
91    let mut found: Option<u32> = None;
92    walk_exprs(program, |expr| {
93        if let Expr::BinOp {
94            op: BinOp::Lt,
95            left,
96            right,
97        } = expr
98        {
99            if let (Expr::Var(name), Expr::LitU32(limit)) = (left.as_ref(), right.as_ref()) {
100                if name.as_str() == var {
101                    found = Some(*limit);
102                }
103            }
104        }
105    });
106    found
107}