use super::substitution::expr_contains_opaque;
use super::{collect_var_reads, rename_var_in_expr};
use crate::ir::{Ident, Node};
use rustc_hash::FxHashSet;
pub(super) fn collect_bound_names(nodes: &[Node], out: &mut FxHashSet<Ident>) {
for node in nodes {
match node {
Node::Let { name, .. } | Node::Assign { name, .. } => {
out.insert(name.clone());
}
Node::If {
then, otherwise, ..
} => {
collect_bound_names(then, out);
collect_bound_names(otherwise, out);
}
Node::Loop { var, body, .. } => {
out.insert(var.clone());
collect_bound_names(body, out);
}
Node::Block(body) => collect_bound_names(body, out),
Node::Region { body, .. } => collect_bound_names(body, out),
_ => {}
}
}
}
pub(super) fn bindings_flow_across(binder: &[Node], reader: &[Node], induction: &Ident) -> bool {
let mut bound: FxHashSet<Ident> = FxHashSet::default();
collect_bound_names(binder, &mut bound);
bound.remove(induction);
if bound.is_empty() {
return false;
}
let mut reads: FxHashSet<Ident> = FxHashSet::default();
collect_var_reads(reader, &mut reads);
!bound.is_disjoint(&reads)
}
pub(super) fn unsummarisable_effect(nodes: &[Node]) -> bool {
nodes.iter().any(node_unsummarisable_effect)
}
fn node_unsummarisable_effect(node: &Node) -> bool {
match node {
Node::Opaque(_) | Node::Trap { .. } | Node::Resume { .. } => true,
Node::Let { value, .. } | Node::Assign { value, .. } => expr_contains_opaque(value),
Node::Store { index, value, .. } => {
expr_contains_opaque(index) || expr_contains_opaque(value)
}
Node::If {
cond,
then,
otherwise,
} => {
expr_contains_opaque(cond)
|| unsummarisable_effect(then)
|| unsummarisable_effect(otherwise)
}
Node::Loop { from, to, body, .. } => {
expr_contains_opaque(from) || expr_contains_opaque(to) || unsummarisable_effect(body)
}
Node::Block(body) => unsummarisable_effect(body),
Node::Region { body, .. } => unsummarisable_effect(body),
Node::AsyncLoad { offset, size, .. } | Node::AsyncStore { offset, size, .. } => {
expr_contains_opaque(offset) || expr_contains_opaque(size)
}
Node::Barrier { .. }
| Node::Return
| Node::IndirectDispatch { .. }
| Node::AsyncWait { .. }
| Node::AllReduce { .. }
| Node::AllGather { .. }
| Node::ReduceScatter { .. }
| Node::Broadcast { .. } => false,
}
}
pub(super) fn rename_var_in_node(node: Node, from: &Ident, to: &Ident) -> Node {
match node {
Node::Let { name, value } => Node::Let {
name: if name == *from { to.clone() } else { name },
value: rename_var_in_expr(value, from, to),
},
Node::Assign { name, value } => Node::Assign {
name: if name == *from { to.clone() } else { name },
value: rename_var_in_expr(value, from, to),
},
Node::Store {
buffer,
index,
value,
} => Node::Store {
buffer,
index: rename_var_in_expr(index, from, to),
value: rename_var_in_expr(value, from, to),
},
Node::If {
cond,
then,
otherwise,
} => Node::If {
cond: rename_var_in_expr(cond, from, to),
then: rename_var_in_body(then, from, to),
otherwise: rename_var_in_body(otherwise, from, to),
},
Node::Loop {
var,
from: lo,
to: hi,
body,
} => Node::Loop {
var,
from: rename_var_in_expr(lo, from, to),
to: rename_var_in_expr(hi, from, to),
body: rename_var_in_body(body, from, to),
},
Node::Block(body) => Node::Block(rename_var_in_body(body, from, to)),
Node::Region {
generator,
source_region,
body,
} => {
let body_vec = std::sync::Arc::try_unwrap(body).unwrap_or_else(|arc| (*arc).clone());
Node::Region {
generator,
source_region,
body: std::sync::Arc::new(rename_var_in_body(body_vec, from, to)),
}
}
other => other,
}
}
fn rename_var_in_body(body: Vec<Node>, from: &Ident, to: &Ident) -> Vec<Node> {
body.into_iter()
.map(|n| rename_var_in_node(n, from, to))
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::optimizer::passes::loops::{loop_fission::LoopFission, loop_fusion::LoopFusion};
use crate::ir::{BufferAccess, BufferDecl, DataType, Expr, ExprNode, Node, NodeExtension, Program};
fn sorted(nodes: &[Node]) -> Vec<String> {
let mut out = FxHashSet::default();
collect_bound_names(nodes, &mut out);
let mut v: Vec<String> = out.iter().map(|n| n.as_str().to_string()).collect();
v.sort();
v
}
#[test]
fn bound_names_cover_lets_assigns_loop_vars_and_nested_scopes() {
let body = vec![
Node::let_bind("v", Expr::u32(1)),
Node::Assign {
name: Ident::from("s"),
value: Expr::u32(2),
},
Node::loop_for(
"k",
Expr::u32(0),
Expr::u32(2),
vec![Node::let_bind("inner", Expr::var("k"))],
),
Node::If {
cond: Expr::LitBool(true),
then: vec![Node::let_bind("t", Expr::u32(3))],
otherwise: vec![Node::Block(vec![Node::let_bind("e", Expr::u32(4))])],
},
];
assert_eq!(sorted(&body), ["e", "inner", "k", "s", "t", "v"]);
}
#[test]
fn bindings_flow_across_ignores_the_shared_induction_variable() {
let binder = vec![Node::let_bind("v", Expr::var("i"))];
let reads_v = vec![Node::store("b", Expr::var("i"), Expr::var("v"))];
let reads_i_only = vec![Node::store("b", Expr::var("i"), Expr::u32(1))];
let i = Ident::from("i");
assert!(bindings_flow_across(&binder, &reads_v, &i));
assert!(!bindings_flow_across(&binder, &reads_i_only, &i));
let binds_i = vec![Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(2),
vec![Node::Return],
)];
assert!(!bindings_flow_across(&binds_i, &reads_i_only, &i));
}
#[test]
fn rename_rewrites_reads_and_binding_occurrences() {
let from = Ident::from("i");
let to = Ident::from("z");
assert_eq!(
rename_var_in_node(Node::let_bind("i", Expr::var("i")), &from, &to),
Node::let_bind("z", Expr::var("z"))
);
assert_eq!(
rename_var_in_node(
Node::Assign {
name: Ident::from("i"),
value: Expr::var("i"),
},
&from,
&to
),
Node::Assign {
name: Ident::from("z"),
value: Expr::var("z"),
}
);
assert_eq!(
rename_var_in_node(
Node::Block(vec![Node::store("b", Expr::var("i"), Expr::var("keep"))]),
&from,
&to
),
Node::Block(vec![Node::store("b", Expr::var("z"), Expr::var("keep"))])
);
}
#[derive(Debug)]
struct OpaqueValue;
impl ExprNode for OpaqueValue {
fn extension_kind(&self) -> &'static str {
"test.legality.opaque_value"
}
fn debug_identity(&self) -> &str {
"opaque_value"
}
fn result_type(&self) -> Option<DataType> {
Some(DataType::U32)
}
fn cse_safe(&self) -> bool {
false
}
fn stable_fingerprint(&self) -> [u8; 32] {
[21; 32]
}
fn validate_extension(&self) -> std::result::Result<(), String> {
Ok(())
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
#[derive(Debug)]
struct OpaqueStatement;
impl NodeExtension for OpaqueStatement {
fn extension_kind(&self) -> &'static str {
"test.legality.opaque_statement"
}
fn debug_identity(&self) -> &str {
"opaque_statement"
}
fn stable_fingerprint(&self) -> [u8; 32] {
[22; 32]
}
fn validate_extension(&self) -> std::result::Result<(), String> {
Ok(())
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
#[test]
fn unsummarisable_effect_sees_through_nested_scopes() {
assert!(!unsummarisable_effect(&[Node::store(
"a",
Expr::u32(0),
Expr::u32(1)
)]));
assert!(unsummarisable_effect(&[Node::Block(vec![Node::opaque(
OpaqueStatement
)])]));
assert!(unsummarisable_effect(&[Node::If {
cond: Expr::LitBool(true),
then: vec![Node::store("a", Expr::u32(0), Expr::opaque(OpaqueValue))],
otherwise: Vec::new(),
}]));
assert!(unsummarisable_effect(&[Node::loop_for(
"k",
Expr::u32(0),
Expr::u32(2),
vec![Node::Trap {
address: Box::new(Expr::u32(0)),
tag: Ident::from("test.legality.trap"),
}],
)]));
assert!(!unsummarisable_effect(&[Node::Barrier {
ordering: crate::ir::MemoryOrdering::Relaxed,
}]));
}
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"), buf("c")],
[1, 1, 1],
entry,
)
}
fn loops_of(nodes: &[Node]) -> Vec<&Node> {
let mut out = Vec::new();
for node in nodes {
match node {
Node::Loop { .. } => out.push(node),
Node::Region { body, .. } => out.extend(loops_of(body)),
Node::Block(body) => out.extend(loops_of(body)),
_ => {}
}
}
out
}
fn shared_fixtures() -> Vec<(&'static str, bool, Vec<Node>, Vec<Node>)> {
vec![
(
"disjoint stores",
true,
vec![Node::store("a", Expr::var("i"), Expr::u32(1))],
vec![Node::store("b", Expr::var("i"), Expr::u32(2))],
),
(
"shared buffer",
false,
vec![Node::store("a", Expr::var("i"), Expr::u32(1))],
vec![Node::store("a", Expr::var("i"), Expr::u32(2))],
),
(
"name flows from first half to second",
false,
vec![Node::let_bind("v", Expr::var("i"))],
vec![Node::store("b", Expr::var("i"), Expr::var("v"))],
),
(
"nested opaque statement in the first half",
false,
vec![Node::Block(vec![Node::opaque(OpaqueStatement)])],
vec![Node::store("b", Expr::var("i"), Expr::u32(2))],
),
(
"opaque value nested in an If arm of the second half",
false,
vec![Node::store("a", Expr::var("i"), Expr::u32(1))],
vec![Node::If {
cond: Expr::LitBool(true),
then: vec![Node::store("b", Expr::var("i"), Expr::opaque(OpaqueValue))],
otherwise: Vec::new(),
}],
),
(
"independent halves each binding a private name",
true,
vec![
Node::let_bind("va", Expr::var("i")),
Node::store("a", Expr::var("i"), Expr::var("va")),
],
vec![
Node::let_bind("vb", Expr::var("i")),
Node::store("b", Expr::var("i"), Expr::var("vb")),
],
),
]
}
#[test]
fn fission_follows_the_shared_legality_verdict() {
for (label, legal, body_a, body_b) in shared_fixtures() {
let mut body = body_a;
body.extend(body_b);
let entry = vec![Node::loop_for("i", Expr::u32(0), Expr::u32(8), body)];
let result = LoopFission::transform(program(entry));
assert_eq!(result.changed, legal, "fission verdict for `{label}`");
let loops = loops_of(result.program.entry());
assert_eq!(
loops.len(),
if legal { 2 } else { 1 },
"fission loop count for `{label}`"
);
}
}
#[test]
fn fusion_follows_the_shared_legality_verdict() {
for (label, legal, body_a, body_b) in shared_fixtures() {
let i = Ident::from("i");
let j = Ident::from("j");
let body_b = rename_var_in_body(body_b, &i, &j);
let entry = vec![
Node::loop_for("i", Expr::u32(0), Expr::u32(8), body_a),
Node::loop_for("j", Expr::u32(0), Expr::u32(8), body_b),
];
let result = LoopFusion::transform(program(entry));
assert_eq!(result.changed, legal, "fusion verdict for `{label}`");
let loops = loops_of(result.program.entry());
assert_eq!(
loops.len(),
if legal { 1 } else { 2 },
"fusion loop count for `{label}`"
);
}
}
#[test]
fn fission_rekeys_the_split_off_half_onto_a_fresh_induction_variable() {
let entry = vec![Node::loop_for(
"i",
Expr::u32(0),
Expr::u32(8),
vec![
Node::store("a", Expr::var("i"), Expr::u32(1)),
Node::store("b", Expr::var("i"), Expr::u32(2)),
],
)];
let result = LoopFission::transform(program(entry));
assert!(result.changed);
let binding = result.program;
let loops = loops_of(binding.entry());
let [first, second] = loops[..] else {
panic!("Fix: fission must produce exactly two sibling loops");
};
let Node::Loop {
var: var_a,
body: body_a,
..
} = first
else {
unreachable!("filtered to Loop above");
};
let Node::Loop {
var: var_b,
body: body_b,
..
} = second
else {
unreachable!("filtered to Loop above");
};
assert_eq!(var_a.as_str(), "i");
assert_ne!(var_a, var_b, "the split-off loop needs a fresh variable");
assert_eq!(
body_a,
&vec![Node::store("a", Expr::var("i"), Expr::u32(1))]
);
assert_eq!(
body_b,
&vec![Node::store("b", Expr::var(var_b.as_str()), Expr::u32(2))],
"the split-off half must read the fresh variable, not the original"
);
}
#[test]
fn fusion_rekeys_the_second_body_onto_the_surviving_induction_variable() {
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 binding = result.program;
let loops = loops_of(binding.entry());
assert_eq!(loops.len(), 1);
let Node::Loop { var, body, .. } = loops[0] else {
unreachable!("filtered to Loop above");
};
assert_eq!(var.as_str(), "i");
assert_eq!(
body,
&vec![
Node::store("a", Expr::var("i"), Expr::u32(1)),
Node::store("b", Expr::var("i"), Expr::u32(2)),
],
"the merged half must be re-keyed onto the surviving variable"
);
}
}