Skip to main content

nxpu_opt/
cse.rs

1//! Common subexpression elimination pass.
2//!
3//! Walks expression arenas, hashes each expression by opcode and operand
4//! indices, and rewrites duplicates to point to the first (canonical)
5//! occurrence. Dead duplicates are then cleaned up by a subsequent DCE pass.
6
7use std::collections::HashMap;
8use std::hash::{DefaultHasher, Hash, Hasher};
9
10use nxpu_ir::{Expression, Function, Handle, Module};
11
12use crate::Pass;
13use crate::dce::expression_operands;
14
15/// Common subexpression elimination pass.
16///
17/// Detects structurally identical expressions within each function's
18/// expression arena and rewrites all references to point to the canonical
19/// (first) occurrence. Subsequent DCE removes the dead duplicates.
20#[derive(Debug)]
21pub struct CommonSubexprElimination;
22
23impl Pass for CommonSubexprElimination {
24    fn name(&self) -> &str {
25        "cse"
26    }
27
28    fn run(&self, module: &mut Module) -> bool {
29        let mut changed = false;
30        // Run on global expressions.
31        changed |= run_on_arena(&mut module.global_expressions);
32        // Run on each function.
33        for (_, func) in module.functions.iter_mut() {
34            changed |= run_on_function(func);
35        }
36        // Run on entry point functions.
37        for ep in &mut module.entry_points {
38            changed |= run_on_function(&mut ep.function);
39        }
40        changed
41    }
42}
43
44/// Run CSE on a single function's expression arena.
45fn run_on_function(func: &mut Function) -> bool {
46    run_on_arena(&mut func.expressions)
47}
48
49/// Run CSE on an expression arena.
50///
51/// Returns `true` if any expressions were rewritten.
52fn run_on_arena(arena: &mut nxpu_ir::Arena<Expression>) -> bool {
53    // Phase 1: compute canonical mapping.
54    // Map hash → list of (handle, expression-discriminant) for collision resolution.
55    let mut seen: HashMap<u64, Handle<Expression>> = HashMap::new();
56    let mut remap: HashMap<Handle<Expression>, Handle<Expression>> = HashMap::new();
57
58    let handles: Vec<Handle<Expression>> = arena.iter().map(|(h, _)| h).collect();
59
60    for &handle in &handles {
61        let hash = hash_expression(&arena[handle], &remap);
62        if let Some(&canonical) = seen.get(&hash) {
63            // Verify structural equality (not just hash match).
64            if expressions_equal(&arena[canonical], &arena[handle], &remap) {
65                remap.insert(handle, canonical);
66                continue;
67            }
68        }
69        seen.insert(hash, handle);
70    }
71
72    if remap.is_empty() {
73        return false;
74    }
75
76    // Phase 2: rewrite operand references using the remap table.
77    for &handle in &handles {
78        if remap.contains_key(&handle) {
79            // This expression is dead (will be cleaned by DCE).
80            continue;
81        }
82        let rewritten = rewrite_operands(&arena[handle], &remap);
83        if let Some(new_expr) = rewritten {
84            arena[handle] = new_expr;
85        }
86    }
87
88    true
89}
90
91/// Hash an expression by its discriminant and operand indices.
92fn hash_expression(
93    expr: &Expression,
94    remap: &HashMap<Handle<Expression>, Handle<Expression>>,
95) -> u64 {
96    let mut hasher = DefaultHasher::new();
97    std::mem::discriminant(expr).hash(&mut hasher);
98
99    match expr {
100        Expression::Literal(lit) => {
101            // Hash the literal bytes for value equality.
102            format!("{lit:?}").hash(&mut hasher);
103        }
104        Expression::Binary { op, left, right } => {
105            std::mem::discriminant(op).hash(&mut hasher);
106            resolve(left, remap).index().hash(&mut hasher);
107            resolve(right, remap).index().hash(&mut hasher);
108        }
109        Expression::Unary { op, expr: inner } => {
110            std::mem::discriminant(op).hash(&mut hasher);
111            resolve(inner, remap).index().hash(&mut hasher);
112        }
113        Expression::Math {
114            fun,
115            arg,
116            arg1,
117            arg2,
118            arg3,
119        } => {
120            std::mem::discriminant(fun).hash(&mut hasher);
121            resolve(arg, remap).index().hash(&mut hasher);
122            if let Some(a) = arg1 {
123                resolve(a, remap).index().hash(&mut hasher);
124            }
125            if let Some(a) = arg2 {
126                resolve(a, remap).index().hash(&mut hasher);
127            }
128            if let Some(a) = arg3 {
129                resolve(a, remap).index().hash(&mut hasher);
130            }
131        }
132        Expression::Load { pointer } => {
133            resolve(pointer, remap).index().hash(&mut hasher);
134        }
135        Expression::Access { base, index } => {
136            resolve(base, remap).index().hash(&mut hasher);
137            resolve(index, remap).index().hash(&mut hasher);
138        }
139        Expression::AccessIndex { base, index } => {
140            resolve(base, remap).index().hash(&mut hasher);
141            index.hash(&mut hasher);
142        }
143        Expression::Compose { ty, components } => {
144            ty.index().hash(&mut hasher);
145            for c in components {
146                resolve(c, remap).index().hash(&mut hasher);
147            }
148        }
149        Expression::Select {
150            condition,
151            accept,
152            reject,
153        } => {
154            resolve(condition, remap).index().hash(&mut hasher);
155            resolve(accept, remap).index().hash(&mut hasher);
156            resolve(reject, remap).index().hash(&mut hasher);
157        }
158        Expression::Splat { size, value } => {
159            size.hash(&mut hasher);
160            resolve(value, remap).index().hash(&mut hasher);
161        }
162        Expression::Swizzle {
163            size,
164            vector,
165            pattern,
166        } => {
167            size.hash(&mut hasher);
168            resolve(vector, remap).index().hash(&mut hasher);
169            pattern.hash(&mut hasher);
170        }
171        Expression::As {
172            expr: inner,
173            kind,
174            convert,
175        } => {
176            resolve(inner, remap).index().hash(&mut hasher);
177            std::mem::discriminant(kind).hash(&mut hasher);
178            convert.hash(&mut hasher);
179        }
180        Expression::ArrayLength(e) => {
181            resolve(e, remap).index().hash(&mut hasher);
182        }
183        // Non-dedupable expressions: hash by handle index to keep them unique.
184        Expression::FunctionArgument(i) => {
185            i.hash(&mut hasher);
186        }
187        Expression::GlobalVariable(h) => {
188            h.index().hash(&mut hasher);
189        }
190        Expression::LocalVariable(h) => {
191            h.index().hash(&mut hasher);
192        }
193        Expression::CallResult(_) | Expression::AtomicResult { .. } | Expression::ZeroValue(_) => {
194            // Side-effectful or unique — never merge.
195            // Use a unique value to prevent collisions.
196            std::ptr::from_ref(expr).hash(&mut hasher);
197        }
198    }
199
200    hasher.finish()
201}
202
203/// Check structural equality of two expressions, resolving handles through the remap.
204fn expressions_equal(
205    a: &Expression,
206    b: &Expression,
207    remap: &HashMap<Handle<Expression>, Handle<Expression>>,
208) -> bool {
209    match (a, b) {
210        (Expression::Literal(la), Expression::Literal(lb)) => {
211            format!("{la:?}") == format!("{lb:?}")
212        }
213        (
214            Expression::Binary {
215                op: op_a,
216                left: la,
217                right: ra,
218            },
219            Expression::Binary {
220                op: op_b,
221                left: lb,
222                right: rb,
223            },
224        ) => {
225            op_a == op_b
226                && resolve(la, remap) == resolve(lb, remap)
227                && resolve(ra, remap) == resolve(rb, remap)
228        }
229        (Expression::Unary { op: op_a, expr: ea }, Expression::Unary { op: op_b, expr: eb }) => {
230            op_a == op_b && resolve(ea, remap) == resolve(eb, remap)
231        }
232        (
233            Expression::Math {
234                fun: fa,
235                arg: a0,
236                arg1: a1,
237                arg2: a2,
238                arg3: a3,
239            },
240            Expression::Math {
241                fun: fb,
242                arg: b0,
243                arg1: b1,
244                arg2: b2,
245                arg3: b3,
246            },
247        ) => {
248            fa == fb
249                && resolve(a0, remap) == resolve(b0, remap)
250                && resolve_opt(a1.as_ref(), remap) == resolve_opt(b1.as_ref(), remap)
251                && resolve_opt(a2.as_ref(), remap) == resolve_opt(b2.as_ref(), remap)
252                && resolve_opt(a3.as_ref(), remap) == resolve_opt(b3.as_ref(), remap)
253        }
254        (Expression::Load { pointer: pa }, Expression::Load { pointer: pb }) => {
255            resolve(pa, remap) == resolve(pb, remap)
256        }
257        (
258            Expression::Compose {
259                ty: ta,
260                components: ca,
261            },
262            Expression::Compose {
263                ty: tb,
264                components: cb,
265            },
266        ) => {
267            ta == tb
268                && ca.len() == cb.len()
269                && ca
270                    .iter()
271                    .zip(cb.iter())
272                    .all(|(x, y)| resolve(x, remap) == resolve(y, remap))
273        }
274        (
275            Expression::Access {
276                base: ba,
277                index: ia,
278            },
279            Expression::Access {
280                base: bb,
281                index: ib,
282            },
283        ) => resolve(ba, remap) == resolve(bb, remap) && resolve(ia, remap) == resolve(ib, remap),
284        (
285            Expression::AccessIndex {
286                base: ba,
287                index: ia,
288            },
289            Expression::AccessIndex {
290                base: bb,
291                index: ib,
292            },
293        ) => resolve(ba, remap) == resolve(bb, remap) && ia == ib,
294        (
295            Expression::Select {
296                condition: ca,
297                accept: aa,
298                reject: ra,
299            },
300            Expression::Select {
301                condition: cb,
302                accept: ab,
303                reject: rb,
304            },
305        ) => {
306            resolve(ca, remap) == resolve(cb, remap)
307                && resolve(aa, remap) == resolve(ab, remap)
308                && resolve(ra, remap) == resolve(rb, remap)
309        }
310        (
311            Expression::Splat {
312                size: sa,
313                value: va,
314            },
315            Expression::Splat {
316                size: sb,
317                value: vb,
318            },
319        ) => sa == sb && resolve(va, remap) == resolve(vb, remap),
320        (
321            Expression::Swizzle {
322                size: sa,
323                vector: va,
324                pattern: pa,
325            },
326            Expression::Swizzle {
327                size: sb,
328                vector: vb,
329                pattern: pb,
330            },
331        ) => sa == sb && resolve(va, remap) == resolve(vb, remap) && pa == pb,
332        (
333            Expression::As {
334                expr: ea,
335                kind: ka,
336                convert: ca,
337            },
338            Expression::As {
339                expr: eb,
340                kind: kb,
341                convert: cb,
342            },
343        ) => resolve(ea, remap) == resolve(eb, remap) && ka == kb && ca == cb,
344        (Expression::ArrayLength(ea), Expression::ArrayLength(eb)) => {
345            resolve(ea, remap) == resolve(eb, remap)
346        }
347        (Expression::FunctionArgument(ia), Expression::FunctionArgument(ib)) => ia == ib,
348        (Expression::GlobalVariable(ha), Expression::GlobalVariable(hb)) => ha == hb,
349        (Expression::LocalVariable(ha), Expression::LocalVariable(hb)) => ha == hb,
350        (Expression::ZeroValue(ta), Expression::ZeroValue(tb)) => ta == tb,
351        _ => false,
352    }
353}
354
355/// Resolve a handle through the remap table.
356fn resolve(
357    h: &Handle<Expression>,
358    remap: &HashMap<Handle<Expression>, Handle<Expression>>,
359) -> Handle<Expression> {
360    remap.get(h).copied().unwrap_or(*h)
361}
362
363fn resolve_opt(
364    h: Option<&Handle<Expression>>,
365    remap: &HashMap<Handle<Expression>, Handle<Expression>>,
366) -> Option<Handle<Expression>> {
367    h.map(|handle| resolve(handle, remap))
368}
369
370/// Rewrite operand references in an expression using the remap table.
371/// Returns `Some(new_expr)` if any operand was remapped, `None` otherwise.
372fn rewrite_operands(
373    expr: &Expression,
374    remap: &HashMap<Handle<Expression>, Handle<Expression>>,
375) -> Option<Expression> {
376    let operands = expression_operands(expr);
377    let any_remapped = operands.iter().any(|h| remap.contains_key(h));
378    if !any_remapped {
379        return None;
380    }
381
382    Some(match expr {
383        Expression::Binary { op, left, right } => Expression::Binary {
384            op: *op,
385            left: resolve(left, remap),
386            right: resolve(right, remap),
387        },
388        Expression::Unary { op, expr: inner } => Expression::Unary {
389            op: *op,
390            expr: resolve(inner, remap),
391        },
392        Expression::Math {
393            fun,
394            arg,
395            arg1,
396            arg2,
397            arg3,
398        } => Expression::Math {
399            fun: *fun,
400            arg: resolve(arg, remap),
401            arg1: resolve_opt(arg1.as_ref(), remap),
402            arg2: resolve_opt(arg2.as_ref(), remap),
403            arg3: resolve_opt(arg3.as_ref(), remap),
404        },
405        Expression::Load { pointer } => Expression::Load {
406            pointer: resolve(pointer, remap),
407        },
408        Expression::Access { base, index } => Expression::Access {
409            base: resolve(base, remap),
410            index: resolve(index, remap),
411        },
412        Expression::AccessIndex { base, index } => Expression::AccessIndex {
413            base: resolve(base, remap),
414            index: *index,
415        },
416        Expression::Compose { ty, components } => Expression::Compose {
417            ty: *ty,
418            components: components.iter().map(|c| resolve(c, remap)).collect(),
419        },
420        Expression::Select {
421            condition,
422            accept,
423            reject,
424        } => Expression::Select {
425            condition: resolve(condition, remap),
426            accept: resolve(accept, remap),
427            reject: resolve(reject, remap),
428        },
429        Expression::Splat { size, value } => Expression::Splat {
430            size: *size,
431            value: resolve(value, remap),
432        },
433        Expression::Swizzle {
434            size,
435            vector,
436            pattern,
437        } => Expression::Swizzle {
438            size: *size,
439            vector: resolve(vector, remap),
440            pattern: *pattern,
441        },
442        Expression::As {
443            expr: inner,
444            kind,
445            convert,
446        } => Expression::As {
447            expr: resolve(inner, remap),
448            kind: *kind,
449            convert: *convert,
450        },
451        Expression::ArrayLength(e) => Expression::ArrayLength(resolve(e, remap)),
452        // These don't have expression operands to remap.
453        other => other.clone(),
454    })
455}
456
457#[cfg(test)]
458mod tests {
459    use super::*;
460    use nxpu_ir::{BinaryOp, Literal, Range, Statement};
461
462    #[test]
463    fn cse_identical_literals() {
464        let mut func = Function::new("test");
465        let _a = func
466            .expressions
467            .append(Expression::Literal(Literal::F32(42.0)));
468        let _b = func
469            .expressions
470            .append(Expression::Literal(Literal::F32(42.0)));
471        let changed = run_on_function(&mut func);
472        assert!(changed);
473    }
474
475    #[test]
476    fn cse_identical_binary_ops() {
477        let mut func = Function::new("test");
478        let a = func
479            .expressions
480            .append(Expression::Literal(Literal::F32(1.0)));
481        let b = func
482            .expressions
483            .append(Expression::Literal(Literal::F32(2.0)));
484        let _add1 = func.expressions.append(Expression::Binary {
485            op: BinaryOp::Add,
486            left: a,
487            right: b,
488        });
489        let _add2 = func.expressions.append(Expression::Binary {
490            op: BinaryOp::Add,
491            left: a,
492            right: b,
493        });
494        let changed = run_on_function(&mut func);
495        assert!(changed);
496    }
497
498    #[test]
499    fn cse_different_ops_not_merged() {
500        let mut func = Function::new("test");
501        let a = func
502            .expressions
503            .append(Expression::Literal(Literal::F32(1.0)));
504        let b = func
505            .expressions
506            .append(Expression::Literal(Literal::F32(2.0)));
507        let _add = func.expressions.append(Expression::Binary {
508            op: BinaryOp::Add,
509            left: a,
510            right: b,
511        });
512        let _mul = func.expressions.append(Expression::Binary {
513            op: BinaryOp::Multiply,
514            left: a,
515            right: b,
516        });
517        let changed = run_on_function(&mut func);
518        assert!(!changed);
519    }
520
521    #[test]
522    fn cse_different_operands_not_merged() {
523        let mut func = Function::new("test");
524        let a = func
525            .expressions
526            .append(Expression::Literal(Literal::F32(1.0)));
527        let b = func
528            .expressions
529            .append(Expression::Literal(Literal::F32(2.0)));
530        let c = func
531            .expressions
532            .append(Expression::Literal(Literal::F32(3.0)));
533        let _add1 = func.expressions.append(Expression::Binary {
534            op: BinaryOp::Add,
535            left: a,
536            right: b,
537        });
538        let _add2 = func.expressions.append(Expression::Binary {
539            op: BinaryOp::Add,
540            left: a,
541            right: c,
542        });
543        let changed = run_on_function(&mut func);
544        assert!(!changed);
545    }
546
547    #[test]
548    fn cse_transitive_dedup() {
549        // x = a + b, y = a + b (dup of x)
550        // z = x * c, w = y * c → after CSE y→x, so w = x * c = z (dup)
551        let mut func = Function::new("test");
552        let a = func
553            .expressions
554            .append(Expression::Literal(Literal::F32(1.0)));
555        let b = func
556            .expressions
557            .append(Expression::Literal(Literal::F32(2.0)));
558        let c = func
559            .expressions
560            .append(Expression::Literal(Literal::F32(3.0)));
561        let x = func.expressions.append(Expression::Binary {
562            op: BinaryOp::Add,
563            left: a,
564            right: b,
565        });
566        let y = func.expressions.append(Expression::Binary {
567            op: BinaryOp::Add,
568            left: a,
569            right: b,
570        });
571        let _z = func.expressions.append(Expression::Binary {
572            op: BinaryOp::Multiply,
573            left: x,
574            right: c,
575        });
576        let _w = func.expressions.append(Expression::Binary {
577            op: BinaryOp::Multiply,
578            left: y,
579            right: c,
580        });
581        let changed = run_on_function(&mut func);
582        assert!(changed);
583    }
584
585    #[test]
586    fn cse_preserves_non_duplicates() {
587        let mut func = Function::new("test");
588        let _a = func
589            .expressions
590            .append(Expression::Literal(Literal::F32(1.0)));
591        let _b = func
592            .expressions
593            .append(Expression::Literal(Literal::F32(2.0)));
594        let _c = func
595            .expressions
596            .append(Expression::Literal(Literal::F32(3.0)));
597        let changed = run_on_function(&mut func);
598        assert!(!changed);
599    }
600
601    #[test]
602    fn cse_empty_module() {
603        let mut module = Module::default();
604        let pass = CommonSubexprElimination;
605        let changed = pass.run(&mut module);
606        assert!(!changed);
607    }
608
609    #[test]
610    fn cse_on_entry_point_functions() {
611        use nxpu_ir::EntryPoint;
612
613        let mut module = Module::default();
614        let mut func = Function::new("ep");
615        let a = func
616            .expressions
617            .append(Expression::Literal(Literal::F32(7.0)));
618        let _b = func
619            .expressions
620            .append(Expression::Literal(Literal::F32(7.0)));
621        func.body
622            .push(Statement::Emit(Range::from_index_range(0..2)));
623        func.body.push(Statement::Return { value: Some(a) });
624        module.entry_points.push(EntryPoint {
625            name: "main".into(),
626            workgroup_size: [1, 1, 1],
627            function: func,
628        });
629        let pass = CommonSubexprElimination;
630        let changed = pass.run(&mut module);
631        assert!(changed);
632    }
633
634    #[test]
635    fn cse_on_global_expressions() {
636        let mut module = Module::default();
637        module
638            .global_expressions
639            .append(Expression::Literal(Literal::F32(99.0)));
640        module
641            .global_expressions
642            .append(Expression::Literal(Literal::F32(99.0)));
643        let pass = CommonSubexprElimination;
644        let changed = pass.run(&mut module);
645        assert!(changed);
646    }
647
648    #[test]
649    fn cse_plus_dce_combined() {
650        use crate::DeadCodeElimination;
651
652        let mut module = Module::default();
653        let mut func = Function::new("ep");
654        let a = func
655            .expressions
656            .append(Expression::Literal(Literal::F32(1.0)));
657        let b = func
658            .expressions
659            .append(Expression::Literal(Literal::F32(2.0)));
660        let add1 = func.expressions.append(Expression::Binary {
661            op: BinaryOp::Add,
662            left: a,
663            right: b,
664        });
665        let _add2 = func.expressions.append(Expression::Binary {
666            op: BinaryOp::Add,
667            left: a,
668            right: b,
669        });
670
671        func.body
672            .push(Statement::Emit(Range::from_index_range(0..4)));
673        func.body.push(Statement::Return { value: Some(add1) });
674
675        module.entry_points.push(nxpu_ir::EntryPoint {
676            name: "main".into(),
677            workgroup_size: [1, 1, 1],
678            function: func,
679        });
680
681        let cse = CommonSubexprElimination;
682        let cse_changed = cse.run(&mut module);
683        assert!(cse_changed);
684
685        let dce = DeadCodeElimination;
686        let dce_changed = dce.run(&mut module);
687        // DCE should be able to clean up the dead duplicate.
688        // The dead emit range covering the duplicate may be removed.
689        let _ = dce_changed;
690    }
691
692    #[test]
693    fn cse_reduces_expression_count() {
694        // Measure that CSE actually reduces the number of unique expressions
695        // referenced by the IR (via remap count).
696        let mut func = Function::new("test");
697        let a = func
698            .expressions
699            .append(Expression::Literal(Literal::F32(1.0)));
700        let b = func
701            .expressions
702            .append(Expression::Literal(Literal::F32(2.0)));
703        // Create 5 identical Add(a,b) expressions.
704        for _ in 0..5 {
705            func.expressions.append(Expression::Binary {
706                op: BinaryOp::Add,
707                left: a,
708                right: b,
709            });
710        }
711        let expr_count_before = func.expressions.len();
712        assert_eq!(expr_count_before, 7); // 2 literals + 5 adds
713
714        let changed = run_on_function(&mut func);
715        assert!(changed);
716
717        // After CSE, 4 of the 5 Add expressions are remapped to the canonical one.
718        // The arena size doesn't shrink (arena is append-only), but all duplicates
719        // now point to the same canonical expression. Verify by counting unique
720        // expression hashes after CSE rewrites.
721        let mut unique = std::collections::HashSet::new();
722        for (_, expr) in func.expressions.iter() {
723            unique.insert(format!("{expr:?}"));
724        }
725        // All 5 Add(a,b) should still be in the arena, but 4 are dead duplicates
726        // that subsequent DCE would remove. The key metric is that CSE reported
727        // a change, meaning duplicates were detected and remapped.
728        assert!(unique.len() <= expr_count_before);
729    }
730}