use std::collections::{HashMap, HashSet};
use onnx_runtime_ir::{Attribute, DataType, Graph, Node, NodeId, ValueId, WeightRef};
use crate::error::Result;
use crate::pass::{OptimizationPass, PassContext};
pub const CONTRIB_DOMAIN: &str = "com.microsoft";
const SQRT_2: f32 = std::f32::consts::SQRT_2;
const FRAC_1_SQRT_2: f32 = std::f32::consts::FRAC_1_SQRT_2;
fn approx(a: f32, expected: f32) -> bool {
(a - expected).abs() <= 1e-6 * expected.abs().max(1.0)
}
type FusedNodeSpec = (Vec<Option<ValueId>>, HashMap<String, Attribute>);
pub struct PatternMatch {
pub nodes: Vec<NodeId>,
pub external_inputs: Vec<ValueId>,
pub output: ValueId,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RewriteKind {
Structural,
LayerNorm,
Attention,
Gelu,
}
#[derive(Clone, Debug)]
pub struct FusionPattern {
name: String,
ops: Vec<String>,
replacement: String,
kind: RewriteKind,
}
impl FusionPattern {
pub fn new(name: &str, ops: &[&str], replacement: &str) -> Self {
assert!(!ops.is_empty(), "fusion pattern must have at least one op");
Self {
name: name.to_string(),
ops: ops.iter().map(|s| s.to_string()).collect(),
replacement: replacement.to_string(),
kind: RewriteKind::Structural,
}
}
pub fn layernorm() -> Self {
Self {
name: "LayerNorm".to_string(),
ops: [
"ReduceMean", "Sub", "Pow", "ReduceMean", "Add", "Sqrt", "Div", "Mul", "Add",
]
.iter()
.map(|s| s.to_string())
.collect(),
replacement: "LayerNormalization".to_string(),
kind: RewriteKind::LayerNorm,
}
}
pub fn kind(&self) -> RewriteKind {
self.kind
}
pub fn attention() -> Self {
Self {
name: "Attention".to_string(),
ops: ["Softmax"].iter().map(|s| s.to_string()).collect(),
replacement: "FusedAttention".to_string(),
kind: RewriteKind::Attention,
}
}
pub fn gelu() -> Self {
Self {
name: "Gelu".to_string(),
ops: ["Erf"].iter().map(|s| s.to_string()).collect(),
replacement: "Gelu".to_string(),
kind: RewriteKind::Gelu,
}
}
pub fn pattern_name(&self) -> &str {
&self.name
}
pub fn find_match(&self, graph: &Graph) -> Option<PatternMatch> {
for start in graph.nodes.keys() {
let m = match self.kind {
RewriteKind::LayerNorm => self.try_match_layernorm(graph, start),
RewriteKind::Attention => self.try_match_attention(graph, start),
RewriteKind::Gelu => self.try_match_gelu(graph, start),
RewriteKind::Structural => self.try_match_from(graph, start),
};
if let Some(m) = m {
return Some(m);
}
}
None
}
fn op_matches(node: &Node, op: &str) -> bool {
node.op_type == op && matches!(node.domain.as_str(), "" | "ai.onnx")
}
fn find_consumer(graph: &Graph, value: ValueId, op: &str) -> Option<NodeId> {
graph
.value(value)
.consumers
.iter()
.copied()
.find(|&c| Self::op_matches(graph.node(c), op))
}
fn try_match_layernorm(&self, graph: &Graph, start: NodeId) -> Option<PatternMatch> {
let mean_rm = graph.try_node(start)?;
if !Self::op_matches(mean_rm, "ReduceMean") || mean_rm.outputs.len() != 1 {
return None;
}
let mean = mean_rm.outputs[0];
let subs: Vec<NodeId> = graph
.value(mean)
.consumers
.iter()
.copied()
.filter(|&c| {
let n = graph.node(c);
Self::op_matches(n, "Sub") && n.input_values().any(|v| v == mean)
})
.collect();
for &sub_pow in &subs {
let sp = graph.node(sub_pow);
if sp.outputs.len() != 1 {
continue;
}
let diff_pow = sp.outputs[0];
let Some(pow) = Self::find_consumer(graph, diff_pow, "Pow") else {
continue;
};
let sq = graph.node(pow).outputs[0];
let Some(var_rm) = Self::find_consumer(graph, sq, "ReduceMean") else {
continue;
};
let var = graph.node(var_rm).outputs[0];
let Some(add_eps) = Self::find_consumer(graph, var, "Add") else {
continue;
};
let vare = graph.node(add_eps).outputs[0];
let Some(sqrt) = Self::find_consumer(graph, vare, "Sqrt") else {
continue;
};
let std = graph.node(sqrt).outputs[0];
let Some(div) = Self::find_consumer(graph, std, "Div") else {
continue;
};
let dn = graph.node(div);
let Some(num) = dn.input_values().find(|&v| v != std) else {
continue;
};
let Some(&sub_div) = subs.iter().find(|&&s| graph.node(s).outputs[0] == num) else {
continue;
};
let norm = dn.outputs[0];
let Some(mul) = Self::find_consumer(graph, norm, "Mul") else {
continue;
};
let scaled = graph.node(mul).outputs[0];
let Some(final_add) = Self::find_consumer(graph, scaled, "Add") else {
continue;
};
let mut nodes = vec![
start, sub_pow, pow, var_rm, add_eps, sqrt, div, mul, final_add,
];
if sub_div != sub_pow {
nodes.push(sub_div);
}
let matched_set: HashSet<NodeId> = nodes.iter().copied().collect();
if matched_set.len() != nodes.len() {
continue;
}
let escapes = nodes.iter().any(|&nid| {
nid != final_add
&& graph.node(nid).outputs.iter().any(|&out| {
graph.outputs.contains(&out)
|| graph
.value(out)
.consumers
.iter()
.any(|c| !matched_set.contains(c))
})
});
if escapes {
continue;
}
let fa = graph.node(final_add);
if fa.outputs.len() != 1 {
continue;
}
let output = fa.outputs[0];
let out_val = graph.value(output);
let survives = graph.outputs.contains(&output)
|| out_val.consumers.iter().any(|c| !matched_set.contains(c));
if !survives {
continue;
}
let produced: HashSet<ValueId> = nodes
.iter()
.flat_map(|&n| graph.node(n).outputs.iter().copied())
.collect();
let mut external = Vec::new();
let mut seen = HashSet::new();
for &nid in &nodes {
for iv in graph.node(nid).input_values() {
if produced.contains(&iv) {
continue;
}
if seen.insert(iv) {
external.push(iv);
}
}
}
let matched = PatternMatch {
nodes,
external_inputs: external,
output,
};
if self.layernorm_spec(graph, &matched).is_none() {
continue;
}
return Some(matched);
}
None
}
fn try_match_attention(&self, graph: &Graph, start: NodeId) -> Option<PatternMatch> {
let p = self.try_parse_attention(graph, start)?;
Some(PatternMatch {
nodes: p.nodes,
external_inputs: p.external_inputs,
output: p.output,
})
}
fn try_match_gelu(&self, graph: &Graph, start: NodeId) -> Option<PatternMatch> {
let p = self.try_parse_gelu(graph, start)?;
Some(PatternMatch {
nodes: p.nodes,
external_inputs: p.external_inputs,
output: p.output,
})
}
fn try_parse_attention(&self, graph: &Graph, sm_start: NodeId) -> Option<AttnParts> {
let sm = graph.try_node(sm_start)?;
if !Self::op_matches(sm, "Softmax") || sm.inputs.len() != 1 || sm.outputs.len() != 1 {
return None;
}
let sm_in = sm.inputs[0]?;
let sm_out = sm.outputs[0];
let rank = graph.value(sm_in).shape.len();
if rank == 0 {
return None;
}
let axis = sm.attr("axis").and_then(Attribute::as_int)?;
let axis = if axis < 0 { axis + rank as i64 } else { axis };
if axis != rank as i64 - 1 {
return None;
}
let out_mm = graph
.value(sm_out)
.consumers
.iter()
.copied()
.find(|&c| {
let n = graph.node(c);
Self::op_matches(n, "MatMul") && n.inputs.first() == Some(&Some(sm_out))
})?;
let out_mm_node = graph.node(out_mm);
if out_mm_node.inputs.len() != 2 || out_mm_node.outputs.len() != 1 {
return None;
}
let v = out_mm_node.inputs[1]?;
let output = out_mm_node.outputs[0];
let sm_in_prod = graph.value(sm_in).producer?;
let prod = graph.node(sm_in_prod);
let (scale_out, mask, mask_add) =
if Self::op_matches(prod, "Add") && prod.inputs.len() == 2 {
let a = prod.inputs[0]?;
let b = prod.inputs[1]?;
let a_scale = graph
.value(a)
.producer
.is_some_and(|p| Self::parse_scale(graph, p).is_some());
let b_scale = graph
.value(b)
.producer
.is_some_and(|p| Self::parse_scale(graph, p).is_some());
match (a_scale, b_scale) {
(true, false) => (a, Some(b), Some(sm_in_prod)),
(false, true) => (b, Some(a), Some(sm_in_prod)),
_ => return None,
}
} else {
(sm_in, None, None)
};
let scale_node_id = graph.value(scale_out).producer?;
let scale_node = graph.node(scale_node_id);
if scale_node.outputs.len() != 1 || scale_node.outputs[0] != scale_out {
return None;
}
let (scores_out, scale) = Self::parse_scale(graph, scale_node_id)?;
let score_mm_id = graph.value(scores_out).producer?;
let score_mm = graph.node(score_mm_id);
if !Self::op_matches(score_mm, "MatMul")
|| score_mm.inputs.len() != 2
|| score_mm.outputs.len() != 1
|| score_mm.outputs[0] != scores_out
{
return None;
}
let q = score_mm.inputs[0]?;
let k_side = score_mm.inputs[1]?;
let (k, k_transposed, transpose_node) = Self::attention_k(graph, k_side, score_mm_id);
let mut nodes = vec![sm_start, score_mm_id, scale_node_id, out_mm];
if let Some(ma) = mask_add {
nodes.push(ma);
}
if let Some(t) = transpose_node {
nodes.push(t);
}
let matched_set: HashSet<NodeId> = nodes.iter().copied().collect();
if matched_set.len() != nodes.len() {
return None;
}
let escapes = nodes.iter().any(|&nid| {
nid != out_mm
&& graph.node(nid).outputs.iter().any(|&o| {
graph.outputs.contains(&o)
|| graph
.value(o)
.consumers
.iter()
.any(|c| !matched_set.contains(c))
})
});
if escapes {
return None;
}
let out_val = graph.value(output);
let survives = graph.outputs.contains(&output)
|| out_val.consumers.iter().any(|c| !matched_set.contains(c));
if !survives {
return None;
}
let mut external = vec![q, k, v];
if let Some(m) = mask {
external.push(m);
}
Some(AttnParts {
nodes,
q,
k,
v,
mask,
scale,
k_transposed,
output,
external_inputs: external,
})
}
fn parse_scale(graph: &Graph, node_id: NodeId) -> Option<(ValueId, f32)> {
let n = graph.node(node_id);
if n.inputs.len() != 2 || n.outputs.len() != 1 {
return None;
}
let (scores_out, scale) = if Self::op_matches(n, "Div") {
let num = n.inputs[0]?;
let den = n.inputs[1]?;
let c = read_scalar_const_f32(graph, den)?;
if c == 0.0 {
return None;
}
(num, 1.0 / c)
} else if Self::op_matches(n, "Mul") {
let x = n.inputs[0]?;
let y = n.inputs[1]?;
match (
read_scalar_const_f32(graph, x),
read_scalar_const_f32(graph, y),
) {
(None, Some(c)) => (x, c),
(Some(c), None) => (y, c),
_ => return None,
}
} else {
return None;
};
let prod = graph.value(scores_out).producer?;
if !Self::op_matches(graph.node(prod), "MatMul") {
return None;
}
Some((scores_out, scale))
}
fn attention_k(
graph: &Graph,
k_side: ValueId,
score_mm_id: NodeId,
) -> (ValueId, bool, Option<NodeId>) {
if let Some(t_id) = graph.value(k_side).producer {
let t = graph.node(t_id);
if Self::op_matches(t, "Transpose")
&& t.inputs.len() == 1
&& t.outputs.len() == 1
&& t.outputs[0] == k_side
&& graph.value(k_side).consumers.as_slice() == [score_mm_id]
&& let Some(perm) = t.attr("perm").and_then(Attribute::as_ints)
&& is_last2_swap_perm(perm)
&& let Some(kin) = t.inputs[0]
{
return (kin, false, Some(t_id));
}
}
(k_side, true, None)
}
fn attention_spec(&self, graph: &Graph, m: &PatternMatch) -> Option<FusedNodeSpec> {
let start = *m.nodes.first()?;
let p = self.try_parse_attention(graph, start)?;
if p.nodes != m.nodes {
return None;
}
let mut inputs: Vec<Option<ValueId>> = vec![Some(p.q), Some(p.k), Some(p.v)];
if let Some(mask) = p.mask {
inputs.push(Some(mask));
}
let mut attributes = HashMap::new();
attributes.insert("scale".to_string(), Attribute::Float(p.scale));
attributes.insert(
"k_transposed".to_string(),
Attribute::Int(if p.k_transposed { 1 } else { 0 }),
);
Some((inputs, attributes))
}
fn try_parse_gelu(&self, graph: &Graph, erf_start: NodeId) -> Option<GeluParts> {
let erf = graph.try_node(erf_start)?;
if !Self::op_matches(erf, "Erf") || erf.inputs.len() != 1 || erf.outputs.len() != 1 {
return None;
}
let erf_in = erf.inputs[0]?;
let erf_out = erf.outputs[0];
let inner_id = graph.value(erf_in).producer?;
let inner = graph.node(inner_id);
if inner.outputs.first() != Some(&erf_in) {
return None;
}
let x = Self::parse_scaled(graph, inner, &[("Div", SQRT_2), ("Mul", FRAC_1_SQRT_2)])?;
let add1_id = Self::find_consumer(graph, erf_out, "Add")?;
let add1 = graph.node(add1_id);
if add1.inputs.len() != 2 || add1.outputs.len() != 1 {
return None;
}
let one = add1.input_values().find(|&v| v != erf_out)?;
if !approx(read_scalar_const_f32(graph, one)?, 1.0) {
return None;
}
let add1_out = add1.outputs[0];
let outer_id = Self::find_consumer(graph, add1_out, "Mul")?;
let outer = graph.node(outer_id);
if outer.inputs.len() != 2 || outer.outputs.len() != 1 {
return None;
}
let half = outer.input_values().find(|&v| v != add1_out)?;
let output = outer.outputs[0];
let half_id = graph.value(half).producer?;
let half_node = graph.node(half_id);
if half_node.outputs.first() != Some(&half) {
return None;
}
let x2 = Self::parse_scaled(graph, half_node, &[("Mul", 0.5), ("Div", 2.0)])?;
if x2 != x {
return None;
}
let nodes = vec![erf_start, inner_id, add1_id, outer_id, half_id];
let matched_set: HashSet<NodeId> = nodes.iter().copied().collect();
if matched_set.len() != nodes.len() {
return None;
}
let escapes = nodes.iter().any(|&nid| {
nid != outer_id
&& graph.node(nid).outputs.iter().any(|&o| {
graph.outputs.contains(&o)
|| graph
.value(o)
.consumers
.iter()
.any(|c| !matched_set.contains(c))
})
});
if escapes {
return None;
}
let out_val = graph.value(output);
let survives = graph.outputs.contains(&output)
|| out_val.consumers.iter().any(|c| !matched_set.contains(c));
if !survives {
return None;
}
Some(GeluParts {
nodes,
x,
output,
external_inputs: vec![x],
})
}
fn parse_scaled(graph: &Graph, node: &Node, forms: &[(&str, f32)]) -> Option<ValueId> {
if node.inputs.len() != 2 || node.outputs.len() != 1 {
return None;
}
let a = node.inputs[0]?;
let b = node.inputs[1]?;
for &(op, k) in forms {
if !Self::op_matches(node, op) {
continue;
}
if read_scalar_const_f32(graph, b).is_some_and(|c| approx(c, k)) {
return Some(a);
}
if op == "Mul" && read_scalar_const_f32(graph, a).is_some_and(|c| approx(c, k)) {
return Some(b);
}
}
None
}
fn gelu_spec(&self, graph: &Graph, m: &PatternMatch) -> Option<FusedNodeSpec> {
let start = *m.nodes.first()?;
let p = self.try_parse_gelu(graph, start)?;
if p.nodes != m.nodes {
return None;
}
Some((vec![Some(p.x)], HashMap::new()))
}
fn try_match_from(&self, graph: &Graph, start: NodeId) -> Option<PatternMatch> {
let start_node = graph.try_node(start)?;
if !Self::op_matches(start_node, &self.ops[0]) {
return None;
}
let mut chain = vec![start];
let mut chain_set: HashSet<NodeId> = HashSet::from([start]);
for op in &self.ops[1..] {
let prev = *chain.last().unwrap();
let mut succ_ids = graph.successors(prev);
succ_ids.sort_by_key(|n| n.0);
let next = succ_ids.into_iter().find(|&s| {
!chain_set.contains(&s) && Self::op_matches(graph.node(s), op)
})?;
chain.push(next);
chain_set.insert(next);
}
for &nid in &chain[..chain.len() - 1] {
for &out in &graph.node(nid).outputs {
if graph.outputs.contains(&out) {
return None;
}
if graph
.value(out)
.consumers
.iter()
.any(|c| !chain_set.contains(c))
{
return None;
}
}
}
let last = *chain.last().unwrap();
let last_node = graph.node(last);
if last_node.outputs.len() != 1 {
return None;
}
let output = last_node.outputs[0];
let out_val = graph.value(output);
let survives = graph.outputs.contains(&output)
|| out_val.consumers.iter().any(|c| !chain_set.contains(c));
if !survives {
return None;
}
let produced: HashSet<ValueId> = chain
.iter()
.flat_map(|&n| graph.node(n).outputs.iter().copied())
.collect();
let mut external = Vec::new();
let mut seen = HashSet::new();
for &nid in &chain {
for iv in graph.node(nid).input_values() {
if produced.contains(&iv) {
continue;
}
if seen.insert(iv) {
external.push(iv);
}
}
}
let matched = PatternMatch {
nodes: chain,
external_inputs: external,
output,
};
if !self.match_is_fusable(graph, &matched) {
return None;
}
Some(matched)
}
fn match_is_fusable(&self, graph: &Graph, m: &PatternMatch) -> bool {
match self.kind {
RewriteKind::LayerNorm => self.layernorm_spec(graph, m).is_some(),
RewriteKind::Attention => self.attention_spec(graph, m).is_some(),
RewriteKind::Gelu => self.gelu_spec(graph, m).is_some(),
RewriteKind::Structural => {
if self.replacement == "FusedMatMulBias" || self.replacement == "FusedGemm" {
self.matmul_bias_broadcast_ok(graph, m)
} else {
true
}
}
}
}
fn matmul_bias_broadcast_ok(&self, graph: &Graph, m: &PatternMatch) -> bool {
let (Some(&matmul), Some(&add)) = (m.nodes.first(), m.nodes.get(1)) else {
return false;
};
let mm_out = graph.node(matmul).outputs[0];
let Some(bias) = graph.node(add).input_values().find(|&v| v != mm_out) else {
return false;
};
let mm_shape = &graph.value(mm_out).shape;
let bias_shape = &graph.value(bias).shape;
if bias_shape.len() > mm_shape.len() {
return false;
}
let offset = mm_shape.len() - bias_shape.len();
for (i, &bdim) in bias_shape.iter().enumerate() {
let mdim = mm_shape[offset + i];
if bdim == mdim {
continue;
}
if bdim.as_static() == Some(1) {
continue;
}
return false;
}
true
}
pub fn apply_fusion(&self, graph: &mut Graph, m: &PatternMatch) -> Result<()> {
let output = m.output;
let (inputs, attributes) = match self.kind {
RewriteKind::Structural => (
m.external_inputs.iter().map(|&v| Some(v)).collect(),
HashMap::new(),
),
RewriteKind::LayerNorm => self
.layernorm_spec(graph, m)
.ok_or_else(|| crate::error::OptimizerError::Fusion(self.name.clone()))?,
RewriteKind::Attention => self
.attention_spec(graph, m)
.ok_or_else(|| crate::error::OptimizerError::Fusion(self.name.clone()))?,
RewriteKind::Gelu => self
.gelu_spec(graph, m)
.ok_or_else(|| crate::error::OptimizerError::Fusion(self.name.clone()))?,
};
for &nid in m.nodes.iter().rev() {
graph.remove_node(nid);
}
if graph.try_value(output).is_none() {
return Err(crate::error::OptimizerError::Fusion(self.name.clone()));
}
let mut fused = Node::new(NodeId(0), self.replacement.clone(), inputs, vec![output]);
fused.attributes = attributes;
fused.domain = CONTRIB_DOMAIN.to_string();
graph.insert_node(fused);
Ok(())
}
fn layernorm_spec(&self, graph: &Graph, m: &PatternMatch) -> Option<FusedNodeSpec> {
let nodes = &m.nodes;
if nodes.len() != 9 && nodes.len() != 10 {
return None;
}
let rm1 = graph.node(nodes[0]);
let sub_pow = graph.node(nodes[1]);
let pow = graph.node(nodes[2]);
let rm2 = graph.node(nodes[3]);
let add_eps = graph.node(nodes[4]);
let div = graph.node(nodes[6]);
let mul = graph.node(nodes[7]);
let final_add = graph.node(nodes[8]);
let sub_div = if nodes.len() == 10 {
graph.node(nodes[9])
} else {
sub_pow
};
let mean = rm1.outputs[0];
let diff_pow = sub_pow.outputs[0];
let diff_div = sub_div.outputs[0];
let var = rm2.outputs[0];
let norm = div.outputs[0];
let scaled = mul.outputs[0];
if !sub_pow.input_values().any(|v| v == mean)
|| !sub_div.input_values().any(|v| v == mean)
|| !pow.input_values().any(|v| v == diff_pow)
|| !div.input_values().any(|v| v == diff_div)
|| !mul.input_values().any(|v| v == norm)
|| !final_add.input_values().any(|v| v == scaled)
{
return None;
}
let x = sub_pow.input_values().find(|&v| v != mean)?;
if !sub_div.input_values().any(|v| v == x) {
return None;
}
let subtracts_x_minus_mean = |sub: &Node| -> bool {
matches!(sub.inputs.as_slice(), [Some(a), Some(b)] if *a == x && *b == mean)
};
if !subtracts_x_minus_mean(sub_pow) || !subtracts_x_minus_mean(sub_div) {
return None;
}
let scale = mul.input_values().find(|&v| v != norm)?;
let bias = final_add.input_values().find(|&v| v != scaled)?;
let eps_val = add_eps.input_values().find(|&v| v != var)?;
let epsilon = read_scalar_f32(graph, eps_val)?;
let axes = rm1.attr("axes").and_then(Attribute::as_ints)?;
let [axis] = axes else {
return None;
};
let mut attributes = HashMap::new();
attributes.insert("axis".to_string(), Attribute::Int(*axis));
attributes.insert("epsilon".to_string(), Attribute::Float(epsilon));
Some((vec![Some(x), Some(scale), Some(bias)], attributes))
}
}
fn read_scalar_f32(graph: &Graph, value: ValueId) -> Option<f32> {
match graph.initializers.get(&value)? {
WeightRef::Inline(t) if t.dtype == DataType::Float32 && t.data.len() >= 4 => {
Some(f32::from_le_bytes(t.data[0..4].try_into().ok()?))
}
_ => None,
}
}
#[derive(Clone, Debug)]
struct AttnParts {
nodes: Vec<NodeId>,
q: ValueId,
k: ValueId,
v: ValueId,
mask: Option<ValueId>,
scale: f32,
k_transposed: bool,
output: ValueId,
external_inputs: Vec<ValueId>,
}
#[derive(Clone, Debug)]
struct GeluParts {
nodes: Vec<NodeId>,
x: ValueId,
output: ValueId,
external_inputs: Vec<ValueId>,
}
fn read_scalar_const_f32(graph: &Graph, value: ValueId) -> Option<f32> {
match graph.initializers.get(&value)? {
WeightRef::Inline(t) if t.dtype == DataType::Float32 => {
let numel: usize = t.dims.iter().product();
if numel != 1 || t.data.len() < 4 {
return None;
}
Some(f32::from_le_bytes(t.data[0..4].try_into().ok()?))
}
_ => None,
}
}
fn is_last2_swap_perm(perm: &[i64]) -> bool {
let r = perm.len();
if r < 2 {
return false;
}
for (i, &p) in perm.iter().enumerate().take(r - 2) {
if p != i as i64 {
return false;
}
}
perm[r - 2] == (r - 1) as i64 && perm[r - 1] == (r - 2) as i64
}
pub fn default_fusion_patterns() -> Vec<FusionPattern> {
vec![
FusionPattern::attention(),
FusionPattern::new("MatMul+Bias+Relu", &["MatMul", "Add", "Relu"], "FusedGemm"),
FusionPattern::layernorm(),
FusionPattern::gelu(),
FusionPattern::new("MatMul+Bias", &["MatMul", "Add"], "FusedMatMulBias"),
]
}
#[derive(Clone, Debug)]
pub struct OpFusion {
patterns: Vec<FusionPattern>,
}
impl Default for OpFusion {
fn default() -> Self {
Self::new()
}
}
impl OpFusion {
pub fn new() -> Self {
Self {
patterns: default_fusion_patterns(),
}
}
pub fn with_patterns(patterns: Vec<FusionPattern>) -> Self {
Self { patterns }
}
}
impl OptimizationPass for OpFusion {
fn name(&self) -> &str {
"OpFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> Result<()> {
for pattern in &self.patterns {
while let Some(m) = pattern.find_match(graph) {
pattern.apply_fusion(graph, &m)?;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use onnx_runtime_ir::{DataType, Node, NodeId, TensorData, static_shape};
fn val(g: &mut Graph, name: &str) -> ValueId {
g.create_named_value(name, DataType::Float32, static_shape([4]))
}
fn matmul_add_graph() -> Graph {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let a = val(&mut g, "a");
let w = val(&mut g, "w");
let bias = val(&mut g, "bias");
g.add_input(a);
g.add_input(w);
g.add_input(bias);
let m = val(&mut g, "m");
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(a), Some(w)], vec![m]));
let out = val(&mut g, "out");
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(m), Some(bias)], vec![out]));
g.add_output(out);
g
}
#[test]
fn fuses_matmul_add() {
let mut g = matmul_add_graph();
assert_eq!(g.num_nodes(), 2);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 1);
let fused = g.nodes.values().next().unwrap();
assert_eq!(fused.op_type, "FusedMatMulBias");
assert_eq!(fused.domain, CONTRIB_DOMAIN);
assert_eq!(fused.inputs.len(), 3);
assert!(g.validate().is_ok());
assert_eq!(g.outputs.len(), 1);
assert_eq!(fused.outputs, g.outputs);
}
#[test]
fn fuses_matmul_add_relu_before_matmul_add() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let a = val(&mut g, "a");
let w = val(&mut g, "w");
let bias = val(&mut g, "bias");
g.add_input(a);
g.add_input(w);
g.add_input(bias);
let m = val(&mut g, "m");
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(a), Some(w)], vec![m]));
let s = val(&mut g, "s");
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(m), Some(bias)], vec![s]));
let out = val(&mut g, "out");
g.insert_node(Node::new(NodeId(0), "Relu", vec![Some(s)], vec![out]));
g.add_output(out);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 1);
let fused = g.nodes.values().next().unwrap();
assert_eq!(fused.op_type, "FusedGemm");
assert_eq!(fused.domain, CONTRIB_DOMAIN);
assert!(g.validate().is_ok());
}
#[test]
fn does_not_fuse_when_intermediate_has_second_consumer() {
let mut g = matmul_add_graph();
let m = g
.values
.iter()
.find(|(_, v)| v.name.as_deref() == Some("m"))
.map(|(id, _)| id)
.unwrap();
let side = val(&mut g, "side");
g.insert_node(Node::new(NodeId(0), "Relu", vec![Some(m)], vec![side]));
g.add_output(side);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert!(
g.nodes.values().any(|n| n.op_type == "MatMul"),
"MatMul must remain — its output has a second consumer"
);
assert!(g.nodes.values().all(|n| n.op_type != "FusedMatMulBias"));
assert!(g.validate().is_ok());
}
#[test]
fn no_match_returns_none() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let a = val(&mut g, "a");
g.add_input(a);
let out = val(&mut g, "out");
g.insert_node(Node::new(NodeId(0), "Relu", vec![Some(a)], vec![out]));
g.add_output(out);
let p = FusionPattern::new("MatMul+Bias", &["MatMul", "Add"], "FusedMatMulBias");
assert!(p.find_match(&g).is_none());
}
fn layernorm_graph() -> Graph {
const EPS: f32 = 1e-12;
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = val(&mut g, "x");
let two = val(&mut g, "two");
let eps = val(&mut g, "eps");
let scale = val(&mut g, "scale");
let bias = val(&mut g, "bias");
g.add_input(x);
g.add_input(two);
g.set_initializer(
eps,
WeightRef::Inline(TensorData::from_raw(
DataType::Float32,
vec![],
EPS.to_le_bytes().to_vec(),
)),
);
g.add_input(scale);
g.add_input(bias);
let reduce_mean = |g: &mut Graph, input: ValueId, out: ValueId| {
let mut n = Node::new(NodeId(0), "ReduceMean", vec![Some(input)], vec![out]);
n.attributes.insert("axes".into(), Attribute::Ints(vec![-1]));
n.attributes.insert("keepdims".into(), Attribute::Int(1));
g.insert_node(n);
};
let mean = val(&mut g, "mean");
reduce_mean(&mut g, x, mean);
let diff = val(&mut g, "diff");
g.insert_node(Node::new(NodeId(0), "Sub", vec![Some(x), Some(mean)], vec![diff]));
let sq = val(&mut g, "sq");
g.insert_node(Node::new(NodeId(0), "Pow", vec![Some(diff), Some(two)], vec![sq]));
let var = val(&mut g, "var");
reduce_mean(&mut g, sq, var);
let vare = val(&mut g, "vare");
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(var), Some(eps)], vec![vare]));
let std = val(&mut g, "std");
g.insert_node(Node::new(NodeId(0), "Sqrt", vec![Some(vare)], vec![std]));
let norm = val(&mut g, "norm");
g.insert_node(Node::new(NodeId(0), "Div", vec![Some(diff), Some(std)], vec![norm]));
let scaled = val(&mut g, "scaled");
g.insert_node(Node::new(
NodeId(0),
"Mul",
vec![Some(norm), Some(scale)],
vec![scaled],
));
let out = val(&mut g, "out");
g.insert_node(Node::new(
NodeId(0),
"Add",
vec![Some(scaled), Some(bias)],
vec![out],
));
g.add_output(out);
g
}
#[test]
fn fuses_layernorm_chain() {
let mut g = layernorm_graph();
assert_eq!(g.num_nodes(), 9);
assert!(g.validate().is_ok());
let vid = |name: &str| {
g.values
.iter()
.find(|(_, v)| v.name.as_deref() == Some(name))
.map(|(id, _)| id)
.unwrap()
};
let x = vid("x");
let scale = vid("scale");
let bias = vid("bias");
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 1, "9-op chain collapses to one node");
let fused = g.nodes.values().next().unwrap();
assert_eq!(fused.op_type, "LayerNormalization");
assert_eq!(fused.domain, CONTRIB_DOMAIN);
assert_eq!(fused.inputs, vec![Some(x), Some(scale), Some(bias)]);
assert_eq!(
fused.attr("axis").and_then(Attribute::as_int),
Some(-1),
"axis extracted from ReduceMean axes"
);
let eps = fused
.attr("epsilon")
.and_then(Attribute::as_float)
.expect("epsilon attribute present");
assert!(
(eps - 1e-12).abs() < 1e-18,
"epsilon extracted from the var+eps constant, got {eps}"
);
assert_eq!(fused.outputs, g.outputs);
assert!(g.validate().is_ok());
}
#[test]
fn layernorm_count_bookkeeping() {
let mut g = layernorm_graph();
let ln_before = g.nodes.values().filter(|n| n.op_type == "LayerNormalization").count();
let rm_before = g.nodes.values().filter(|n| n.op_type == "ReduceMean").count();
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
let ln_after = g.nodes.values().filter(|n| n.op_type == "LayerNormalization").count();
let rm_after = g.nodes.values().filter(|n| n.op_type == "ReduceMean").count();
assert_eq!(ln_before, 0);
assert_eq!(ln_after, 1);
assert_eq!(rm_before, 2);
assert_eq!(rm_after, 0);
}
fn layernorm_split_graph(reverse_num_sub: bool) -> Graph {
const EPS: f32 = 1e-12;
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = val(&mut g, "x");
let two = val(&mut g, "two");
let eps = val(&mut g, "eps");
let scale = val(&mut g, "scale");
let bias = val(&mut g, "bias");
g.add_input(x);
g.add_input(two);
g.set_initializer(
eps,
WeightRef::Inline(TensorData::from_raw(
DataType::Float32,
vec![],
EPS.to_le_bytes().to_vec(),
)),
);
g.add_input(scale);
g.add_input(bias);
let reduce_mean = |g: &mut Graph, input: ValueId, out: ValueId| {
let mut n = Node::new(NodeId(0), "ReduceMean", vec![Some(input)], vec![out]);
n.attributes.insert("axes".into(), Attribute::Ints(vec![-1]));
n.attributes.insert("keepdims".into(), Attribute::Int(1));
g.insert_node(n);
};
let mean = val(&mut g, "mean");
reduce_mean(&mut g, x, mean);
let diff_pow = val(&mut g, "diff_pow");
g.insert_node(Node::new(
NodeId(0),
"Sub",
vec![Some(x), Some(mean)],
vec![diff_pow],
));
let diff_div = val(&mut g, "diff_div");
let num_inputs = if reverse_num_sub {
vec![Some(mean), Some(x)]
} else {
vec![Some(x), Some(mean)]
};
g.insert_node(Node::new(NodeId(0), "Sub", num_inputs, vec![diff_div]));
let sq = val(&mut g, "sq");
g.insert_node(Node::new(
NodeId(0),
"Pow",
vec![Some(diff_pow), Some(two)],
vec![sq],
));
let var = val(&mut g, "var");
reduce_mean(&mut g, sq, var);
let vare = val(&mut g, "vare");
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(var), Some(eps)], vec![vare]));
let std = val(&mut g, "std");
g.insert_node(Node::new(NodeId(0), "Sqrt", vec![Some(vare)], vec![std]));
let norm = val(&mut g, "norm");
g.insert_node(Node::new(
NodeId(0),
"Div",
vec![Some(diff_div), Some(std)],
vec![norm],
));
let scaled = val(&mut g, "scaled");
g.insert_node(Node::new(
NodeId(0),
"Mul",
vec![Some(norm), Some(scale)],
vec![scaled],
));
let out = val(&mut g, "out");
g.insert_node(Node::new(
NodeId(0),
"Add",
vec![Some(scaled), Some(bias)],
vec![out],
));
g.add_output(out);
g
}
#[test]
fn fuses_layernorm_split_chain() {
let mut g = layernorm_split_graph(false);
assert_eq!(g.num_nodes(), 10, "split-diff shape has two distinct Subs");
assert!(g.validate().is_ok());
let vid = |name: &str| {
g.values
.iter()
.find(|(_, v)| v.name.as_deref() == Some(name))
.map(|(id, _)| id)
.unwrap()
};
let x = vid("x");
let scale = vid("scale");
let bias = vid("bias");
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 1, "10-op split chain collapses to one node");
let fused = g.nodes.values().next().unwrap();
assert_eq!(fused.op_type, "LayerNormalization");
assert_eq!(fused.domain, CONTRIB_DOMAIN);
assert_eq!(fused.inputs, vec![Some(x), Some(scale), Some(bias)]);
assert_eq!(
fused.attr("axis").and_then(Attribute::as_int),
Some(-1),
"axis extracted from ReduceMean axes"
);
let eps = fused
.attr("epsilon")
.and_then(Attribute::as_float)
.expect("epsilon attribute present");
assert!(
(eps - 1e-12).abs() < 1e-18,
"epsilon extracted from the var+eps constant, got {eps}"
);
assert_eq!(fused.outputs, g.outputs);
assert!(g.validate().is_ok());
}
#[test]
fn declines_layernorm_when_numerator_sub_reversed() {
let mut g = layernorm_split_graph(true);
assert_eq!(g.num_nodes(), 10);
assert!(g.validate().is_ok());
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert!(
g.nodes.values().all(|n| n.op_type != "LayerNormalization"),
"reversed Sub(mean, x) must NOT fuse — sign-flip over-match"
);
assert_eq!(g.num_nodes(), 10, "all 10 ops remain (declined)");
assert_eq!(
g.nodes.values().filter(|n| n.op_type == "Sub").count(),
2,
"both centering Subs preserved"
);
assert!(g.validate().is_ok());
}
#[test]
fn does_not_fuse_partial_layernorm() {
let mut g = layernorm_graph();
let p = FusionPattern::layernorm();
let diff = g
.values
.iter()
.find(|(_, v)| v.name.as_deref() == Some("diff"))
.map(|(id, _)| id)
.unwrap();
let side = val(&mut g, "side");
g.insert_node(Node::new(NodeId(0), "Neg", vec![Some(diff)], vec![side]));
g.add_output(side);
assert!(
p.find_match(&g).is_none(),
"external consumer on `diff` blocks the fusion"
);
}
#[test]
fn fuses_two_independent_matmul_adds() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
for i in 0..2 {
let a = val(&mut g, &format!("a{i}"));
let w = val(&mut g, &format!("w{i}"));
let bias = val(&mut g, &format!("bias{i}"));
g.add_input(a);
g.add_input(w);
g.add_input(bias);
let m = val(&mut g, &format!("m{i}"));
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(a), Some(w)], vec![m]));
let out = val(&mut g, &format!("out{i}"));
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(m), Some(bias)], vec![out]));
g.add_output(out);
}
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 2);
assert!(g.nodes.values().all(|n| n.op_type == "FusedMatMulBias"));
assert!(g.validate().is_ok());
}
#[test]
fn find_match_reports_correct_shape() {
let g = matmul_add_graph();
let p = FusionPattern::new("MatMul+Bias", &["MatMul", "Add"], "FusedMatMulBias");
let m = p.find_match(&g).expect("should match");
assert_eq!(m.nodes.len(), 2);
assert_eq!(m.external_inputs.len(), 3);
assert_eq!(p.pattern_name(), "MatMul+Bias");
}
#[test]
fn declines_layernorm_when_axes_is_input() {
let mut g = layernorm_graph();
let mean = g
.values
.iter()
.find(|(_, v)| v.name.as_deref() == Some("mean"))
.map(|(id, _)| id)
.unwrap();
let rm1 = g.value(mean).producer.unwrap();
g.node_mut(rm1).attributes.remove("axes");
let axes_in = g.create_named_value("axes_in", DataType::Int64, static_shape([1]));
g.set_initializer(
axes_in,
WeightRef::Inline(TensorData::from_raw(
DataType::Int64,
vec![1],
(-1i64).to_le_bytes().to_vec(),
)),
);
g.node_mut(rm1).inputs.push(Some(axes_in));
g.value_mut(axes_in).consumers.push(rm1);
assert!(g.validate().is_ok());
assert_eq!(g.num_nodes(), 9);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(
g.num_nodes(),
9,
"axes-as-input LayerNorm must NOT fuse — all original ops kept"
);
assert!(
g.nodes.values().all(|n| n.op_type != "LayerNormalization"),
"no fused LayerNormalization must be emitted"
);
assert_eq!(
g.nodes.values().filter(|n| n.op_type == "ReduceMean").count(),
2,
"both ReduceMean ops remain"
);
assert!(g.validate().is_ok());
}
#[test]
fn declines_layernorm_when_epsilon_not_constant() {
let mut g = layernorm_graph();
let eps = g
.values
.iter()
.find(|(_, v)| v.name.as_deref() == Some("eps"))
.map(|(id, _)| id)
.unwrap();
g.initializers.remove(&eps);
g.add_input(eps);
assert!(g.validate().is_ok());
assert_eq!(g.num_nodes(), 9);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(
g.num_nodes(),
9,
"non-constant epsilon LayerNorm must NOT fuse"
);
assert!(g.nodes.values().all(|n| n.op_type != "LayerNormalization"));
assert!(g.validate().is_ok());
}
#[test]
fn declines_matmul_add_when_bias_expands() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let a = g.create_named_value("a", DataType::Float32, static_shape([4, 4]));
let w = g.create_named_value("w", DataType::Float32, static_shape([4]));
let bias = g.create_named_value("bias", DataType::Float32, static_shape([2, 4]));
g.add_input(a);
g.add_input(w);
g.add_input(bias);
let m = g.create_named_value("m", DataType::Float32, static_shape([4]));
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(a), Some(w)], vec![m]));
let out = g.create_named_value("out", DataType::Float32, static_shape([2, 4]));
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(m), Some(bias)], vec![out]));
g.add_output(out);
assert_eq!(g.num_nodes(), 2);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 2, "expanding bias must NOT fuse");
assert!(g.nodes.values().any(|n| n.op_type == "MatMul"));
assert!(g.nodes.values().any(|n| n.op_type == "Add"));
assert!(g.nodes.values().all(|n| n.op_type != "FusedMatMulBias"));
assert!(g.validate().is_ok());
}
#[test]
fn fuses_matmul_add_with_trailing_broadcast_bias() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let a = g.create_named_value("a", DataType::Float32, static_shape([3, 4]));
let w = g.create_named_value("w", DataType::Float32, static_shape([4, 4]));
let bias = g.create_named_value("bias", DataType::Float32, static_shape([1, 4]));
g.add_input(a);
g.add_input(w);
g.add_input(bias);
let m = g.create_named_value("m", DataType::Float32, static_shape([3, 4]));
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(a), Some(w)], vec![m]));
let out = g.create_named_value("out", DataType::Float32, static_shape([3, 4]));
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(m), Some(bias)], vec![out]));
g.add_output(out);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 1, "trailing-broadcast bias must fuse");
assert_eq!(
g.nodes.values().next().unwrap().op_type,
"FusedMatMulBias"
);
assert!(g.validate().is_ok());
}
#[test]
fn declines_matmul_add_when_shape_unknown() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let a = g.create_named_value("a", DataType::Float32, Vec::new());
let w = g.create_named_value("w", DataType::Float32, Vec::new());
let bias = g.create_named_value("bias", DataType::Float32, static_shape([4]));
g.add_input(a);
g.add_input(w);
g.add_input(bias);
let m = g.create_named_value("m", DataType::Float32, Vec::new());
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(a), Some(w)], vec![m]));
let out = g.create_named_value("out", DataType::Float32, static_shape([4]));
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(m), Some(bias)], vec![out]));
g.add_output(out);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 2, "unknown matmul shape must NOT fuse");
assert!(g.nodes.values().all(|n| n.op_type != "FusedMatMulBias"));
assert!(g.validate().is_ok());
}
#[test]
fn declines_fused_gemm_when_bias_expands() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let a = g.create_named_value("a", DataType::Float32, static_shape([4, 4]));
let w = g.create_named_value("w", DataType::Float32, static_shape([4]));
let bias = g.create_named_value("bias", DataType::Float32, static_shape([2, 4]));
g.add_input(a);
g.add_input(w);
g.add_input(bias);
let m = g.create_named_value("m", DataType::Float32, static_shape([4]));
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(a), Some(w)], vec![m]));
let biased = g.create_named_value("biased", DataType::Float32, static_shape([2, 4]));
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(m), Some(bias)], vec![biased]));
let out = g.create_named_value("out", DataType::Float32, static_shape([2, 4]));
g.insert_node(Node::new(NodeId(0), "Relu", vec![Some(biased)], vec![out]));
g.add_output(out);
assert_eq!(g.num_nodes(), 3);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 3, "expanding bias must NOT fuse to FusedGemm");
assert!(g.nodes.values().any(|n| n.op_type == "MatMul"));
assert!(g.nodes.values().any(|n| n.op_type == "Add"));
assert!(g.nodes.values().any(|n| n.op_type == "Relu"));
assert!(g.nodes.values().all(|n| n.op_type != "FusedGemm"));
assert!(g.validate().is_ok());
}
fn scalar_init(g: &mut Graph, name: &str, v: f32) -> ValueId {
let vid = g.create_named_value(name, DataType::Float32, Vec::new());
g.set_initializer(
vid,
WeightRef::Inline(TensorData::from_raw(
DataType::Float32,
vec![],
v.to_le_bytes().to_vec(),
)),
);
vid
}
fn fval(g: &mut Graph, name: &str, dims: &[usize]) -> ValueId {
g.create_named_value(name, DataType::Float32, static_shape(dims.iter().copied()))
}
fn value_id(g: &Graph, name: &str) -> ValueId {
g.values
.iter()
.find(|(_, v)| v.name.as_deref() == Some(name))
.map(|(id, _)| id)
.unwrap_or_else(|| panic!("no value named {name}"))
}
fn sdpa_graph(masked: bool, axis: i64) -> Graph {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 12);
let q = fval(&mut g, "Q", &[1, 2, 3, 4]);
let kt = fval(&mut g, "K", &[1, 2, 4, 3]); let v = fval(&mut g, "V", &[1, 2, 3, 4]);
g.add_input(q);
g.add_input(kt);
g.add_input(v);
let c = scalar_init(&mut g, "scale_c", 2.0);
let scores = fval(&mut g, "scores", &[1, 2, 3, 3]);
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(q), Some(kt)], vec![scores]));
let scaled = fval(&mut g, "scaled", &[1, 2, 3, 3]);
g.insert_node(Node::new(NodeId(0), "Div", vec![Some(scores), Some(c)], vec![scaled]));
let sm_in = if masked {
let mask = fval(&mut g, "mask", &[1, 1, 3, 3]);
g.add_input(mask);
let masked_v = fval(&mut g, "masked", &[1, 2, 3, 3]);
g.insert_node(Node::new(
NodeId(0),
"Add",
vec![Some(scaled), Some(mask)],
vec![masked_v],
));
masked_v
} else {
scaled
};
let probs = fval(&mut g, "probs", &[1, 2, 3, 3]);
let mut sm = Node::new(NodeId(0), "Softmax", vec![Some(sm_in)], vec![probs]);
sm.attributes.insert("axis".into(), Attribute::Int(axis));
g.insert_node(sm);
let out = fval(&mut g, "out", &[1, 2, 3, 4]);
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(probs), Some(v)], vec![out]));
g.add_output(out);
g
}
fn fused_attention_node(g: &Graph) -> Option<&Node> {
g.nodes.values().find(|n| n.op_type == "FusedAttention")
}
#[test]
fn fuses_sdpa_unmasked_pretransposed_k() {
let mut g = sdpa_graph(false, 3);
let q = value_id(&g, "Q");
let k = value_id(&g, "K");
let v = value_id(&g, "V");
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(
g.nodes.values().filter(|n| n.op_type == "FusedAttention").count(),
1
);
assert!(g.nodes.values().all(|n| n.op_type != "Softmax"));
assert!(g.nodes.values().all(|n| n.op_type != "MatMul"));
assert!(g.nodes.values().all(|n| n.op_type != "Div"));
let fa = fused_attention_node(&g).unwrap();
assert_eq!(fa.domain, CONTRIB_DOMAIN);
assert_eq!(fa.inputs, vec![Some(q), Some(k), Some(v)]);
assert_eq!(fa.attr("scale").and_then(Attribute::as_float), Some(0.5));
assert_eq!(fa.attr("k_transposed").and_then(Attribute::as_int), Some(1));
assert!(g.validate().is_ok());
}
#[test]
fn fuses_sdpa_masked() {
let mut g = sdpa_graph(true, 3);
let (q, k, v, mask) = (
value_id(&g, "Q"),
value_id(&g, "K"),
value_id(&g, "V"),
value_id(&g, "mask"),
);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(
g.nodes.values().filter(|n| n.op_type == "FusedAttention").count(),
1
);
assert!(g.nodes.values().all(|n| n.op_type != "Softmax"));
assert!(g.nodes.values().all(|n| n.op_type != "Add"));
let fa = fused_attention_node(&g).unwrap();
assert_eq!(fa.inputs, vec![Some(q), Some(k), Some(v), Some(mask)]);
assert_eq!(fa.attr("k_transposed").and_then(Attribute::as_int), Some(1));
assert!(g.validate().is_ok());
}
#[test]
fn fuses_sdpa_absorbing_clean_transpose_sets_k_transposed_0() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 12);
let q = fval(&mut g, "Q", &[1, 2, 3, 4]);
let k = fval(&mut g, "K", &[1, 2, 3, 4]); let v = fval(&mut g, "V", &[1, 2, 3, 4]);
g.add_input(q);
g.add_input(k);
g.add_input(v);
let c = scalar_init(&mut g, "scale_c", 4.0);
let kt = fval(&mut g, "Kt", &[1, 2, 4, 3]);
let mut tr = Node::new(NodeId(0), "Transpose", vec![Some(k)], vec![kt]);
tr.attributes.insert("perm".into(), Attribute::Ints(vec![0, 1, 3, 2]));
g.insert_node(tr);
let scores = fval(&mut g, "scores", &[1, 2, 3, 3]);
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(q), Some(kt)], vec![scores]));
let scaled = fval(&mut g, "scaled", &[1, 2, 3, 3]);
g.insert_node(Node::new(NodeId(0), "Div", vec![Some(scores), Some(c)], vec![scaled]));
let probs = fval(&mut g, "probs", &[1, 2, 3, 3]);
let mut sm = Node::new(NodeId(0), "Softmax", vec![Some(scaled)], vec![probs]);
sm.attributes.insert("axis".into(), Attribute::Int(-1));
g.insert_node(sm);
let out = fval(&mut g, "out", &[1, 2, 3, 4]);
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(probs), Some(v)], vec![out]));
g.add_output(out);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert!(g.nodes.values().all(|n| n.op_type != "Transpose"), "clean Kᵀ Transpose absorbed");
let fa = fused_attention_node(&g).unwrap();
assert_eq!(fa.inputs, vec![Some(q), Some(k), Some(v)], "K input is the natural (un-transposed) K");
assert_eq!(fa.attr("k_transposed").and_then(Attribute::as_int), Some(0));
assert_eq!(fa.attr("scale").and_then(Attribute::as_float), Some(0.25));
assert!(g.validate().is_ok());
}
#[test]
fn declines_sdpa_when_softmax_axis_not_last() {
let mut g = sdpa_graph(false, 1);
let before = g.num_nodes();
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert!(g.nodes.values().all(|n| n.op_type != "FusedAttention"));
assert!(g.nodes.values().any(|n| n.op_type == "Softmax"));
assert_eq!(g.num_nodes(), before, "no fusion when axis is not last");
}
#[test]
fn declines_sdpa_when_scale_is_not_scalar_constant() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 12);
let q = fval(&mut g, "Q", &[1, 2, 3, 4]);
let kt = fval(&mut g, "K", &[1, 2, 4, 3]);
let v = fval(&mut g, "V", &[1, 2, 3, 4]);
let c = fval(&mut g, "scale_c", &[]); g.add_input(q);
g.add_input(kt);
g.add_input(v);
g.add_input(c);
let scores = fval(&mut g, "scores", &[1, 2, 3, 3]);
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(q), Some(kt)], vec![scores]));
let scaled = fval(&mut g, "scaled", &[1, 2, 3, 3]);
g.insert_node(Node::new(NodeId(0), "Div", vec![Some(scores), Some(c)], vec![scaled]));
let probs = fval(&mut g, "probs", &[1, 2, 3, 3]);
let mut sm = Node::new(NodeId(0), "Softmax", vec![Some(scaled)], vec![probs]);
sm.attributes.insert("axis".into(), Attribute::Int(3));
g.insert_node(sm);
let out = fval(&mut g, "out", &[1, 2, 3, 4]);
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(probs), Some(v)], vec![out]));
g.add_output(out);
let before = g.num_nodes();
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert!(g.nodes.values().all(|n| n.op_type != "FusedAttention"));
assert_eq!(g.num_nodes(), before, "non-constant scale must NOT fuse");
}
#[test]
fn declines_sdpa_when_softmax_is_right_operand_of_output_matmul() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 12);
let q = fval(&mut g, "Q", &[1, 2, 3, 4]);
let kt = fval(&mut g, "K", &[1, 2, 4, 3]);
let v = fval(&mut g, "V", &[1, 2, 3, 3]);
g.add_input(q);
g.add_input(kt);
g.add_input(v);
let c = scalar_init(&mut g, "scale_c", 2.0);
let scores = fval(&mut g, "scores", &[1, 2, 3, 3]);
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(q), Some(kt)], vec![scores]));
let scaled = fval(&mut g, "scaled", &[1, 2, 3, 3]);
g.insert_node(Node::new(NodeId(0), "Div", vec![Some(scores), Some(c)], vec![scaled]));
let probs = fval(&mut g, "probs", &[1, 2, 3, 3]);
let mut sm = Node::new(NodeId(0), "Softmax", vec![Some(scaled)], vec![probs]);
sm.attributes.insert("axis".into(), Attribute::Int(3));
g.insert_node(sm);
let out = fval(&mut g, "out", &[1, 2, 3, 3]);
g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(v), Some(probs)], vec![out]));
g.add_output(out);
let before = g.num_nodes();
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert!(g.nodes.values().all(|n| n.op_type != "FusedAttention"));
assert!(g.nodes.values().any(|n| n.op_type == "Softmax"));
assert_eq!(g.num_nodes(), before);
}
fn gelu_graph(inner_div_sqrt2: bool, half_mul: bool) -> Graph {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = val(&mut g, "x");
g.add_input(x);
let half = val(&mut g, "half");
if half_mul {
let c = scalar_init(&mut g, "c_half", 0.5);
g.insert_node(Node::new(NodeId(0), "Mul", vec![Some(x), Some(c)], vec![half]));
} else {
let c = scalar_init(&mut g, "c_two", 2.0);
g.insert_node(Node::new(NodeId(0), "Div", vec![Some(x), Some(c)], vec![half]));
}
let scaled = val(&mut g, "scaled");
if inner_div_sqrt2 {
let c = scalar_init(&mut g, "c_sqrt2", std::f32::consts::SQRT_2);
g.insert_node(Node::new(NodeId(0), "Div", vec![Some(x), Some(c)], vec![scaled]));
} else {
let c = scalar_init(&mut g, "c_isqrt2", std::f32::consts::FRAC_1_SQRT_2);
g.insert_node(Node::new(NodeId(0), "Mul", vec![Some(x), Some(c)], vec![scaled]));
}
let e = val(&mut g, "e");
g.insert_node(Node::new(NodeId(0), "Erf", vec![Some(scaled)], vec![e]));
let one = scalar_init(&mut g, "c_one", 1.0);
let a = val(&mut g, "a");
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(e), Some(one)], vec![a]));
let out = val(&mut g, "out");
g.insert_node(Node::new(NodeId(0), "Mul", vec![Some(half), Some(a)], vec![out]));
g.add_output(out);
g
}
#[test]
fn fuses_gelu_div_sqrt2() {
let mut g = gelu_graph(true, true);
assert_eq!(g.num_nodes(), 5);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
let gelu: Vec<_> = g.nodes.values().filter(|n| n.op_type == "Gelu").collect();
assert_eq!(gelu.len(), 1, "the Erf decomposition must fuse to one Gelu");
let fused = gelu[0];
assert_eq!(fused.domain, CONTRIB_DOMAIN);
assert_eq!(fused.inputs.len(), 1, "Gelu takes the single input x");
assert!(fused.attributes.is_empty(), "exact Gelu has no attributes");
let x = g.values.iter().find(|(_, v)| v.name.as_deref() == Some("x")).map(|(id, _)| id).unwrap();
assert_eq!(fused.inputs[0], Some(x));
assert_eq!(fused.outputs, g.outputs);
assert!(g.nodes.values().all(|n| n.op_type != "Erf"));
assert!(g.validate().is_ok());
}
#[test]
fn fuses_gelu_mul_reciprocal_and_div_two() {
let mut g = gelu_graph(false, false);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert_eq!(
g.nodes.values().filter(|n| n.op_type == "Gelu").count(),
1,
"the reciprocal/half-divisor encoding must also fuse"
);
assert!(g.validate().is_ok());
}
#[test]
fn declines_gelu_wrong_inner_constant() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = val(&mut g, "x");
g.add_input(x);
let half = val(&mut g, "half");
let ch = scalar_init(&mut g, "c_half", 0.5);
g.insert_node(Node::new(NodeId(0), "Mul", vec![Some(x), Some(ch)], vec![half]));
let scaled = val(&mut g, "scaled");
let cbad = scalar_init(&mut g, "c_bad", 2.0);
g.insert_node(Node::new(NodeId(0), "Div", vec![Some(x), Some(cbad)], vec![scaled]));
let e = val(&mut g, "e");
g.insert_node(Node::new(NodeId(0), "Erf", vec![Some(scaled)], vec![e]));
let one = scalar_init(&mut g, "c_one", 1.0);
let a = val(&mut g, "a");
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(e), Some(one)], vec![a]));
let out = val(&mut g, "out");
g.insert_node(Node::new(NodeId(0), "Mul", vec![Some(half), Some(a)], vec![out]));
g.add_output(out);
let before = g.num_nodes();
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert!(g.nodes.values().all(|n| n.op_type != "Gelu"));
assert!(g.nodes.values().any(|n| n.op_type == "Erf"));
assert_eq!(g.num_nodes(), before);
}
#[test]
fn declines_gelu_wrong_half_constant() {
let mut g = gelu_graph(true, true);
let ch = g.values.iter().find(|(_, v)| v.name.as_deref() == Some("c_half")).map(|(id, _)| id).unwrap();
g.set_initializer(
ch,
WeightRef::Inline(TensorData::from_raw(DataType::Float32, vec![], 0.4f32.to_le_bytes().to_vec())),
);
let before = g.num_nodes();
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert!(g.nodes.values().all(|n| n.op_type != "Gelu"));
assert_eq!(g.num_nodes(), before);
}
#[test]
fn declines_gelu_when_half_uses_different_x() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = val(&mut g, "x");
let y = val(&mut g, "y");
g.add_input(x);
g.add_input(y);
let half = val(&mut g, "half");
let ch = scalar_init(&mut g, "c_half", 0.5);
g.insert_node(Node::new(NodeId(0), "Mul", vec![Some(y), Some(ch)], vec![half]));
let scaled = val(&mut g, "scaled");
let cs = scalar_init(&mut g, "c_sqrt2", std::f32::consts::SQRT_2);
g.insert_node(Node::new(NodeId(0), "Div", vec![Some(x), Some(cs)], vec![scaled]));
let e = val(&mut g, "e");
g.insert_node(Node::new(NodeId(0), "Erf", vec![Some(scaled)], vec![e]));
let one = scalar_init(&mut g, "c_one", 1.0);
let a = val(&mut g, "a");
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(e), Some(one)], vec![a]));
let out = val(&mut g, "out");
g.insert_node(Node::new(NodeId(0), "Mul", vec![Some(half), Some(a)], vec![out]));
g.add_output(out);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert!(g.nodes.values().all(|n| n.op_type != "Gelu"));
assert!(g.nodes.values().any(|n| n.op_type == "Erf"));
}
#[test]
fn declines_gelu_when_interior_escapes() {
let mut g = gelu_graph(true, true);
let e = g.values.iter().find(|(_, v)| v.name.as_deref() == Some("e")).map(|(id, _)| id).unwrap();
let side = val(&mut g, "side");
g.insert_node(Node::new(NodeId(0), "Erf", vec![Some(e)], vec![side]));
g.add_output(side);
OpFusion::new().run(&mut g, &PassContext::new()).unwrap();
assert!(
g.nodes.values().all(|n| n.op_type != "Gelu"),
"must not fuse when an interior value escapes"
);
}
}