vyre_foundation/ir_inner/model/program/
canonical.rs1use std::sync::Arc;
2
3use crate::ir_inner::model::expr::Expr;
4use crate::ir_inner::model::node::Node;
5use crate::ir_inner::model::spec_types::BinOp;
6
7use super::{meta::buffer_decl_canonical_key, BufferDecl, Program};
8
9impl Program {
10 #[must_use]
17 pub fn canonicalized(&self) -> Self {
18 let mut buffers = self.buffers().to_vec();
19 sort_buffers(&mut buffers);
20 let mut ctx = CanonicalCtx::default();
21 self.with_rewritten_entry(ctx.canonicalize_nodes(self.entry()))
22 .with_rewritten_buffers(buffers)
23 }
24
25 #[must_use]
32 pub fn canonical_wire_bytes(&self) -> Result<Vec<u8>, crate::error::IrError> {
33 let canonical = self.canonicalized();
34 let stats = canonical.stats();
40 let estimate = 256
41 + stats.node_count.saturating_mul(48)
42 + canonical.buffers().len().saturating_mul(64);
43 let mut out = Vec::with_capacity(estimate);
44 crate::serial::wire::encode::to_wire_into(&canonical, &mut out)
45 .map_err(|message| crate::error::IrError::WireFormatValidation { message })?;
46 Ok(out)
47 }
48
49 pub fn canonical_wire_hash(&self) -> Result<blake3::Hash, crate::error::IrError> {
56 self.canonical_wire_bytes()
57 .map(|bytes| blake3::hash(&bytes))
58 }
59}
60
61fn sort_buffers(buffers: &mut [BufferDecl]) {
62 buffers.sort_by_cached_key(buffer_decl_canonical_key);
63}
64
65#[derive(Default)]
66struct CanonicalCtx {
67 left_key: Vec<u8>,
68 right_key: Vec<u8>,
69}
70
71impl CanonicalCtx {
72 fn canonicalize_nodes(&mut self, nodes: &[Node]) -> Vec<Node> {
73 let mut out = Vec::with_capacity(nodes.len());
74 for node in nodes {
75 push_canonical_node(&mut out, self.canonicalize_node(node));
76 }
77 out
78 }
79
80 fn canonicalize_node(&mut self, node: &Node) -> Node {
81 match node {
82 Node::Let { name, value } => Node::Let {
83 name: name.clone(),
84 value: self.canonicalize_expr(value),
85 },
86 Node::Assign { name, value } => Node::Assign {
87 name: name.clone(),
88 value: self.canonicalize_expr(value),
89 },
90 Node::Store {
91 buffer,
92 index,
93 value,
94 } => Node::Store {
95 buffer: buffer.clone(),
96 index: self.canonicalize_expr(index),
97 value: self.canonicalize_expr(value),
98 },
99 Node::If {
100 cond,
101 then,
102 otherwise,
103 } => Node::If {
104 cond: self.canonicalize_expr(cond),
105 then: self.canonicalize_nodes(then),
106 otherwise: self.canonicalize_nodes(otherwise),
107 },
108 Node::Loop {
109 var,
110 from,
111 to,
112 body,
113 } => Node::Loop {
114 var: var.clone(),
115 from: self.canonicalize_expr(from),
116 to: self.canonicalize_expr(to),
117 body: self.canonicalize_nodes(body),
118 },
119 Node::Block(children) => Node::Block(self.canonicalize_nodes(children)),
120 Node::Region {
121 generator,
122 source_region,
123 body,
124 } => Node::Region {
125 generator: generator.clone(),
126 source_region: source_region.clone(),
127 body: Arc::new(self.canonicalize_nodes(body)),
128 },
129 Node::AsyncLoad {
130 source,
131 destination,
132 offset,
133 size,
134 tag,
135 } => Node::AsyncLoad {
136 source: source.clone(),
137 destination: destination.clone(),
138 offset: Box::new(self.canonicalize_expr(offset)),
139 size: Box::new(self.canonicalize_expr(size)),
140 tag: tag.clone(),
141 },
142 Node::AsyncStore {
143 source,
144 destination,
145 offset,
146 size,
147 tag,
148 } => Node::AsyncStore {
149 source: source.clone(),
150 destination: destination.clone(),
151 offset: Box::new(self.canonicalize_expr(offset)),
152 size: Box::new(self.canonicalize_expr(size)),
153 tag: tag.clone(),
154 },
155 Node::Trap { address, tag } => Node::Trap {
156 address: Box::new(self.canonicalize_expr(address)),
157 tag: tag.clone(),
158 },
159 Node::IndirectDispatch {
160 count_buffer,
161 count_offset,
162 } => Node::IndirectDispatch {
163 count_buffer: count_buffer.clone(),
164 count_offset: *count_offset,
165 },
166 Node::AllReduce { buffer, op, group } => Node::AllReduce {
167 buffer: buffer.clone(),
168 op: *op,
169 group: *group,
170 },
171 Node::AllGather {
172 input,
173 output,
174 group,
175 } => Node::AllGather {
176 input: input.clone(),
177 output: output.clone(),
178 group: *group,
179 },
180 Node::ReduceScatter {
181 input,
182 output,
183 op,
184 group,
185 } => Node::ReduceScatter {
186 input: input.clone(),
187 output: output.clone(),
188 op: *op,
189 group: *group,
190 },
191 Node::Broadcast {
192 buffer,
193 root,
194 group,
195 } => Node::Broadcast {
196 buffer: buffer.clone(),
197 root: *root,
198 group: *group,
199 },
200 Node::AsyncWait { tag } => Node::AsyncWait { tag: tag.clone() },
201 Node::Resume { tag } => Node::Resume { tag: tag.clone() },
202 Node::Return => Node::Return,
203 Node::Barrier { ordering } => Node::barrier_with_ordering(*ordering),
204 Node::Opaque(extension) => Node::Opaque(Arc::clone(extension)),
205 }
206 }
207
208 fn canonicalize_expr(&mut self, expr: &Expr) -> Expr {
209 crate::optimizer::rewrite::rewrite_expr(expr, &mut |candidate| {
210 let Expr::BinOp { op, left, right } = candidate else {
211 return None;
212 };
213 if !should_swap_operands(*op, left, right, &mut self.left_key, &mut self.right_key) {
214 return None;
215 }
216 Some(Expr::BinOp {
217 op: *op,
218 left: Box::new((**right).clone()),
219 right: Box::new((**left).clone()),
220 })
221 })
222 .into_owned()
223 }
224}
225
226fn push_canonical_node(out: &mut Vec<Node>, node: Node) {
227 match node {
228 Node::Block(children) if can_splice_block(&children) => out.extend(children),
229 other => out.push(other),
230 }
231}
232
233fn can_splice_block(nodes: &[Node]) -> bool {
234 nodes.iter().all(|node| !matches!(node, Node::Let { .. }))
235}
236
237fn should_swap_operands(
238 op: BinOp,
239 left: &Expr,
240 right: &Expr,
241 left_key: &mut Vec<u8>,
242 right_key: &mut Vec<u8>,
243) -> bool {
244 if !is_commutative_binop(op) {
245 return false;
246 }
247 match (is_literal(left), is_literal(right)) {
248 (true, false) => true,
249 (false, true) => false,
250 (true, true) => {
251 expr_wire_key_cmp(left, right, left_key, right_key).is_gt()
257 }
258 (false, false) => {
259 can_sort_all_operands(op) && expr_wire_key_cmp(left, right, left_key, right_key).is_gt()
260 }
261 }
262}
263
264fn expr_wire_key_cmp(
265 left: &Expr,
266 right: &Expr,
267 left_key: &mut Vec<u8>,
268 right_key: &mut Vec<u8>,
269) -> std::cmp::Ordering {
270 left_key.clear();
271 right_key.clear();
272 append_expr_wire_key(left_key, left);
273 append_expr_wire_key(right_key, right);
274 left_key.as_slice().cmp(right_key.as_slice())
275}
276
277fn append_expr_wire_key(key: &mut Vec<u8>, expr: &Expr) {
278 if let Err(error) = crate::serial::wire::encode::put_expr(key, expr) {
279 key.clear();
280 key.extend_from_slice(b"VYRE-CANONICAL-EXPR-WIRE-ERROR\0");
281 key.extend_from_slice(error.as_bytes());
282 }
283}
284
285fn is_commutative_binop(op: BinOp) -> bool {
286 matches!(
287 op,
288 BinOp::Add
289 | BinOp::WrappingAdd
290 | BinOp::SaturatingAdd
291 | BinOp::Mul
292 | BinOp::SaturatingMul
293 | BinOp::BitAnd
294 | BinOp::BitOr
295 | BinOp::BitXor
296 | BinOp::Eq
297 | BinOp::Ne
298 | BinOp::And
299 | BinOp::Or
300 | BinOp::Min
301 | BinOp::Max
302 | BinOp::AbsDiff
303 )
304}
305
306fn can_sort_all_operands(op: BinOp) -> bool {
307 matches!(
314 op,
315 BinOp::WrappingAdd
316 | BinOp::SaturatingAdd
317 | BinOp::SaturatingMul
318 | BinOp::BitAnd
319 | BinOp::BitOr
320 | BinOp::BitXor
321 | BinOp::Eq
322 | BinOp::Ne
323 | BinOp::And
324 | BinOp::Or
325 | BinOp::AbsDiff
326 )
327}
328
329fn is_literal(expr: &Expr) -> bool {
330 matches!(
331 expr,
332 Expr::LitU32(_) | Expr::LitI32(_) | Expr::LitF32(_) | Expr::LitBool(_)
333 )
334}