use onnx_runtime_ir::{Attribute, DataType, Graph, NodeId, ValueId};
use onnx_runtime_optimizer::{
OptimizationPass, OptimizerError, PassContext, Result as OptimizerResult,
};
pub(crate) const SILU_MUL_FUSION_ATTR: &str = "_cuda_silu_mul";
pub(crate) const MATMUL_NBITS_FOLDED_BIAS_ATTR: &str = "_cuda_matmul_nbits_folded_bias";
pub(crate) const GATE_UP_SWIGLU_FUSION_ATTR: &str = "_cuda_gate_up_swiglu";
const MICROSOFT_DOMAIN: &str = "com.microsoft";
const GATE_UP_SWIGLU_SUPPORTED_BLOCK_SIZE: usize = 32;
const GATE_UP_SWIGLU_SUPPORTED_BITS: i64 = 4;
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaSwiGluFusion;
pub(crate) fn cuda_optimization_passes() -> Vec<Box<dyn OptimizationPass>> {
vec![
Box::new(CudaMatMulNBitsBiasFusion),
Box::new(CudaSwiGluFusion),
Box::new(CudaGateUpSwiGluFusion),
]
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaMatMulNBitsBiasFusion;
impl OptimizationPass for CudaMatMulNBitsBiasFusion {
fn name(&self) -> &str {
"CudaMatMulNBitsBiasFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
let add_nodes: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "Add"
&& matches!(node.domain.as_str(), "" | "ai.onnx")
&& node.inputs.len() == 2
&& node.outputs.len() == 1
&& node.attributes.is_empty())
.then_some(id)
})
.collect();
let mut plans: Vec<BiasFoldPlan> = Vec::new();
for add_id in add_nodes {
if let Some(plan) = self.plan_fold(graph, add_id) {
plans.push(plan);
}
}
let changed = !plans.is_empty();
for plan in plans {
let downstream_name = graph.value(plan.add_out).name.clone();
graph.replace_all_uses(plan.add_out, plan.matmul_out);
graph.remove_node(plan.add_id);
if downstream_name.is_some() {
graph.value_mut(plan.matmul_out).name = downstream_name;
}
let mut fused = graph.node(plan.matmul_id).clone();
fused.inputs = vec![
plan.matmul_inputs[0],
plan.matmul_inputs[1],
plan.matmul_inputs[2],
None,
None,
Some(plan.bias),
];
fused
.attributes
.insert(MATMUL_NBITS_FOLDED_BIAS_ATTR.into(), Attribute::Int(1));
graph.replace_node(plan.matmul_id, fused);
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
struct BiasFoldPlan {
add_id: NodeId,
matmul_id: NodeId,
matmul_inputs: [Option<ValueId>; 3],
matmul_out: ValueId,
add_out: ValueId,
bias: ValueId,
}
impl CudaMatMulNBitsBiasFusion {
fn plan_fold(&self, graph: &Graph, add_id: NodeId) -> Option<BiasFoldPlan> {
let add = graph.try_node(add_id)?;
let lhs = add.inputs[0]?;
let rhs = add.inputs[1]?;
let add_out = add.outputs[0];
let (matmul_out, bias) = match (
self.matmul_producer(graph, lhs),
self.matmul_producer(graph, rhs),
) {
(Some(_), Some(_)) => return None,
(Some(_), None) => (lhs, rhs),
(None, Some(_)) => (rhs, lhs),
(None, None) => return None,
};
let matmul_id = self.matmul_producer(graph, matmul_out)?;
if graph.consumers(matmul_out).len() != 1 || graph.value(matmul_out).is_graph_output {
return None;
}
let matmul = graph.node(matmul_id);
let present: Vec<ValueId> = matmul.input_values().collect();
if present.len() != 3 || matmul.inputs.iter().skip(3).any(Option::is_some) {
return None;
}
let n = matmul.attr("N").and_then(Attribute::as_int)? as usize;
if !graph.initializers.contains_key(&bias) {
return None;
}
let bias_value = graph.value(bias);
let out_value = graph.value(matmul_out);
if bias_value.dtype != out_value.dtype
|| !matches!(
bias_value.dtype,
DataType::Float16 | DataType::Float32 | DataType::BFloat16
)
{
return None;
}
let bias_dims = onnx_runtime_ir::as_static_shape(&bias_value.shape)?;
if bias_dims != [n] {
return None;
}
Some(BiasFoldPlan {
add_id,
matmul_id,
matmul_inputs: [matmul.inputs[0], matmul.inputs[1], matmul.inputs[2]],
matmul_out,
add_out,
bias,
})
}
fn matmul_producer(&self, graph: &Graph, value: ValueId) -> Option<NodeId> {
let producer = graph.try_value(value)?.producer?;
let node = graph.try_node(producer)?;
(node.op_type == "MatMulNBits" && node.domain == MICROSOFT_DOMAIN).then_some(producer)
}
}
impl OptimizationPass for CudaSwiGluFusion {
fn name(&self) -> &str {
"CudaSwiGluFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
let silu_nodes: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "Silu"
&& node.domain == MICROSOFT_DOMAIN
&& node.inputs.len() == 1
&& node.outputs.len() == 1)
.then_some(id)
})
.collect();
let mut changed = false;
for silu_id in silu_nodes {
let Some(silu) = graph.try_node(silu_id) else {
continue;
};
let Some(gate) = silu.inputs[0] else {
continue;
};
let silu_output = silu.outputs[0];
if graph.outputs.contains(&silu_output) {
continue;
}
let consumers = graph.consumers(silu_output);
if consumers.len() != 1 {
continue;
}
let mul_id = consumers[0];
let mul = graph.node(mul_id);
if mul.op_type != "Mul"
|| !matches!(mul.domain.as_str(), "" | "ai.onnx")
|| mul.inputs.len() != 2
|| mul.outputs.len() != 1
|| !mul.attributes.is_empty()
{
continue;
}
let up = if mul.inputs[0] == Some(silu_output) {
mul.inputs[1]
} else if mul.inputs[1] == Some(silu_output) {
mul.inputs[0]
} else {
None
};
let Some(up) = up else {
continue;
};
let gate_value = graph.value(gate);
let up_value = graph.value(up);
if gate_value.dtype != up_value.dtype
|| gate_value.shape != up_value.shape
|| !matches!(
gate_value.dtype,
DataType::Float16 | DataType::Float32 | DataType::BFloat16
)
{
continue;
}
let mut fused = mul.clone();
fused.inputs = vec![Some(gate), Some(up)];
fused
.attributes
.insert(SILU_MUL_FUSION_ATTR.into(), Attribute::Int(1));
graph.replace_node(mul_id, fused);
graph.remove_node(silu_id);
changed = true;
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaGateUpSwiGluFusion;
struct GateUpSwiGluPlan {
mul_id: NodeId,
gate_matmul_id: NodeId,
up_matmul_id: NodeId,
activation: ValueId,
gate_weight: ValueId,
gate_scales: ValueId,
up_weight: ValueId,
up_scales: ValueId,
gate_out: ValueId,
up_out: ValueId,
}
impl OptimizationPass for CudaGateUpSwiGluFusion {
fn name(&self) -> &str {
"CudaGateUpSwiGluFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
let mul_nodes: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "Mul"
&& matches!(node.domain.as_str(), "" | "ai.onnx")
&& node.attr(SILU_MUL_FUSION_ATTR).and_then(Attribute::as_int) == Some(1)
&& node.inputs.len() == 2
&& node.outputs.len() == 1)
.then_some(id)
})
.collect();
let mut plans: Vec<GateUpSwiGluPlan> = Vec::new();
for mul_id in mul_nodes {
if let Some(plan) = self.plan_fuse(graph, mul_id) {
plans.push(plan);
}
}
let changed = !plans.is_empty();
for plan in plans {
let mut fused = graph.node(plan.gate_matmul_id).clone();
fused.id = plan.mul_id;
fused.inputs = vec![
Some(plan.activation),
Some(plan.gate_weight),
Some(plan.gate_scales),
Some(plan.up_weight),
Some(plan.up_scales),
];
fused.outputs = graph.node(plan.mul_id).outputs.clone();
fused
.attributes
.insert(GATE_UP_SWIGLU_FUSION_ATTR.into(), Attribute::Int(1));
graph.replace_node(plan.mul_id, fused);
debug_assert_eq!(graph.consumers(plan.gate_out).len(), 0);
debug_assert_eq!(graph.consumers(plan.up_out).len(), 0);
graph.remove_node(plan.gate_matmul_id);
graph.remove_node(plan.up_matmul_id);
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
impl CudaGateUpSwiGluFusion {
fn plan_fuse(&self, graph: &Graph, mul_id: NodeId) -> Option<GateUpSwiGluPlan> {
let mul = graph.try_node(mul_id)?;
let gate_out = mul.inputs[0]?;
let up_out = mul.inputs[1]?;
if gate_out == up_out {
return None;
}
let gate_matmul_id = self.matmul_producer(graph, gate_out)?;
let up_matmul_id = self.matmul_producer(graph, up_out)?;
if gate_matmul_id == up_matmul_id {
return None;
}
for out in [gate_out, up_out] {
if graph.consumers(out).len() != 1 || graph.value(out).is_graph_output {
return None;
}
}
let gate = self.eligible_projection(graph, gate_matmul_id)?;
let up = self.eligible_projection(graph, up_matmul_id)?;
if gate.activation != up.activation || gate.n != up.n || gate.k != up.k {
return None;
}
Some(GateUpSwiGluPlan {
mul_id,
gate_matmul_id,
up_matmul_id,
activation: gate.activation,
gate_weight: gate.weight,
gate_scales: gate.scales,
up_weight: up.weight,
up_scales: up.scales,
gate_out,
up_out,
})
}
fn eligible_projection(&self, graph: &Graph, matmul_id: NodeId) -> Option<Projection> {
let matmul = graph.try_node(matmul_id)?;
let present: Vec<ValueId> = matmul.input_values().collect();
if present.len() != 3 || matmul.inputs.iter().skip(3).any(Option::is_some) {
return None;
}
let n = matmul.attr("N").and_then(Attribute::as_int)? as usize;
let k = matmul.attr("K").and_then(Attribute::as_int)? as usize;
let block_size = matmul.attr("block_size").and_then(Attribute::as_int)? as usize;
let bits = matmul.attr("bits").and_then(Attribute::as_int).unwrap_or(4);
if block_size != GATE_UP_SWIGLU_SUPPORTED_BLOCK_SIZE
|| bits != GATE_UP_SWIGLU_SUPPORTED_BITS
{
return None;
}
let activation = matmul.inputs[0]?;
let weight = matmul.inputs[1]?;
let scales = matmul.inputs[2]?;
if graph.value(activation).dtype != DataType::Float16
|| graph.value(matmul.outputs[0]).dtype != DataType::Float16
|| graph.value(scales).dtype != DataType::Float16
{
return None;
}
if !graph.initializers.contains_key(&weight) || !graph.initializers.contains_key(&scales) {
return None;
}
Some(Projection {
activation,
weight,
scales,
n,
k,
})
}
fn matmul_producer(&self, graph: &Graph, value: ValueId) -> Option<NodeId> {
let producer = graph.try_value(value)?.producer?;
let node = graph.try_node(producer)?;
(node.op_type == "MatMulNBits" && node.domain == MICROSOFT_DOMAIN).then_some(producer)
}
}
struct Projection {
activation: ValueId,
weight: ValueId,
scales: ValueId,
n: usize,
k: usize,
}
#[cfg(test)]
mod tests {
use super::*;
use onnx_runtime_ir::{Dim, Node, NodeId, ValueId};
fn value(graph: &mut Graph, name: &str, dtype: DataType, width: usize) -> ValueId {
graph.create_named_value(name, dtype, vec![Dim::Static(1), Dim::Static(width)])
}
fn swiglu_graph(dtype: DataType, gate_width: usize, up_width: usize) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let gate = value(&mut graph, "gate", dtype, gate_width);
let up = value(&mut graph, "up", dtype, up_width);
let silu_output = value(&mut graph, "silu", dtype, gate_width);
let output = value(&mut graph, "output", dtype, gate_width);
graph.add_input(gate);
graph.add_input(up);
let mut silu = Node::new(NodeId(0), "Silu", vec![Some(gate)], vec![silu_output]);
silu.domain = MICROSOFT_DOMAIN.into();
graph.insert_node(silu);
graph.insert_node(Node::new(
NodeId(0),
"Mul",
vec![Some(silu_output), Some(up)],
vec![output],
));
graph.add_output(output);
graph
}
#[test]
fn fuses_equal_shape_silu_mul() {
let mut graph = swiglu_graph(DataType::Float16, 7, 7);
CudaSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.num_nodes(), 1);
let fused = graph.nodes.values().next().unwrap();
assert_eq!(fused.op_type, "Mul");
assert_eq!(
fused.attr(SILU_MUL_FUSION_ATTR).and_then(Attribute::as_int),
Some(1)
);
assert_eq!(fused.inputs.len(), 2);
assert!(graph.validate().is_ok());
}
#[test]
fn leaves_broadcast_silu_mul_separate() {
let mut graph = swiglu_graph(DataType::Float16, 7, 1);
CudaSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.num_nodes(), 2);
assert!(
graph
.nodes
.values()
.all(|node| node.attr(SILU_MUL_FUSION_ATTR).is_none())
);
}
use onnx_runtime_ir::{TensorData, WeightRef};
fn vec1d(graph: &mut Graph, name: &str, dtype: DataType, width: usize) -> ValueId {
graph.create_named_value(name, dtype, vec![Dim::Static(width)])
}
fn matmul_nbits(inputs: Vec<Option<ValueId>>, output: ValueId, k: usize, n: usize) -> Node {
let mut node = Node::new(NodeId(0), "MatMulNBits", inputs, vec![output]);
node.domain = MICROSOFT_DOMAIN.into();
node.attributes.insert("K".into(), Attribute::Int(k as i64));
node.attributes.insert("N".into(), Attribute::Int(n as i64));
node.attributes
.insert("block_size".into(), Attribute::Int(32));
node.attributes.insert("bits".into(), Attribute::Int(4));
node
}
fn qkv_bias_graph(dtype: DataType, n: usize, bias_is_initializer: bool) -> Graph {
let k = 896usize;
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let x = value(&mut graph, "x", dtype, k);
let packed = vec1d(&mut graph, "packed", DataType::Uint8, n * (k / 32) * 16);
let scales = vec1d(&mut graph, "scales", dtype, n * (k / 32));
let mm_out = value(&mut graph, "mm_out", dtype, n);
let bias = vec1d(&mut graph, "bias", dtype, n);
let out = value(&mut graph, "out", dtype, n);
graph.add_input(x);
graph.set_initializer(
packed,
WeightRef::Inline(TensorData::from_raw(
DataType::Uint8,
vec![n * (k / 32) * 16],
vec![0u8; n * (k / 32) * 16],
)),
);
graph.set_initializer(
scales,
WeightRef::Inline(TensorData::from_raw(
dtype,
vec![n * (k / 32)],
vec![0u8; n * (k / 32) * 2],
)),
);
if bias_is_initializer {
graph.set_initializer(
bias,
WeightRef::Inline(TensorData::from_raw(dtype, vec![n], vec![0u8; n * 2])),
);
} else {
graph.add_input(bias);
}
graph.insert_node(matmul_nbits(
vec![Some(x), Some(packed), Some(scales)],
mm_out,
k,
n,
));
graph.insert_node(Node::new(
NodeId(0),
"Add",
vec![Some(mm_out), Some(bias)],
vec![out],
));
graph.add_output(out);
graph
}
#[test]
fn folds_qkv_bias_into_matmul_nbits() {
let mut graph = qkv_bias_graph(DataType::Float16, 1152, true);
CudaMatMulNBitsBiasFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.num_nodes(), 1, "the Add must be folded away");
let fused = graph.nodes.values().next().unwrap();
assert_eq!(fused.op_type, "MatMulNBits");
assert_eq!(
fused
.attr(MATMUL_NBITS_FOLDED_BIAS_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert_eq!(fused.inputs.len(), 6, "bias occupies input slot 5");
assert!(fused.inputs[3].is_none() && fused.inputs[4].is_none());
assert!(fused.inputs[5].is_some(), "bias must be wired at index 5");
let out = fused.outputs[0];
assert_eq!(graph.outputs, vec![out]);
assert_eq!(graph.value(out).name.as_deref(), Some("out"));
assert!(graph.validate().is_ok());
}
#[test]
fn does_not_fold_non_initializer_bias() {
let mut graph = qkv_bias_graph(DataType::Float16, 1152, false);
CudaMatMulNBitsBiasFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.num_nodes(), 2, "a runtime bias must not be folded");
assert!(
graph
.nodes
.values()
.all(|node| node.attr(MATMUL_NBITS_FOLDED_BIAS_ATTR).is_none())
);
}
#[test]
fn does_not_fold_wrong_shape_bias() {
let mut graph = qkv_bias_graph(DataType::Float16, 1152, true);
let bias = graph
.values
.iter()
.find_map(|(id, v)| (v.name.as_deref() == Some("bias")).then_some(id))
.unwrap();
graph.value_mut(bias).shape = vec![Dim::Static(2), Dim::Static(1152)];
CudaMatMulNBitsBiasFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.num_nodes(), 2, "a non-[N] bias must not be folded");
}
#[test]
fn does_not_fold_when_matmul_output_is_shared() {
let mut graph = qkv_bias_graph(DataType::Float16, 1152, true);
let mm_out = graph
.values
.iter()
.find_map(|(id, v)| (v.name.as_deref() == Some("mm_out")).then_some(id))
.unwrap();
let sink = value(&mut graph, "sink", DataType::Float16, 1152);
graph.insert_node(Node::new(NodeId(0), "Neg", vec![Some(mm_out)], vec![sink]));
graph.add_output(sink);
CudaMatMulNBitsBiasFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.all(|node| node.attr(MATMUL_NBITS_FOLDED_BIAS_ATTR).is_none()),
"a shared GEMV output must not be folded"
);
}
const QWEN_GATE_UP_K: usize = 896;
const QWEN_GATE_UP_N: usize = 4864;
const NON_QWEN_GATE_UP_SHAPES: [(usize, usize); 2] = [
(2048, 5632), (2048, 16384), ];
fn projection(graph: &mut Graph, tag: &str, x: ValueId, k: usize, n: usize) -> ValueId {
projection_dtype(graph, tag, x, k, n, DataType::Float16)
}
fn projection_dtype(
graph: &mut Graph,
tag: &str,
x: ValueId,
k: usize,
n: usize,
dtype: DataType,
) -> ValueId {
let scale_bytes = if dtype == DataType::Float16 { 2 } else { 4 };
let packed = vec1d(
graph,
&format!("{tag}_packed"),
DataType::Uint8,
n * (k / 32) * 16,
);
let scales = vec1d(graph, &format!("{tag}_scales"), dtype, n * (k / 32));
let out = value(graph, &format!("{tag}_out"), dtype, n);
graph.set_initializer(
packed,
WeightRef::Inline(TensorData::from_raw(
DataType::Uint8,
vec![n * (k / 32) * 16],
vec![0u8; n * (k / 32) * 16],
)),
);
graph.set_initializer(
scales,
WeightRef::Inline(TensorData::from_raw(
dtype,
vec![n * (k / 32)],
vec![0u8; n * (k / 32) * scale_bytes],
)),
);
graph.insert_node(matmul_nbits(
vec![Some(x), Some(packed), Some(scales)],
out,
k,
n,
));
out
}
fn gate_up_graph(k: usize, n: usize, shared: bool) -> Graph {
gate_up_graph_dtype_impl(k, n, shared, DataType::Float16)
}
fn gate_up_graph_dtype(k: usize, n: usize, dtype: DataType) -> Graph {
gate_up_graph_dtype_impl(k, n, true, dtype)
}
fn gate_up_graph_dtype_impl(k: usize, n: usize, shared: bool, dtype: DataType) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let x = value(&mut graph, "x", dtype, k);
graph.add_input(x);
let up_x = if shared {
x
} else {
let x2 = value(&mut graph, "x2", dtype, k);
graph.add_input(x2);
x2
};
let gate_out = projection_dtype(&mut graph, "gate", x, k, n, dtype);
let up_out = projection_dtype(&mut graph, "up", up_x, k, n, dtype);
let out = value(&mut graph, "output", dtype, n);
let mut mul = Node::new(
NodeId(0),
"Mul",
vec![Some(gate_out), Some(up_out)],
vec![out],
);
mul.attributes
.insert(SILU_MUL_FUSION_ATTR.into(), Attribute::Int(1));
graph.insert_node(mul);
graph.add_output(out);
graph
}
fn gate_up_graph_asymmetric(k: usize, n_gate: usize, n_up: usize) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let x = value(&mut graph, "x", DataType::Float16, k);
graph.add_input(x);
let gate_out = projection(&mut graph, "gate", x, k, n_gate);
let up_out = projection(&mut graph, "up", x, k, n_up);
let out = value(&mut graph, "output", DataType::Float16, n_gate);
let mut mul = Node::new(
NodeId(0),
"Mul",
vec![Some(gate_out), Some(up_out)],
vec![out],
);
mul.attributes
.insert(SILU_MUL_FUSION_ATTR.into(), Attribute::Int(1));
graph.insert_node(mul);
graph.add_output(out);
graph
}
#[test]
fn fuses_paired_gate_up_swiglu() {
let mut graph = gate_up_graph(QWEN_GATE_UP_K, QWEN_GATE_UP_N, true);
CudaGateUpSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph.num_nodes(),
1,
"both projections and the Mul collapse into one node"
);
let fused = graph.nodes.values().next().unwrap();
assert_eq!(fused.op_type, "MatMulNBits");
assert_eq!(fused.domain, MICROSOFT_DOMAIN);
assert_eq!(
fused
.attr(GATE_UP_SWIGLU_FUSION_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert!(
fused.attr(SILU_MUL_FUSION_ATTR).is_none(),
"the silu_mul marker must not leak onto the fused MatMulNBits"
);
assert_eq!(
fused.inputs.len(),
5,
"inputs are [x, W_gate, scales_gate, W_up, scales_up]"
);
assert!(fused.inputs.iter().all(Option::is_some));
assert_eq!(
fused.attr("N").and_then(Attribute::as_int),
Some(QWEN_GATE_UP_N as i64)
);
let out = fused.outputs[0];
assert_eq!(graph.outputs, vec![out]);
assert_eq!(graph.value(out).name.as_deref(), Some("output"));
assert!(graph.validate().is_ok());
}
#[test]
fn fuses_paired_gate_up_swiglu_for_non_qwen_shapes() {
for (k, n) in NON_QWEN_GATE_UP_SHAPES {
let mut graph = gate_up_graph(k, n, true);
CudaGateUpSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph.num_nodes(),
1,
"gate/up SwiGLU must fuse for non-Qwen shape K={k}, N={n}"
);
let fused = graph.nodes.values().next().unwrap();
assert_eq!(
fused
.attr(GATE_UP_SWIGLU_FUSION_ATTR)
.and_then(Attribute::as_int),
Some(1),
"fused marker missing for K={k}, N={n}"
);
assert_eq!(fused.attr("N").and_then(Attribute::as_int), Some(n as i64));
assert!(graph.validate().is_ok());
}
}
#[test]
fn does_not_fuse_when_activation_differs() {
let mut graph = gate_up_graph(QWEN_GATE_UP_K, QWEN_GATE_UP_N, false);
CudaGateUpSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph.num_nodes(),
3,
"projections on different activations must not be paired"
);
assert!(
graph
.nodes
.values()
.all(|node| node.attr(GATE_UP_SWIGLU_FUSION_ATTR).is_none())
);
}
fn set_matmul_attr(graph: &mut Graph, name: &str, value: i64) {
let ids: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| (node.op_type == "MatMulNBits").then_some(id))
.collect();
for id in ids {
graph
.node_mut(id)
.attributes
.insert(name.into(), Attribute::Int(value));
}
}
fn asserts_not_fused(graph: &Graph) {
assert_eq!(
graph.num_nodes(),
3,
"incompatible projections must stay separate"
);
assert!(
graph
.nodes
.values()
.all(|node| node.attr(GATE_UP_SWIGLU_FUSION_ATTR).is_none())
);
}
#[test]
fn does_not_fuse_incompatible_block_size() {
let mut graph = gate_up_graph(QWEN_GATE_UP_K, QWEN_GATE_UP_N, true);
set_matmul_attr(&mut graph, "block_size", 64);
CudaGateUpSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
asserts_not_fused(&graph);
}
#[test]
fn does_not_fuse_incompatible_bits() {
let mut graph = gate_up_graph(QWEN_GATE_UP_K, QWEN_GATE_UP_N, true);
set_matmul_attr(&mut graph, "bits", 8);
CudaGateUpSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
asserts_not_fused(&graph);
}
#[test]
fn does_not_fuse_non_fp16_projection() {
let mut graph = gate_up_graph_dtype(QWEN_GATE_UP_K, QWEN_GATE_UP_N, DataType::Float32);
CudaGateUpSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
asserts_not_fused(&graph);
}
#[test]
fn does_not_fuse_mismatched_output_width() {
let mut graph =
gate_up_graph_asymmetric(QWEN_GATE_UP_K, QWEN_GATE_UP_N, QWEN_GATE_UP_N / 2);
CudaGateUpSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.all(|node| node.attr(GATE_UP_SWIGLU_FUSION_ATTR).is_none())
);
}
#[test]
fn does_not_fuse_untagged_mul() {
let mut graph = gate_up_graph(QWEN_GATE_UP_K, QWEN_GATE_UP_N, true);
let mul_id = graph
.nodes
.iter()
.find_map(|(id, node)| (node.op_type == "Mul").then_some(id))
.unwrap();
graph
.node_mut(mul_id)
.attributes
.remove(SILU_MUL_FUSION_ATTR);
CudaGateUpSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.num_nodes(), 3, "an untagged Mul must not be fused");
}
#[test]
fn gate_up_pass_chains_after_swiglu_fusion() {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let x = value(&mut graph, "x", DataType::Float16, QWEN_GATE_UP_K);
graph.add_input(x);
let gate_out = projection(&mut graph, "gate", x, QWEN_GATE_UP_K, QWEN_GATE_UP_N);
let up_out = projection(&mut graph, "up", x, QWEN_GATE_UP_K, QWEN_GATE_UP_N);
let silu_out = value(&mut graph, "silu", DataType::Float16, QWEN_GATE_UP_N);
let out = value(&mut graph, "output", DataType::Float16, QWEN_GATE_UP_N);
let mut silu = Node::new(NodeId(0), "Silu", vec![Some(gate_out)], vec![silu_out]);
silu.domain = MICROSOFT_DOMAIN.into();
graph.insert_node(silu);
graph.insert_node(Node::new(
NodeId(0),
"Mul",
vec![Some(silu_out), Some(up_out)],
vec![out],
));
graph.add_output(out);
for pass in cuda_optimization_passes() {
pass.run(&mut graph, &PassContext::new()).unwrap();
}
assert_eq!(graph.num_nodes(), 1);
let fused = graph.nodes.values().next().unwrap();
assert_eq!(fused.op_type, "MatMulNBits");
assert_eq!(
fused
.attr(GATE_UP_SWIGLU_FUSION_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert!(graph.validate().is_ok());
}
}