use std::collections::{BTreeSet, HashMap, HashSet};
use onnx_runtime_ir::{Attribute, DataType, Graph, Node, NodeId, TensorData, 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,
#[cfg(test)]
replacement_domain: 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(),
#[cfg(test)]
replacement_domain: CONTRIB_DOMAIN.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(),
#[cfg(test)]
replacement_domain: CONTRIB_DOMAIN.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(),
#[cfg(test)]
replacement_domain: CONTRIB_DOMAIN.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(),
#[cfg(test)]
replacement_domain: CONTRIB_DOMAIN.to_string(),
kind: RewriteKind::Gelu,
}
}
#[cfg(test)]
fn with_replacement_domain(mut self, domain: &str) -> Self {
self.replacement_domain = domain.to_string();
self
}
pub fn pattern_name(&self) -> &str {
&self.name
}
pub fn find_match(&self, graph: &Graph) -> Option<PatternMatch> {
for start in graph.nodes.keys() {
if let Some(m) = self.try_match_at(graph, start) {
return Some(m);
}
}
None
}
fn try_match_at(&self, graph: &Graph, start: NodeId) -> Option<PatternMatch> {
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),
}
}
fn affected_candidate_starts(&self, graph: &Graph, matched: &PatternMatch) -> Vec<NodeId> {
let max_depth = match self.kind {
RewriteKind::LayerNorm => 10,
RewriteKind::Attention => 6,
RewriteKind::Gelu => 5,
RewriteKind::Structural => self.ops.len(),
};
let mut affected = HashSet::new();
let mut frontier: Vec<(NodeId, usize)> = matched
.external_inputs
.iter()
.filter_map(|&value| graph.value(value).producer)
.map(|producer| (producer, 0))
.collect();
while let Some((node_id, depth)) = frontier.pop() {
if !affected.insert(node_id) || depth >= max_depth.saturating_sub(1) {
continue;
}
frontier.extend(
graph
.node(node_id)
.input_values()
.filter_map(|value| graph.value(value).producer)
.map(|producer| (producer, depth + 1)),
);
}
affected.into_iter().collect()
}
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
.consumers(value)
.into_iter()
.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
.consumers(mean)
.into_iter()
.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
.consumers(out)
.into_iter()
.any(|consumer| !matched_set.contains(&consumer))
})
});
if escapes {
continue;
}
let fa = graph.node(final_add);
if fa.outputs.len() != 1 {
continue;
}
let output = fa.outputs[0];
let survives = graph.outputs.contains(&output)
|| graph
.consumers(output)
.into_iter()
.any(|consumer| !matched_set.contains(&consumer));
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.consumers(sm_out).into_iter().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
.consumers(o)
.into_iter()
.any(|consumer| !matched_set.contains(&consumer))
})
});
if escapes {
return None;
}
let survives = graph.outputs.contains(&output)
|| graph
.consumers(output)
.into_iter()
.any(|consumer| !matched_set.contains(&consumer));
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.consumers(k_side) == [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
.consumers(o)
.into_iter()
.any(|consumer| !matched_set.contains(&consumer))
})
});
if escapes {
return None;
}
let survives = graph.outputs.contains(&output)
|| graph
.consumers(output)
.into_iter()
.any(|consumer| !matched_set.contains(&consumer));
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
.consumers(out)
.into_iter()
.any(|consumer| !chain_set.contains(&consumer))
{
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 survives = graph.outputs.contains(&output)
|| graph
.consumers(output)
.into_iter()
.any(|consumer| !chain_set.contains(&consumer));
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<()> {
self.apply_fusion_returning_id(graph, m).map(|_| ())
}
fn apply_fusion_returning_id(&self, graph: &mut Graph, m: &PatternMatch) -> Result<NodeId> {
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;
#[cfg(not(test))]
{
fused.domain = CONTRIB_DOMAIN.to_string();
}
#[cfg(test)]
{
fused.domain = self.replacement_domain.clone();
}
if !fused.domain.is_empty() {
graph.opset_imports.entry(fused.domain.clone()).or_insert(1);
}
Ok(graph.insert_node(fused))
}
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 axis = reduce_single_axis(graph, rm1)?;
if reduce_single_axis(graph, rm2)? != axis {
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 reduce_single_axis(graph: &Graph, rm: &Node) -> Option<i64> {
if let Some(keepdims) = rm.attr("keepdims").and_then(Attribute::as_int)
&& keepdims != 1
{
return None;
}
let axes: Vec<i64> = if let Some(axes) = rm.attr("axes").and_then(Attribute::as_ints) {
axes.to_vec()
} else {
let axes_value = rm.inputs.get(1).copied().flatten()?;
read_i64_vector(graph, axes_value)?
};
let [axis] = axes.as_slice() else {
return None;
};
Some(*axis)
}
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,
}
}
fn read_i64_vector(graph: &Graph, value: ValueId) -> Option<Vec<i64>> {
if let Some(WeightRef::Inline(tensor)) = graph.initializers.get(&value) {
return i64_axes_from_tensor(tensor);
}
let producer = graph.value(value).producer?;
let node = graph.node(producer);
if node.op_type != "Constant" || !(node.domain.is_empty() || node.domain == "ai.onnx") {
return None;
}
if let Some(Attribute::Tensor(tensor)) = node.attr("value") {
return i64_axes_from_tensor(tensor);
}
if let Some(ints) = node.attr("value_ints").and_then(Attribute::as_ints) {
return Some(ints.to_vec());
}
None
}
fn i64_axes_from_tensor(tensor: &TensorData) -> Option<Vec<i64>> {
if tensor.dtype != DataType::Int64 || !tensor.data.len().is_multiple_of(8) {
return None;
}
let numel = tensor.data.len() / 8;
if tensor.dims.len() != 1 || tensor.dims[0] != numel {
return None;
}
tensor
.data
.chunks_exact(8)
.map(|chunk| Some(i64::from_le_bytes(chunk.try_into().ok()?)))
.collect()
}
#[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>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum ScanCandidateSource {
Initial,
Revisit,
}
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 }
}
fn run_resumable(
&self,
graph: &mut Graph,
mut observe_fusion: impl FnMut(&str, ScanCandidateSource, NodeId, &[NodeId], &[NodeId], NodeId),
) -> Result<()> {
for pattern in &self.patterns {
let candidates: Vec<u32> = graph.nodes.keys().map(|id| id.0).collect();
let mut cursor = 0;
let mut revisits = BTreeSet::new();
loop {
let initial = candidates.get(cursor).copied();
let revisit = revisits.first().copied();
let (raw_id, source) = match (initial, revisit) {
(None, None) => break,
(Some(id), None) => {
cursor += 1;
(id, ScanCandidateSource::Initial)
}
(None, Some(_)) => {
(revisits.pop_first().unwrap(), ScanCandidateSource::Revisit)
}
(Some(id), Some(revisit)) if id <= revisit => {
cursor += 1;
if id == revisit {
revisits.pop_first();
}
(id, ScanCandidateSource::Initial)
}
(Some(_), Some(_)) => {
(revisits.pop_first().unwrap(), ScanCandidateSource::Revisit)
}
};
let start = NodeId(raw_id);
let Some(matched) = pattern.try_match_at(graph, start) else {
continue;
};
let affected = pattern.affected_candidate_starts(graph, &matched);
let fused_id = pattern.apply_fusion_returning_id(graph, &matched)?;
observe_fusion(
pattern.pattern_name(),
source,
start,
&matched.nodes,
&affected,
fused_id,
);
revisits.insert(fused_id.0);
for candidate in affected {
if graph.try_node(candidate).is_some() {
revisits.insert(candidate.0);
}
}
}
}
Ok(())
}
#[cfg(test)]
fn run_with_fusion_observer(
&self,
graph: &mut Graph,
observe_fusion: impl FnMut(&str, ScanCandidateSource, NodeId, &[NodeId], &[NodeId], NodeId),
) -> Result<()> {
self.run_resumable(graph, observe_fusion)
}
}
impl OptimizationPass for OpFusion {
fn name(&self) -> &str {
"OpFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> Result<()> {
self.run_resumable(graph, |_, _, _, _, _, _| {})
}
}
#[cfg(test)]
mod tests;