use super::{collect_touched_buffers, collect_var_reads, legality};
use crate::ir::{Expr, Ident, Node, Program};
use crate::optimizer::{vyre_pass, PassAnalysis, PassResult};
use crate::visit::node_map;
use rustc_hash::FxHashSet;
#[derive(Debug, Default)]
#[vyre_pass(
name = "loop_fusion",
requires = [],
invalidates = [],
phase = "loop",
boundary_class = "abi_preserving",
cost_model_family = "loop"
)]
pub struct LoopFusion;
impl LoopFusion {
#[must_use]
fn analyze_impl(program: &Program) -> PassAnalysis {
if !program.stats().has_node_loop() {
return PassAnalysis::SKIP;
}
if body_has_fusable_pair(program.entry())
|| program
.entry()
.iter()
.any(|n| node_map::any_descendant(n, &mut has_fusable_pair))
{
PassAnalysis::RUN
} else {
PassAnalysis::SKIP
}
}
#[must_use]
pub fn transform(program: Program) -> PassResult {
let mut changed = false;
let program = program.map_entry(|entry| fuse_in_body(entry, &mut changed));
PassResult { program, changed }
}
}
fn fuse_in_body(body: Vec<Node>, changed: &mut bool) -> Vec<Node> {
let body: Vec<Node> = body.into_iter().map(|n| recurse(n, changed)).collect();
let mut out: Vec<Node> = Vec::with_capacity(body.len());
let mut iter = body.into_iter().peekable();
while let Some(node) = iter.next() {
let Node::Loop {
var: var_a,
from: from_a,
to: to_a,
body: body_a,
} = node
else {
out.push(node);
continue;
};
let next_is_fusable = matches!(iter.peek(), Some(Node::Loop { .. }));
if !next_is_fusable {
out.push(Node::Loop {
var: var_a,
from: from_a,
to: to_a,
body: body_a,
});
continue;
}
let Some(Node::Loop {
var: var_b,
from: from_b,
to: to_b,
body: body_b,
}) = iter.next()
else {
unreachable!("peek confirmed Loop above");
};
if !pair_is_fusable(
&LoopRef {
var: &var_a,
from: &from_a,
to: &to_a,
body: &body_a,
},
&LoopRef {
var: &var_b,
from: &from_b,
to: &to_b,
body: &body_b,
},
) {
out.push(Node::Loop {
var: var_a,
from: from_a,
to: to_a,
body: body_a,
});
out.push(Node::Loop {
var: var_b,
from: from_b,
to: to_b,
body: body_b,
});
continue;
}
let mut fused = body_a;
let renamed_body_b: Vec<Node> = body_b
.into_iter()
.map(|n| legality::rename_var_in_node(n, &var_b, &var_a))
.collect();
fused.extend(renamed_body_b);
*changed = true;
out.push(Node::Loop {
var: var_a,
from: from_a,
to: to_a,
body: fused,
});
}
out
}
fn recurse(node: Node, changed: &mut bool) -> Node {
let recursed = node_map::map_children(node, &mut |child| recurse(child, changed));
node_map::map_body(recursed, &mut |body| fuse_in_body(body, changed))
}
fn bounds_match(from_a: &Expr, to_a: &Expr, from_b: &Expr, to_b: &Expr) -> bool {
matches!(
(from_a, to_a, from_b, to_b),
(
Expr::LitU32(_),
Expr::LitU32(_),
Expr::LitU32(_),
Expr::LitU32(_)
)
) && from_a == from_b
&& to_a == to_b
}
struct LoopRef<'a> {
var: &'a Ident,
from: &'a Expr,
to: &'a Expr,
body: &'a [Node],
}
fn pair_is_fusable(a: &LoopRef<'_>, b: &LoopRef<'_>) -> bool {
bounds_match(a.from, a.to, b.from, b.to)
&& a.var != b.var
&& super::buffers_disjoint_with(a.body, b.body, collect_touched_buffers)
&& !legality::unsummarisable_effect(a.body)
&& !legality::unsummarisable_effect(b.body)
&& !legality::bindings_flow_across(a.body, b.body, b.var)
&& !fusion_collides_bindings(a.body, b.body, a.var, b.var)
&& !fusion_has_scalar_dependency(a.body, b.body)
}
fn fusion_collides_bindings(
body_a: &[Node],
body_b: &[Node],
var_a: &Ident,
var_b: &Ident,
) -> bool {
let mut a_lets: FxHashSet<Ident> = FxHashSet::default();
legality::collect_bound_names(body_a, &mut a_lets);
let mut b_lets: FxHashSet<Ident> = FxHashSet::default();
legality::collect_bound_names(body_b, &mut b_lets);
b_lets.iter().any(|name| {
let fused_name = if name == var_b { var_a } else { name };
fused_name == var_a || a_lets.contains(fused_name)
})
}
fn collect_assign_targets(nodes: &[Node], out: &mut FxHashSet<Ident>) {
for node in nodes {
match node {
Node::Assign { name, .. } => {
out.insert(name.clone());
}
Node::If {
then, otherwise, ..
} => {
collect_assign_targets(then, out);
collect_assign_targets(otherwise, out);
}
Node::Loop { body, .. } | Node::Block(body) => collect_assign_targets(body, out),
Node::Region { body, .. } => collect_assign_targets(body, out),
_ => {}
}
}
}
fn fusion_has_scalar_dependency(body_a: &[Node], body_b: &[Node]) -> bool {
let mut writes_a: FxHashSet<Ident> = FxHashSet::default();
collect_assign_targets(body_a, &mut writes_a);
let mut writes_b: FxHashSet<Ident> = FxHashSet::default();
collect_assign_targets(body_b, &mut writes_b);
if writes_a.is_empty() && writes_b.is_empty() {
return false; }
let mut refs_a: FxHashSet<Ident> = FxHashSet::default();
collect_var_reads(body_a, &mut refs_a);
refs_a.extend(writes_a.iter().cloned());
let mut refs_b: FxHashSet<Ident> = FxHashSet::default();
collect_var_reads(body_b, &mut refs_b);
refs_b.extend(writes_b.iter().cloned());
!writes_a.is_disjoint(&refs_b) || !writes_b.is_disjoint(&refs_a)
}
fn has_fusable_pair(node: &Node) -> bool {
let body: &[Node] = match node {
Node::If {
then, otherwise, ..
} => {
return body_has_fusable_pair(then) || body_has_fusable_pair(otherwise);
}
Node::Loop { body, .. } | Node::Block(body) => body,
Node::Region { body, .. } => body.as_ref(),
_ => return false,
};
body_has_fusable_pair(body)
}
fn body_has_fusable_pair(body: &[Node]) -> bool {
body.windows(2).any(|window| {
let (
Node::Loop {
var: var_a,
from: from_a,
to: to_a,
body: body_a,
},
Node::Loop {
var: var_b,
from: from_b,
to: to_b,
body: body_b,
},
) = (&window[0], &window[1])
else {
return false;
};
pair_is_fusable(
&LoopRef {
var: var_a,
from: from_a,
to: to_a,
body: body_a,
},
&LoopRef {
var: var_b,
from: from_b,
to: to_b,
body: body_b,
},
)
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ir::{BufferAccess, BufferDecl, DataType, Expr, ExprNode, Node, NodeExtension};
fn buf(name: &str) -> BufferDecl {
BufferDecl::storage(name, 0, BufferAccess::ReadWrite, DataType::U32).with_count(8)
}
fn program(entry: Vec<Node>) -> Program {
Program::wrapped(vec![buf("a"), buf("b")], [1, 1, 1], entry)
}
#[derive(Debug)]
struct OpaqueReader;
impl ExprNode for OpaqueReader {
fn extension_kind(&self) -> &'static str {
"test.fusion.opaque_buffer_reader"
}
fn debug_identity(&self) -> &str {
"opaque_reader"
}
fn result_type(&self) -> Option<DataType> {
Some(DataType::U32)
}
fn cse_safe(&self) -> bool {
false
}
fn stable_fingerprint(&self) -> [u8; 32] {
[13; 32]
}
fn validate_extension(&self) -> std::result::Result<(), String> {
Ok(())
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
#[derive(Debug)]
struct OpaqueWriter;
impl NodeExtension for OpaqueWriter {
fn extension_kind(&self) -> &'static str {
"test.fusion.opaque_buffer_writer"
}
fn debug_identity(&self) -> &str {
"opaque_writer"
}
fn stable_fingerprint(&self) -> [u8; 32] {
[14; 32]
}
fn validate_extension(&self) -> std::result::Result<(), String> {
Ok(())
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
#[test]
fn does_not_fuse_when_a_body_holds_an_opaque_expr() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("a", Expr::var("i"), Expr::opaque(OpaqueReader))],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("b", Expr::var("j"), Expr::u32(2))],
),
];
let result = LoopFusion::transform(program(entry));
assert!(
!result.changed,
"an opaque expression's unknowable buffer effect must block fusion"
);
assert_eq!(
count_loops(®ion_body(result.program.entry())),
2,
"loops bracketing an opaque-valued store must not fuse"
);
}
#[test]
fn does_not_fuse_when_a_body_holds_an_opaque_node() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::opaque(OpaqueWriter)],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("b", Expr::var("j"), Expr::u32(2))],
),
];
let result = LoopFusion::transform(program(entry));
assert!(
!result.changed,
"an opaque node's unknowable buffer effect must block fusion"
);
assert_eq!(
count_loops(®ion_body(result.program.entry())),
2,
"a loop whose body is an opaque node must not fuse with its sibling"
);
}
#[test]
fn does_not_fuse_when_a_shuffle_lane_loads_the_siblings_buffer() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store(
"a",
Expr::var("i"),
Expr::subgroup_shuffle(Expr::u32(5), Expr::load("b", Expr::var("i"))),
)],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("b", Expr::var("j"), Expr::u32(2))],
),
];
let result = LoopFusion::transform(program(entry));
assert!(
!result.changed,
"a buffer load hidden in a shuffle lane must block fusion"
);
assert_eq!(
count_loops(®ion_body(result.program.entry())),
2,
"loops sharing buffer `b` through a shuffle lane load must not fuse"
);
}
#[test]
fn does_not_fuse_when_a_shuffle_lane_reads_a_cross_loop_scalar() {
let entry = vec![
Node::let_bind("s", Expr::u32(0)),
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![
Node::assign("s", Expr::var("i")),
Node::store("a", Expr::var("i"), Expr::var("s")),
],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![Node::store(
"b",
Expr::var("j"),
Expr::subgroup_shuffle(Expr::u32(5), Expr::var("s")),
)],
),
];
let result = LoopFusion::transform(program(entry));
assert!(
!result.changed,
"a scalar read hidden in a shuffle lane is a cross-loop dependency; the loops must not fuse"
);
assert_eq!(
count_loops(®ion_body(result.program.entry())),
2,
"loops with a shuffle-lane scalar dependency must stay separate"
);
}
fn region_body(program_entry: &[Node]) -> Vec<Node> {
for n in program_entry {
if let Node::Region { body, .. } = n {
return body.as_ref().clone();
}
}
program_entry.to_vec()
}
fn count_loops(nodes: &[Node]) -> usize {
nodes
.iter()
.map(|n| match n {
Node::Loop { body, .. } => 1 + count_loops(body),
Node::If {
then, otherwise, ..
} => count_loops(then) + count_loops(otherwise),
Node::Block(body) => count_loops(body),
Node::Region { body, .. } => count_loops(body),
_ => 0,
})
.sum()
}
#[test]
fn fuses_two_disjoint_buffer_loops_with_matching_bounds() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("a", Expr::var("i"), Expr::u32(1))],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("b", Expr::var("j"), Expr::u32(2))],
),
];
let result = LoopFusion::transform(program(entry));
assert!(result.changed);
assert_eq!(
count_loops(®ion_body(result.program.entry())),
1,
"two loops fused into one"
);
}
#[test]
fn does_not_fuse_when_bounds_differ() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("a", Expr::var("i"), Expr::u32(1))],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(16),
vec![Node::store("b", Expr::var("j"), Expr::u32(2))],
),
];
let result = LoopFusion::transform(program(entry));
assert!(!result.changed);
}
#[test]
fn does_not_fuse_when_buffers_overlap() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("a", Expr::var("i"), Expr::u32(1))],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("a", Expr::var("j"), Expr::u32(2))],
),
];
let result = LoopFusion::transform(program(entry));
assert!(
!result.changed,
"shared buffer blocks fusion under disjoint-only conservatism"
);
}
#[test]
fn does_not_fuse_when_loop_vars_match() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("a", Expr::var("i"), Expr::u32(1))],
),
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("b", Expr::var("i"), Expr::u32(2))],
),
];
let result = LoopFusion::transform(program(entry));
assert!(!result.changed, "same loop var name blocks fusion");
}
#[test]
fn renames_second_loop_var_in_fused_body() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("a", Expr::var("i"), Expr::u32(1))],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("b", Expr::var("j"), Expr::u32(2))],
),
];
let result = LoopFusion::transform(program(entry));
assert!(result.changed);
let body = region_body(result.program.entry());
let Node::Loop { body: fused, .. } = &body[0] else {
panic!("Fix: fusion reported a change, so body[0] must be the fused Node::Loop");
};
assert_eq!(fused.len(), 2);
if let Node::Store { index, .. } = &fused[1] {
assert_eq!(
index,
&Expr::var("i"),
"second store's index must be renamed to outer var"
);
} else {
panic!("Fix: fusion concatenates body_b verbatim, so fused[1] must be body_b's Store");
}
}
#[test]
fn does_not_fuse_when_body_b_reads_a_let_bound_in_body_a() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![
Node::let_bind("tmp", Expr::u32(7)),
Node::store("a", Expr::var("i"), Expr::var("tmp")),
],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("b", Expr::var("j"), Expr::var("tmp"))],
),
];
let result = LoopFusion::transform(program(entry));
assert!(!result.changed, "shared name `tmp` blocks fusion");
}
#[test]
fn does_not_fuse_when_both_bodies_bind_same_local() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![
Node::let_bind("x", Expr::u32(1)),
Node::store("a", Expr::var("i"), Expr::var("x")),
],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![
Node::let_bind("x", Expr::u32(2)),
Node::store("b", Expr::var("j"), Expr::u32(9)),
],
),
];
let result = LoopFusion::transform(program(entry));
assert!(
!result.changed,
"a local name bound by both bodies blocks fusion (duplicate sibling)"
);
assert_eq!(
count_loops(®ion_body(result.program.entry())),
2,
"both loops must survive unfused"
);
}
#[test]
fn fuses_when_bodies_bind_distinct_locals() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![
Node::let_bind("x", Expr::u32(1)),
Node::store("a", Expr::var("i"), Expr::var("x")),
],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![
Node::let_bind("y", Expr::u32(2)),
Node::store("b", Expr::var("j"), Expr::var("y")),
],
),
];
let result = LoopFusion::transform(program(entry));
assert!(
result.changed,
"distinct local names must still allow fusion"
);
assert_eq!(
count_loops(®ion_body(result.program.entry())),
1,
"the two loops fuse into one"
);
}
#[test]
fn does_not_fuse_when_body_b_writes_a_scalar_body_a_reads() {
let entry = vec![
Node::let_bind("s", Expr::u32(0)),
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("a", Expr::var("i"), Expr::var("s"))],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![Node::assign("s", Expr::var("j"))],
),
];
let result = LoopFusion::transform(program(entry));
assert!(
!result.changed,
"a cross-loop scalar read/write dependency blocks fusion"
);
assert_eq!(
count_loops(®ion_body(result.program.entry())),
2,
"both loops survive unfused"
);
}
#[test]
fn fuses_when_bodies_mutate_independent_scalars() {
let entry = vec![
Node::let_bind("acc1", Expr::u32(0)),
Node::let_bind("acc2", Expr::u32(0)),
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![
Node::assign("acc1", Expr::add(Expr::var("acc1"), Expr::var("i"))),
Node::store("a", Expr::var("i"), Expr::var("acc1")),
],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![
Node::assign("acc2", Expr::add(Expr::var("acc2"), Expr::var("j"))),
Node::store("b", Expr::var("j"), Expr::var("acc2")),
],
),
];
let result = LoopFusion::transform(program(entry));
assert!(
result.changed,
"independent scalar accumulators must still allow fusion"
);
assert_eq!(
count_loops(®ion_body(result.program.entry())),
1,
"the two loops fuse into one"
);
}
#[test]
fn analyze_skips_when_no_fusable_pair() {
let entry = vec![Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("a", Expr::var("i"), Expr::u32(1))],
)];
assert_eq!(
crate::optimizer::ProgramPass::analyze(&LoopFusion, &program(entry)),
PassAnalysis::SKIP
);
}
#[test]
fn analyze_runs_when_fusable_pair_exists() {
let entry = vec![
Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("a", Expr::var("i"), Expr::u32(1))],
),
Node::loop_for(
"j",
Expr::u32(0),
Expr::u32(8),
vec![Node::store("b", Expr::var("j"), Expr::u32(2))],
),
];
assert_eq!(
crate::optimizer::ProgramPass::analyze(&LoopFusion, &program(entry)),
PassAnalysis::RUN
);
}
}