1use 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#[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 changed |= run_on_arena(&mut module.global_expressions);
32 for (_, func) in module.functions.iter_mut() {
34 changed |= run_on_function(func);
35 }
36 for ep in &mut module.entry_points {
38 changed |= run_on_function(&mut ep.function);
39 }
40 changed
41 }
42}
43
44fn run_on_function(func: &mut Function) -> bool {
46 run_on_arena(&mut func.expressions)
47}
48
49fn run_on_arena(arena: &mut nxpu_ir::Arena<Expression>) -> bool {
53 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 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 for &handle in &handles {
78 if remap.contains_key(&handle) {
79 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
91fn 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 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 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 std::ptr::from_ref(expr).hash(&mut hasher);
197 }
198 }
199
200 hasher.finish()
201}
202
203fn 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
355fn 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
370fn 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 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 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 let _ = dce_changed;
690 }
691
692 #[test]
693 fn cse_reduces_expression_count() {
694 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 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); let changed = run_on_function(&mut func);
715 assert!(changed);
716
717 let mut unique = std::collections::HashSet::new();
722 for (_, expr) in func.expressions.iter() {
723 unique.insert(format!("{expr:?}"));
724 }
725 assert!(unique.len() <= expr_count_before);
729 }
730}