vyre_primitives/graph/exploded/
program_key.rs1#[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#[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}