use onnx_runtime_ir::{
Attribute, DataType, Graph, NodeId, TensorData, ValueId, WeightRef, static_shape,
};
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";
pub(crate) const MATMUL_NBITS_RMSNORM_PROLOGUE_ATTR: &str = "_cuda_matmul_nbits_rmsnorm_prologue";
pub(crate) const MATMUL_NBITS_RMSNORM_EPSILON_ATTR: &str = "_cuda_matmul_nbits_rmsnorm_epsilon";
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(CudaFoldConstantTranspose),
Box::new(CudaDropNormalizationCasts),
Box::new(CudaMatMulNBitsBiasFusion),
Box::new(CudaSwiGluFusion),
Box::new(CudaGateUpSwiGluFusion),
Box::new(CudaSkipRmsNormMatMulFusion),
Box::new(CudaOnDeviceConstantSelect),
]
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaDropNormalizationCasts;
const NORM_CAST_FOLD_DISABLE_ENV: &str = "ONNX_GENAI_CUDA_DISABLE_NORM_CAST_FOLD";
fn norm_cast_fold_disabled() -> bool {
std::env::var_os(NORM_CAST_FOLD_DISABLE_ENV)
.is_some_and(|value| value != "0" && !value.is_empty())
}
struct NormCastFoldPlan {
node_id: NodeId,
new_inputs: Vec<Option<ValueId>>,
retyped_outputs: Vec<ValueId>,
output_cast_bypass: Vec<(ValueId, ValueId, NodeId)>,
dead_input_casts: Vec<NodeId>,
}
impl OptimizationPass for CudaDropNormalizationCasts {
fn name(&self) -> &str {
"CudaDropNormalizationCasts"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
if norm_cast_fold_disabled() {
return Ok(());
}
let mut changed = false;
loop {
let candidates: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| Self::activation_input_indices(node).map(|_| id))
.collect();
let Some(plan) = candidates
.into_iter()
.find_map(|id| self.plan_fold(graph, id))
else {
break;
};
self.apply_fold(graph, plan);
changed = true;
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
impl CudaDropNormalizationCasts {
fn apply_fold(&self, graph: &mut Graph, plan: NormCastFoldPlan) {
let mut node = graph.node(plan.node_id).clone();
node.inputs = plan.new_inputs;
graph.replace_node(plan.node_id, node);
for output in plan.retyped_outputs {
graph.value_mut(output).dtype = DataType::Float16;
}
for (cast_output, norm_output, cast_id) in plan.output_cast_bypass {
graph.replace_all_uses(cast_output, norm_output);
graph.remove_node(cast_id);
}
for cast_id in plan.dead_input_casts {
if let Some(cast) = graph.try_node(cast_id) {
let cast_output = *cast.outputs.first().expect("Cast has one output");
if graph.consumers(cast_output).is_empty()
&& !graph.value(cast_output).is_graph_output
{
graph.remove_node(cast_id);
}
}
}
}
}
impl CudaDropNormalizationCasts {
fn activation_input_indices(node: &onnx_runtime_ir::Node) -> Option<&'static [usize]> {
match (node.op_type.as_str(), node.domain.as_str()) {
("SkipSimplifiedLayerNormalization", MICROSOFT_DOMAIN) => Some(&[0, 1]),
("SimplifiedLayerNormalization", "" | "ai.onnx") => Some(&[0]),
_ => None,
}
}
fn fp32_cast_from_fp16(&self, graph: &Graph, value: ValueId) -> Option<(NodeId, ValueId)> {
let producer = graph.try_value(value)?.producer?;
let node = graph.try_node(producer)?;
if node.op_type != "Cast" || !matches!(node.domain.as_str(), "" | "ai.onnx") {
return None;
}
if cast_target(node)? != DataType::Float32 {
return None;
}
let source = node.inputs.first().copied().flatten()?;
if graph.try_value(source)?.dtype != DataType::Float16 {
return None;
}
Some((producer, source))
}
fn plan_fold(&self, graph: &Graph, node_id: NodeId) -> Option<NormCastFoldPlan> {
let node = graph.try_node(node_id)?;
let activation_indices = Self::activation_input_indices(node)?;
if node.op_type == "SkipSimplifiedLayerNormalization"
&& node.inputs.get(3).copied().flatten().is_some()
{
return None;
}
let mut new_inputs = node.inputs.clone();
let mut dead_input_casts = Vec::new();
for &index in activation_indices {
let value = node.inputs.get(index).copied().flatten()?;
let (cast_id, source) = self.fp32_cast_from_fp16(graph, value)?;
new_inputs[index] = Some(source);
dead_input_casts.push(cast_id);
}
let mut retyped_outputs = Vec::new();
let mut output_cast_bypass = Vec::new();
for &output in &node.outputs {
if graph.value(output).is_graph_output {
return None;
}
let consumers = graph.consumers(output);
if consumers.is_empty() {
continue;
}
for consumer in &consumers {
let cast = graph.try_node(*consumer)?;
if cast.op_type != "Cast" || !matches!(cast.domain.as_str(), "" | "ai.onnx") {
return None;
}
if cast_target(cast)? != DataType::Float16 {
return None;
}
let cast_output = *cast.outputs.first()?;
output_cast_bypass.push((cast_output, output, *consumer));
}
retyped_outputs.push(output);
}
let primary = *node.outputs.first()?;
if !retyped_outputs.contains(&primary) {
return None;
}
Some(NormCastFoldPlan {
node_id,
new_inputs,
retyped_outputs,
output_cast_bypass,
dead_input_casts,
})
}
}
fn cast_target(node: &onnx_runtime_ir::Node) -> Option<DataType> {
let raw = node.attr("to").and_then(Attribute::as_int)?;
DataType::from_onnx(raw as i32)
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaFoldConstantTranspose;
struct TransposeFoldPlan {
node: NodeId,
output: ValueId,
dtype: DataType,
out_dims: Vec<usize>,
bytes: Vec<u8>,
}
impl OptimizationPass for CudaFoldConstantTranspose {
fn name(&self) -> &str {
"CudaFoldConstantTranspose"
}
fn run(&self, graph: &mut Graph, ctx: &PassContext) -> OptimizerResult<()> {
let candidates: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "Transpose"
&& matches!(node.domain.as_str(), "" | "ai.onnx")
&& node.inputs.len() == 1
&& node.outputs.len() == 1)
.then_some(id)
})
.collect();
let mut plans: Vec<TransposeFoldPlan> = Vec::new();
for node_id in candidates {
if let Some(plan) = self.plan_fold(graph, ctx, node_id) {
plans.push(plan);
}
}
let changed = !plans.is_empty();
for plan in plans {
graph.remove_node(plan.node);
if graph.try_value(plan.output).is_none() {
continue;
}
let value = graph.value_mut(plan.output);
value.dtype = plan.dtype;
value.shape = static_shape(plan.out_dims.clone());
let tensor = TensorData::from_raw(plan.dtype, plan.out_dims, plan.bytes);
graph.set_initializer(plan.output, WeightRef::Inline(tensor));
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
impl CudaFoldConstantTranspose {
fn plan_fold(
&self,
graph: &Graph,
ctx: &PassContext,
node_id: NodeId,
) -> Option<TransposeFoldPlan> {
let node = graph.try_node(node_id)?;
let input = node.inputs[0]?;
let output = node.outputs[0];
if graph.try_value(input)?.producer.is_some() {
return None;
}
let weight = graph.initializers.get(&input)?;
let dtype = weight.dtype();
let elem = dtype.byte_size();
if elem == 0 || dtype.is_sub_byte() {
return None;
}
let dims = weight.dims().to_vec();
let rank = dims.len();
let perm = transpose_perm(node, rank)?;
let src = ctx.initializer_bytes(weight)?;
let expected = dims.iter().product::<usize>().checked_mul(elem)?;
if src.len() != expected {
return None;
}
let out_dims: Vec<usize> = perm.iter().map(|&p| dims[p]).collect();
let bytes = permute_bytes(src, &dims, &perm, elem);
Some(TransposeFoldPlan {
node: node_id,
output,
dtype,
out_dims,
bytes,
})
}
}
fn transpose_perm(node: &onnx_runtime_ir::Node, rank: usize) -> Option<Vec<usize>> {
match node.attr("perm").and_then(Attribute::as_ints) {
None => Some((0..rank).rev().collect()),
Some(perm) => {
if perm.len() != rank {
return None;
}
let mut axes: Vec<usize> = Vec::with_capacity(rank);
let mut seen = vec![false; rank];
for &p in perm {
let p = usize::try_from(p).ok()?;
if p >= rank || seen[p] {
return None;
}
seen[p] = true;
axes.push(p);
}
Some(axes)
}
}
}
fn permute_bytes(src: &[u8], dims: &[usize], perm: &[usize], elem: usize) -> Vec<u8> {
let rank = dims.len();
let out_dims: Vec<usize> = perm.iter().map(|&p| dims[p]).collect();
let total: usize = out_dims.iter().product();
let mut dst = vec![0u8; total * elem];
if total == 0 {
return dst;
}
let mut in_strides = vec![0usize; rank];
let mut stride = 1usize;
for axis in (0..rank).rev() {
in_strides[axis] = stride;
stride *= dims[axis];
}
let out_in_stride: Vec<usize> = perm.iter().map(|&p| in_strides[p]).collect();
let mut coord = vec![0usize; rank];
let mut in_off = 0usize;
for out_index in 0..total {
let dst_off = out_index * elem;
let src_off = in_off * elem;
dst[dst_off..dst_off + elem].copy_from_slice(&src[src_off..src_off + elem]);
for axis in (0..rank).rev() {
coord[axis] += 1;
in_off += out_in_stride[axis];
if coord[axis] == out_dims[axis] {
coord[axis] = 0;
in_off -= out_in_stride[axis] * out_dims[axis];
} else {
break;
}
}
}
dst
}
#[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)
}
}
const RMSNORM_FUSION_SUPPORTED_BITS: [i64; 2] = [4, 8];
const RMSNORM_FUSION_SUPPORTED_BLOCK_SIZE: i64 = 32;
const RMSNORM_FUSION_WARP_HALF4_MULTIPLE: usize = 128;
const RMSNORM_FUSION_DISABLE_ENV: &str = "ONNX_GENAI_CUDA_DISABLE_RMSNORM_FUSION";
fn rmsnorm_fusion_disabled() -> bool {
std::env::var_os(RMSNORM_FUSION_DISABLE_ENV)
.is_some_and(|value| value != "0" && !value.is_empty())
}
const RMSNORM_FUSION_MIN_HIDDEN: usize = 10 * RMSNORM_FUSION_WARP_HALF4_MULTIPLE;
const RMSNORM_FUSION_MIN_HIDDEN_ENV: &str = "ONNX_GENAI_RMSNORM_MIN_HIDDEN";
fn env_usize(name: &str, default: usize) -> usize {
std::env::var(name)
.ok()
.and_then(|value| value.trim().parse::<usize>().ok())
.unwrap_or(default)
}
fn fusion_benefit_is_positive(norm_size: usize, _fanout: usize, _following_min_n: usize) -> bool {
let floor = env_usize(RMSNORM_FUSION_MIN_HIDDEN_ENV, RMSNORM_FUSION_MIN_HIDDEN);
norm_size >= floor
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaSkipRmsNormMatMulFusion;
struct SkipRmsNormPlan {
skip_id: NodeId,
preceding_id: NodeId,
preceding_inputs: [ValueId; 3],
preceding_zero_points: Option<ValueId>,
preceding_out: ValueId,
residual: ValueId,
gamma: ValueId,
epsilon: f32,
normalized_out: ValueId,
sum_out: Option<ValueId>,
following_ids: Vec<NodeId>,
}
impl OptimizationPass for CudaSkipRmsNormMatMulFusion {
fn name(&self) -> &str {
"CudaSkipRmsNormMatMulFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
if rmsnorm_fusion_disabled() {
return Ok(());
}
let skip_nodes: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "SkipSimplifiedLayerNormalization"
&& node.domain == MICROSOFT_DOMAIN)
.then_some(id)
})
.collect();
let mut plans: Vec<SkipRmsNormPlan> = Vec::new();
let mut used_nodes: std::collections::HashSet<NodeId> = std::collections::HashSet::new();
for skip_id in skip_nodes {
if let Some(plan) = self.plan_fusion(graph, skip_id) {
if std::iter::once(plan.preceding_id)
.chain(std::iter::once(plan.skip_id))
.chain(plan.following_ids.iter().copied())
.any(|id| used_nodes.contains(&id))
{
continue;
}
used_nodes.insert(plan.preceding_id);
used_nodes.insert(plan.skip_id);
used_nodes.extend(plan.following_ids.iter().copied());
plans.push(plan);
}
}
let changed = !plans.is_empty();
let mut value_redirects: std::collections::HashMap<ValueId, ValueId> =
std::collections::HashMap::new();
let resolve = |redirects: &std::collections::HashMap<ValueId, ValueId>,
mut value: ValueId| {
while let Some(&next) = redirects.get(&value) {
if next == value {
break;
}
value = next;
}
value
};
for plan in plans {
let residual = resolve(&value_redirects, plan.residual);
graph.replace_all_uses(plan.normalized_out, plan.preceding_out);
value_redirects.insert(plan.normalized_out, plan.preceding_out);
if let Some(sum_out) = plan.sum_out {
graph.replace_all_uses(sum_out, plan.preceding_out);
value_redirects.insert(sum_out, plan.preceding_out);
}
for following_id in &plan.following_ids {
let mut fused = graph.node(*following_id).clone();
let gamma_slot = if fused.attr(GATE_UP_SWIGLU_FUSION_ATTR).is_some() {
5
} else {
6
};
while fused.inputs.len() <= gamma_slot {
fused.inputs.push(None);
}
fused.inputs[gamma_slot] = Some(plan.gamma);
fused
.attributes
.insert(MATMUL_NBITS_RMSNORM_PROLOGUE_ATTR.into(), Attribute::Int(1));
fused.attributes.insert(
MATMUL_NBITS_RMSNORM_EPSILON_ATTR.into(),
Attribute::Float(plan.epsilon),
);
graph.replace_node(*following_id, fused);
}
let mut preceding = graph.node(plan.preceding_id).clone();
preceding.inputs = vec![
Some(plan.preceding_inputs[0]),
Some(plan.preceding_inputs[1]),
Some(plan.preceding_inputs[2]),
plan.preceding_zero_points,
None,
Some(residual),
];
preceding
.attributes
.insert(MATMUL_NBITS_FOLDED_BIAS_ATTR.into(), Attribute::Int(1));
graph.replace_node(plan.preceding_id, preceding);
graph.remove_node(plan.skip_id);
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
impl CudaSkipRmsNormMatMulFusion {
fn plan_fusion(&self, graph: &Graph, skip_id: NodeId) -> Option<SkipRmsNormPlan> {
let skip = graph.try_node(skip_id)?;
if skip.inputs.len() < 3 {
return None;
}
let input_value = skip.inputs[0]?;
let skip_value = skip.inputs[1]?;
let gamma = skip.inputs[2]?;
if skip.inputs.get(3).copied().flatten().is_some() {
return None;
}
let normalized_out = *skip.outputs.first()?;
for stat in skip.outputs.iter().skip(1).take(2) {
if !graph.consumers(*stat).is_empty() || graph.value(*stat).is_graph_output {
return None;
}
}
let sum_out = skip.outputs.get(3).copied();
if graph.value(normalized_out).is_graph_output {
return None;
}
if let Some(sum_out) = sum_out
&& graph.value(sum_out).is_graph_output
{
return None;
}
let input_meta = graph.value(input_value);
let skip_meta = graph.value(skip_value);
let gamma_meta = graph.value(gamma);
if input_meta.dtype != DataType::Float16
|| skip_meta.dtype != DataType::Float16
|| (gamma_meta.dtype != DataType::Float16 && gamma_meta.dtype != DataType::Float32)
{
return None;
}
if input_meta.shape != skip_meta.shape {
return None;
}
let norm_size = input_meta.shape.last()?.as_static()?;
if norm_size == 0 || !norm_size.is_multiple_of(RMSNORM_FUSION_WARP_HALF4_MULTIPLE) {
return None;
}
let gamma_dims = onnx_runtime_ir::as_static_shape(&gamma_meta.shape)?;
if gamma_dims != [norm_size] {
return None;
}
let (preceding_out, residual) = match (
self.preceding_gemv(graph, input_value, norm_size),
self.preceding_gemv(graph, skip_value, norm_size),
) {
(Some(_), Some(_)) => return None,
(Some(_), None) => (input_value, skip_value),
(None, Some(_)) => (skip_value, input_value),
(None, None) => return None,
};
let preceding_id = self.preceding_gemv(graph, preceding_out, norm_size)?;
if graph.value(residual).dtype != DataType::Float16 {
return None;
}
let preceding = graph.node(preceding_id);
let preceding_inputs = [
preceding.inputs[0]?,
preceding.inputs[1]?,
preceding.inputs[2]?,
];
let preceding_zero_points = preceding.inputs.get(3).copied().flatten();
let following_ids = graph.consumers(normalized_out);
if following_ids.is_empty() {
return None;
}
for following_id in &following_ids {
if !self.following_gemv_is_fusable(graph, *following_id) {
return None;
}
}
if following_ids.contains(&preceding_id) {
return None;
}
let following_min_n = following_ids
.iter()
.filter_map(|id| graph.node(*id).attr("N").and_then(Attribute::as_int))
.min()
.unwrap_or(0)
.max(0) as usize;
if !fusion_benefit_is_positive(norm_size, following_ids.len(), following_min_n) {
return None;
}
let epsilon = skip
.attr("epsilon")
.and_then(Attribute::as_float)
.unwrap_or(1e-5);
Some(SkipRmsNormPlan {
skip_id,
preceding_id,
preceding_inputs,
preceding_zero_points,
preceding_out,
residual,
gamma,
epsilon,
normalized_out,
sum_out,
following_ids,
})
}
fn preceding_gemv(&self, graph: &Graph, value: ValueId, norm_size: usize) -> Option<NodeId> {
let producer = graph.try_value(value)?.producer?;
let node = graph.try_node(producer)?;
if node.op_type != "MatMulNBits" || node.domain != MICROSOFT_DOMAIN {
return None;
}
if node.attr(GATE_UP_SWIGLU_FUSION_ATTR).is_some()
|| node.attr(MATMUL_NBITS_RMSNORM_PROLOGUE_ATTR).is_some()
{
return None;
}
let value_count = node.input_values().count();
if !(value_count == 3 || value_count == 4)
|| node.inputs.iter().skip(4).any(Option::is_some)
{
return None;
}
if let Some(zero_points) = node.inputs.get(3).copied().flatten()
&& graph.try_value(zero_points).map(|value| value.dtype) != Some(DataType::Uint8)
{
return None;
}
if !self.is_fusable_bits_fp16_matmul(graph, node) {
return None;
}
if node.attr("N").and_then(Attribute::as_int)? as usize != norm_size {
return None;
}
if graph.consumers(value).len() != 1 || graph.value(value).is_graph_output {
return None;
}
Some(producer)
}
fn following_gemv_is_fusable(&self, graph: &Graph, id: NodeId) -> bool {
let Some(node) = graph.try_node(id) else {
return false;
};
if node.op_type != "MatMulNBits" || node.domain != MICROSOFT_DOMAIN {
return false;
}
if node.attr(MATMUL_NBITS_RMSNORM_PROLOGUE_ATTR).is_some() {
return false;
}
if !self.is_fusable_bits_fp16_matmul(graph, node) {
return false;
}
let (Some(k), Some(n)) = (
node.attr("K").and_then(Attribute::as_int),
node.attr("N").and_then(Attribute::as_int),
) else {
return false;
};
if node.attr(GATE_UP_SWIGLU_FUSION_ATTR).is_some() {
let value_count = node.input_values().count();
return (value_count == 5 || value_count == 7) && k <= n;
}
if let Some(zero_points) = node.inputs.get(3).copied().flatten()
&& graph.try_value(zero_points).map(|value| value.dtype) != Some(DataType::Uint8)
{
return false;
}
if node.inputs.get(4).copied().flatten().is_some()
|| node.inputs.get(6).copied().flatten().is_some()
{
return false;
}
k <= n
}
fn is_fusable_bits_fp16_matmul(&self, graph: &Graph, node: &onnx_runtime_ir::Node) -> bool {
if !RMSNORM_FUSION_SUPPORTED_BITS
.contains(&node.attr("bits").and_then(Attribute::as_int).unwrap_or(4))
{
return false;
}
if node.attr("block_size").and_then(Attribute::as_int)
!= Some(RMSNORM_FUSION_SUPPORTED_BLOCK_SIZE)
{
return false;
}
let Some(scales) = node.inputs.get(2).copied().flatten() else {
return false;
};
graph
.try_value(scales)
.is_some_and(|value| value.dtype == DataType::Float16)
}
}
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,
gate_zero_points: Option<ValueId>,
up_weight: ValueId,
up_scales: ValueId,
up_zero_points: Option<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),
];
if let (Some(gate_zp), Some(up_zp)) = (plan.gate_zero_points, plan.up_zero_points) {
fused.inputs.push(None);
fused.inputs.push(Some(gate_zp));
fused.inputs.push(Some(up_zp));
}
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;
}
if gate.zero_points.is_some() != up.zero_points.is_some() {
return None;
}
Some(GateUpSwiGluPlan {
mul_id,
gate_matmul_id,
up_matmul_id,
activation: gate.activation,
gate_weight: gate.weight,
gate_scales: gate.scales,
gate_zero_points: gate.zero_points,
up_weight: up.weight,
up_scales: up.scales,
up_zero_points: up.zero_points,
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 || present.len() == 4)
|| matmul.inputs.iter().skip(4).any(Option::is_some)
{
return None;
}
let zero_points = matmul.inputs.get(3).copied().flatten();
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;
}
if let Some(zero_points) = zero_points
&& (graph.value(zero_points).dtype != DataType::Uint8
|| !graph.initializers.contains_key(&zero_points))
{
return None;
}
Some(Projection {
activation,
weight,
scales,
zero_points,
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,
zero_points: Option<ValueId>,
n: usize,
k: usize,
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaOnDeviceConstantSelect;
struct SelectOutput {
value: ValueId,
dtype: DataType,
out_dims: Vec<usize>,
x_bytes: Vec<u8>,
y_bytes: Vec<u8>,
}
struct SelectPlan {
if_node: NodeId,
then_key: (NodeId, String),
else_key: (NodeId, String),
cond: ValueId,
name: String,
outputs: Vec<SelectOutput>,
}
impl OptimizationPass for CudaOnDeviceConstantSelect {
fn name(&self) -> &str {
"CudaOnDeviceConstantSelect"
}
fn run(&self, graph: &mut Graph, ctx: &PassContext) -> OptimizerResult<()> {
let candidates: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "If" && matches!(node.domain.as_str(), "" | "ai.onnx"))
.then_some(id)
})
.collect();
let mut plans: Vec<SelectPlan> = Vec::new();
for if_node in candidates {
if let Some(plan) = self.plan_select(graph, ctx, if_node) {
plans.push(plan);
}
}
let changed = !plans.is_empty();
for plan in plans {
graph.remove_node(plan.if_node);
graph.subgraphs.remove(&plan.then_key);
graph.subgraphs.remove(&plan.else_key);
for (index, out) in plan.outputs.into_iter().enumerate() {
if graph.try_value(out.value).is_none() {
continue;
}
let shape = static_shape(out.out_dims.clone());
let x = graph.create_value(out.dtype, shape.clone());
graph.set_initializer(
x,
WeightRef::Inline(TensorData::from_raw(
out.dtype,
out.out_dims.clone(),
out.x_bytes,
)),
);
let y = graph.create_value(out.dtype, shape.clone());
graph.set_initializer(
y,
WeightRef::Inline(TensorData::from_raw(
out.dtype,
out.out_dims.clone(),
out.y_bytes,
)),
);
let value = graph.value_mut(out.value);
value.dtype = out.dtype;
value.shape = shape;
let mut node = onnx_runtime_ir::Node::new(
NodeId(0),
"Where",
vec![Some(plan.cond), Some(x), Some(y)],
vec![out.value],
);
node.name = format!("{}/on_device_select_{index}", plan.name);
graph.insert_node(node);
}
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
impl CudaOnDeviceConstantSelect {
fn plan_select(&self, graph: &Graph, ctx: &PassContext, if_node: NodeId) -> Option<SelectPlan> {
let node = graph.try_node(if_node)?;
let cond = node.inputs.first().copied().flatten()?;
let then_key = (if_node, "then_branch".to_string());
let else_key = (if_node, "else_branch".to_string());
let then_branch = graph.subgraphs.get(&then_key)?;
let else_branch = graph.subgraphs.get(&else_key)?;
if !branch_is_pure_constants(then_branch) || !branch_is_pure_constants(else_branch) {
return None;
}
let output_count = node.outputs.len();
if output_count == 0
|| then_branch.outputs.len() != output_count
|| else_branch.outputs.len() != output_count
{
return None;
}
let threshold = greater_threshold(graph, ctx, cond);
let mut outputs = Vec::with_capacity(output_count);
for i in 0..output_count {
let then_tensor = branch_constant(then_branch, then_branch.outputs[i])?;
let else_tensor = branch_constant(else_branch, else_branch.outputs[i])?;
if then_tensor.dtype != else_tensor.dtype {
return None;
}
let dtype = then_tensor.dtype;
let elem = dtype.byte_size();
if elem == 0 || dtype.is_sub_byte() {
return None;
}
let then_dims = then_tensor.dims.clone();
let else_dims = else_tensor.dims.clone();
let then_bytes: &[u8] = &then_tensor.data;
let else_bytes: &[u8] = &else_tensor.data;
if then_bytes.len() != dims_bytes(&then_dims, elem)?
|| else_bytes.len() != dims_bytes(&else_dims, elem)?
{
return None;
}
let plan = if then_dims == else_dims {
SelectOutput {
value: node.outputs[i],
dtype,
out_dims: then_dims,
x_bytes: then_bytes.to_vec(),
y_bytes: else_bytes.to_vec(),
}
} else {
let (then_lead, then_trail) = split_leading(&then_dims)?;
let (else_lead, else_trail) = split_leading(&else_dims)?;
if then_trail != else_trail || then_lead <= else_lead {
return None;
}
let threshold = threshold?;
if i64::try_from(else_lead).ok()? != threshold {
return None;
}
let row_bytes = else_trail.iter().product::<usize>().checked_mul(elem)?;
let mut y_bytes = else_bytes.to_vec();
let pad_rows = then_lead.checked_sub(else_lead)?;
y_bytes.resize(else_bytes.len() + pad_rows.checked_mul(row_bytes)?, 0);
SelectOutput {
value: node.outputs[i],
dtype,
out_dims: then_dims,
x_bytes: then_bytes.to_vec(),
y_bytes,
}
};
outputs.push(plan);
}
Some(SelectPlan {
if_node,
then_key,
else_key,
cond,
name: node.name.clone(),
outputs,
})
}
}
fn branch_is_pure_constants(branch: &Graph) -> bool {
if !branch.inputs.is_empty() {
return false;
}
branch.nodes.iter().all(|(_, node)| {
node.op_type == "Constant"
&& matches!(node.domain.as_str(), "" | "ai.onnx")
&& node.inputs.is_empty()
&& node.outputs.len() == 1
&& matches!(node.attr("value"), Some(Attribute::Tensor(_)))
})
}
fn branch_constant(branch: &Graph, out: ValueId) -> Option<&TensorData> {
let producer = branch.try_value(out)?.producer?;
let node = branch.try_node(producer)?;
if node.op_type != "Constant" || !matches!(node.domain.as_str(), "" | "ai.onnx") {
return None;
}
match node.attr("value") {
Some(Attribute::Tensor(tensor)) => Some(tensor),
_ => None,
}
}
fn greater_threshold(graph: &Graph, ctx: &PassContext, cond: ValueId) -> Option<i64> {
let producer = graph.try_value(cond)?.producer?;
let node = graph.try_node(producer)?;
if !matches!(node.op_type.as_str(), "Greater" | "GreaterOrEqual")
|| !matches!(node.domain.as_str(), "" | "ai.onnx")
{
return None;
}
let threshold_value = node.inputs.get(1).copied().flatten()?;
scalar_int(graph, ctx, threshold_value)
}
fn scalar_int(graph: &Graph, ctx: &PassContext, value: ValueId) -> Option<i64> {
let tensor_owned;
let tensor: &TensorData = if let Some(weight) = graph.initializers.get(&value) {
match weight {
WeightRef::Inline(tensor) => tensor,
WeightRef::External { .. } => {
tensor_owned = TensorData::from_raw(
weight.dtype(),
weight.dims().to_vec(),
ctx.initializer_bytes(weight)?.to_vec(),
);
&tensor_owned
}
}
} else {
let producer = graph.try_value(value)?.producer?;
match graph.try_node(producer)?.attr("value") {
Some(Attribute::Tensor(tensor)) => tensor,
_ => return None,
}
};
if tensor.dims.iter().product::<usize>() != 1 {
return None;
}
match tensor.dtype {
DataType::Int64 => Some(i64::from_le_bytes(tensor.data.get(..8)?.try_into().ok()?)),
DataType::Int32 => Some(i64::from(i32::from_le_bytes(
tensor.data.get(..4)?.try_into().ok()?,
))),
_ => None,
}
}
fn split_leading(dims: &[usize]) -> Option<(usize, Vec<usize>)> {
let (&lead, trail) = dims.split_first()?;
Some((lead, trail.to_vec()))
}
fn dims_bytes(dims: &[usize], elem: usize) -> Option<usize> {
dims.iter().product::<usize>().checked_mul(elem)
}
#[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 const_transpose_graph(rows: usize, cols: usize, perm: Option<Vec<i64>>) -> (Graph, ValueId) {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let weight = graph.create_named_value(
"weight",
DataType::Float16,
vec![Dim::Static(rows), Dim::Static(cols)],
);
let mut bytes = Vec::with_capacity(rows * cols * 2);
for r in 0..rows {
for c in 0..cols {
let v = half::f16::from_f32((r * cols + c) as f32);
bytes.extend_from_slice(&v.to_le_bytes());
}
}
graph.set_initializer(
weight,
WeightRef::Inline(TensorData::from_raw(
DataType::Float16,
vec![rows, cols],
bytes,
)),
);
let transposed = graph.create_named_value(
"transposed",
DataType::Float16,
vec![Dim::Static(cols), Dim::Static(rows)],
);
let mut node = Node::new(NodeId(0), "Transpose", vec![Some(weight)], vec![transposed]);
if let Some(perm) = perm {
node.attributes.insert("perm".into(), Attribute::Ints(perm));
}
graph.insert_node(node);
let out = graph.create_named_value(
"out",
DataType::Float16,
vec![Dim::Static(cols), Dim::Static(rows)],
);
graph.insert_node(Node::new(
NodeId(0),
"Identity",
vec![Some(transposed)],
vec![out],
));
graph.add_output(out);
(graph, transposed)
}
fn f16_at(bytes: &[u8], index: usize) -> f32 {
half::f16::from_le_bytes([bytes[index * 2], bytes[index * 2 + 1]]).to_f32()
}
fn static_shape_of(graph: &Graph, value: ValueId) -> Vec<usize> {
onnx_runtime_ir::as_static_shape(&graph.value(value).shape).unwrap()
}
#[test]
fn folds_constant_transpose_into_initializer() {
let (mut graph, transposed) = const_transpose_graph(3, 4, Some(vec![1, 0]));
assert!(graph.value(transposed).producer.is_some());
CudaFoldConstantTranspose
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(graph.nodes.values().all(|node| node.op_type != "Transpose"));
let value = graph.value(transposed);
assert!(value.producer.is_none());
assert_eq!(static_shape_of(&graph, transposed), vec![4, 3]);
let WeightRef::Inline(tensor) = graph.initializers.get(&transposed).unwrap() else {
panic!("expected inline initializer");
};
assert_eq!(tensor.dims, vec![4, 3]);
for r in 0..3usize {
for c in 0..4usize {
assert_eq!(f16_at(&tensor.data, c * 3 + r), (r * 4 + c) as f32);
}
}
assert!(graph.validate().is_ok());
}
#[test]
fn folds_constant_transpose_default_perm() {
let (mut graph, transposed) = const_transpose_graph(2, 5, None);
CudaFoldConstantTranspose
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(graph.nodes.values().all(|node| node.op_type != "Transpose"));
assert_eq!(static_shape_of(&graph, transposed), vec![5, 2]);
}
#[test]
fn leaves_transpose_of_non_constant() {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let input =
graph.create_named_value("x", DataType::Float16, vec![Dim::Static(3), Dim::Static(4)]);
graph.add_input(input);
let out =
graph.create_named_value("y", DataType::Float16, vec![Dim::Static(4), Dim::Static(3)]);
let mut node = Node::new(NodeId(0), "Transpose", vec![Some(input)], vec![out]);
node.attributes
.insert("perm".into(), Attribute::Ints(vec![1, 0]));
graph.insert_node(node);
graph.add_output(out);
CudaFoldConstantTranspose
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph
.nodes
.values()
.filter(|n| n.op_type == "Transpose")
.count(),
1
);
assert!(!graph.initializers.contains_key(&out));
}
#[test]
fn leaves_sub_byte_constant_transpose() {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let weight =
graph.create_named_value("w", DataType::Int4, vec![Dim::Static(4), Dim::Static(4)]);
graph.set_initializer(
weight,
WeightRef::Inline(TensorData::from_raw(
DataType::Int4,
vec![4, 4],
vec![0u8; 8],
)),
);
let out =
graph.create_named_value("wt", DataType::Int4, vec![Dim::Static(4), Dim::Static(4)]);
let mut node = Node::new(NodeId(0), "Transpose", vec![Some(weight)], vec![out]);
node.attributes
.insert("perm".into(), Attribute::Ints(vec![1, 0]));
graph.insert_node(node);
let consumer_out =
graph.create_named_value("o", DataType::Int4, vec![Dim::Static(4), Dim::Static(4)]);
graph.insert_node(Node::new(
NodeId(0),
"Identity",
vec![Some(out)],
vec![consumer_out],
));
graph.add_output(consumer_out);
CudaFoldConstantTranspose
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph
.nodes
.values()
.filter(|n| n.op_type == "Transpose")
.count(),
1,
"sub-byte Transpose must be left intact"
);
}
#[test]
fn folds_rank3_constant_transpose() {
let (rows, mid, cols) = (2usize, 3usize, 4usize);
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let weight = graph.create_named_value(
"w",
DataType::Float16,
vec![Dim::Static(rows), Dim::Static(mid), Dim::Static(cols)],
);
let mut bytes = Vec::new();
for i in 0..rows * mid * cols {
bytes.extend_from_slice(&half::f16::from_f32(i as f32).to_le_bytes());
}
graph.set_initializer(
weight,
WeightRef::Inline(TensorData::from_raw(
DataType::Float16,
vec![rows, mid, cols],
bytes,
)),
);
let out = graph.create_named_value(
"wt",
DataType::Float16,
vec![Dim::Static(cols), Dim::Static(rows), Dim::Static(mid)],
);
let mut node = Node::new(NodeId(0), "Transpose", vec![Some(weight)], vec![out]);
node.attributes
.insert("perm".into(), Attribute::Ints(vec![2, 0, 1]));
graph.insert_node(node);
let consumer_out = graph.create_named_value(
"o",
DataType::Float16,
vec![Dim::Static(cols), Dim::Static(rows), Dim::Static(mid)],
);
graph.insert_node(Node::new(
NodeId(0),
"Identity",
vec![Some(out)],
vec![consumer_out],
));
CudaFoldConstantTranspose
.run(&mut graph, &PassContext::new())
.unwrap();
let WeightRef::Inline(tensor) = graph.initializers.get(&out).unwrap() else {
panic!("expected inline initializer");
};
assert_eq!(tensor.dims, vec![cols, rows, mid]);
for c in 0..cols {
for r in 0..rows {
for m in 0..mid {
let out_flat = (c * rows + r) * mid + m;
let expected = (r * mid + m) * cols + c;
assert_eq!(f16_at(&tensor.data, out_flat), expected as f32);
}
}
}
}
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());
}
fn skip_rms_graph(norm_size: usize, following_n: usize) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let pre_x = value(&mut graph, "pre_x", DataType::Float16, norm_size + 128);
graph.add_input(pre_x);
let pre_out = projection(&mut graph, "pre", pre_x, norm_size + 128, norm_size);
let residual = value(&mut graph, "residual", DataType::Float16, norm_size);
graph.add_input(residual);
let gamma = vec1d(&mut graph, "gamma", DataType::Float16, norm_size);
graph.set_initializer(
gamma,
WeightRef::Inline(TensorData::from_raw(
DataType::Float16,
vec![norm_size],
vec![0u8; norm_size * 2],
)),
);
let normalized = value(&mut graph, "normalized", DataType::Float16, norm_size);
let stat_mean = graph.create_value(DataType::Float32, Vec::new());
let stat_inv_std = graph.create_value(DataType::Float32, Vec::new());
let sum = value(&mut graph, "sum", DataType::Float16, norm_size);
let mut skip = Node::new(
NodeId(0),
"SkipSimplifiedLayerNormalization",
vec![Some(pre_out), Some(residual), Some(gamma)],
vec![normalized, stat_mean, stat_inv_std, sum],
);
skip.domain = MICROSOFT_DOMAIN.into();
skip.attributes
.insert("epsilon".into(), Attribute::Float(9.999_999e-7));
graph.insert_node(skip);
let post_out = projection(&mut graph, "post", normalized, norm_size, following_n);
graph.add_output(post_out);
let sum_sink = value(&mut graph, "sum_sink", DataType::Float16, norm_size);
graph.insert_node(Node::new(
NodeId(0),
"Identity",
vec![Some(sum)],
vec![sum_sink],
));
graph.add_output(sum_sink);
graph
}
fn value_id_by_name(graph: &Graph, name: &str) -> ValueId {
graph
.values
.iter()
.find_map(|(id, v)| (v.name.as_deref() == Some(name)).then_some(id))
.unwrap()
}
fn node_producing(graph: &Graph, output: ValueId) -> &Node {
let producer = graph.value(output).producer.unwrap();
graph.node(producer)
}
#[test]
fn folds_skip_rmsnorm_into_neighbouring_gemvs() {
let mut graph = skip_rms_graph(RMSNORM_FUSION_MIN_HIDDEN, RMSNORM_FUSION_MIN_HIDDEN + 256);
let pre_out = value_id_by_name(&graph, "pre_out");
let residual = value_id_by_name(&graph, "residual");
let gamma = value_id_by_name(&graph, "gamma");
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.all(|node| node.op_type != "SkipSimplifiedLayerNormalization"),
"the SkipSimplifiedLayerNormalization must be deleted"
);
let preceding = node_producing(&graph, pre_out);
assert_eq!(preceding.op_type, "MatMulNBits");
assert_eq!(
preceding
.attr(MATMUL_NBITS_FOLDED_BIAS_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert_eq!(preceding.inputs.len(), 6);
assert_eq!(preceding.inputs[5], Some(residual), "residual at slot 5");
assert!(preceding.inputs[3].is_none() && preceding.inputs[4].is_none());
let post_out = value_id_by_name(&graph, "post_out");
let following = node_producing(&graph, post_out);
assert_eq!(
following.inputs[0],
Some(pre_out),
"activation is residual sum"
);
assert_eq!(following.inputs.get(6).copied().flatten(), Some(gamma));
assert_eq!(
following
.attr(MATMUL_NBITS_RMSNORM_PROLOGUE_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert_eq!(
following
.attr(MATMUL_NBITS_RMSNORM_EPSILON_ATTR)
.and_then(Attribute::as_float),
Some(9.999_999e-7)
);
let identity = graph
.nodes
.values()
.find(|node| node.op_type == "Identity")
.unwrap();
assert_eq!(identity.inputs[0], Some(pre_out));
assert!(graph.validate().is_ok());
}
#[test]
fn folds_skip_rmsnorm_into_int8_neighbouring_gemvs() {
let mut graph = skip_rms_graph(RMSNORM_FUSION_MIN_HIDDEN, RMSNORM_FUSION_MIN_HIDDEN + 256);
set_matmul_attr(&mut graph, "bits", 8);
let pre_out = value_id_by_name(&graph, "pre_out");
let gamma = value_id_by_name(&graph, "gamma");
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.all(|node| node.op_type != "SkipSimplifiedLayerNormalization"),
"int8 skip-rmsnorm must fuse just like int4"
);
let preceding = node_producing(&graph, pre_out);
assert_eq!(
preceding
.attr(MATMUL_NBITS_FOLDED_BIAS_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
let post_out = value_id_by_name(&graph, "post_out");
let following = node_producing(&graph, post_out);
assert_eq!(following.inputs.get(6).copied().flatten(), Some(gamma));
assert_eq!(
following
.attr(MATMUL_NBITS_RMSNORM_PROLOGUE_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert!(graph.validate().is_ok());
}
#[test]
fn leaves_skip_rmsnorm_when_hidden_not_multiple_of_128() {
let mut graph = skip_rms_graph(1288, RMSNORM_FUSION_MIN_HIDDEN + 256);
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.any(|node| node.op_type == "SkipSimplifiedLayerNormalization"),
"an unaligned hidden size must keep the standalone norm"
);
}
#[test]
fn leaves_skip_rmsnorm_when_following_is_down_variant() {
let mut graph = skip_rms_graph(RMSNORM_FUSION_MIN_HIDDEN, RMSNORM_FUSION_MIN_HIDDEN - 128);
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.any(|node| node.op_type == "SkipSimplifiedLayerNormalization"),
"a down-variant following GEMV must block the fusion"
);
}
#[test]
fn leaves_skip_rmsnorm_with_norm_bias() {
let hidden = RMSNORM_FUSION_MIN_HIDDEN;
let mut graph = skip_rms_graph(hidden, hidden + 256);
let bias = vec1d(&mut graph, "norm_bias", DataType::Float16, hidden);
graph.set_initializer(
bias,
WeightRef::Inline(TensorData::from_raw(
DataType::Float16,
vec![hidden],
vec![0u8; hidden * 2],
)),
);
let skip_id = graph
.nodes
.iter()
.find_map(|(id, n)| (n.op_type == "SkipSimplifiedLayerNormalization").then_some(id))
.unwrap();
let mut skip = graph.node(skip_id).clone();
skip.inputs.push(Some(bias));
graph.replace_node(skip_id, skip);
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.any(|node| node.op_type == "SkipSimplifiedLayerNormalization"),
"a norm bias must block the fusion"
);
}
#[test]
fn leaves_skip_rmsnorm_when_preceding_gemv_shared() {
let hidden = RMSNORM_FUSION_MIN_HIDDEN;
let mut graph = skip_rms_graph(hidden, hidden + 256);
let pre_out = value_id_by_name(&graph, "pre_out");
let sink = value(&mut graph, "pre_sink", DataType::Float16, hidden);
graph.insert_node(Node::new(NodeId(0), "Neg", vec![Some(pre_out)], vec![sink]));
graph.add_output(sink);
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.any(|node| node.op_type == "SkipSimplifiedLayerNormalization"),
"a shared preceding GEMV must block the fusion"
);
}
#[test]
fn leaves_skip_rmsnorm_when_broadcast_skip() {
let hidden = RMSNORM_FUSION_MIN_HIDDEN;
let mut graph = skip_rms_graph(hidden, hidden + 256);
let residual = value_id_by_name(&graph, "residual");
graph.value_mut(residual).shape = vec![Dim::Static(1), Dim::Static(1)];
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.any(|node| node.op_type == "SkipSimplifiedLayerNormalization"),
"a broadcast skip must block the fusion"
);
}
#[test]
fn folds_skip_rmsnorm_with_symbolic_batch_and_sequence_dims() {
let mut graph = skip_rms_graph(RMSNORM_FUSION_MIN_HIDDEN, RMSNORM_FUSION_MIN_HIDDEN + 256);
let batch = graph.create_symbol(Some("batch".into()));
let sequence = graph.create_symbol(Some("sequence".into()));
let symbolic = vec![
Dim::Symbolic(batch),
Dim::Symbolic(sequence),
Dim::Static(RMSNORM_FUSION_MIN_HIDDEN),
];
for name in ["pre_out", "residual", "normalized", "sum"] {
let id = value_id_by_name(&graph, name);
graph.value_mut(id).shape = symbolic.clone();
}
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.all(|node| node.op_type != "SkipSimplifiedLayerNormalization"),
"symbolic batch/sequence dims must not block the fusion"
);
let post_out = value_id_by_name(&graph, "post_out");
let following = node_producing(&graph, post_out);
assert_eq!(
following
.attr(MATMUL_NBITS_RMSNORM_PROLOGUE_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert!(graph.validate().is_ok());
}
#[test]
fn skip_rmsnorm_fires_through_full_cuda_pass_list() {
let mut graph = skip_rms_graph(RMSNORM_FUSION_MIN_HIDDEN, RMSNORM_FUSION_MIN_HIDDEN + 256);
for pass in cuda_optimization_passes() {
pass.run(&mut graph, &PassContext::new()).unwrap();
}
assert!(
graph
.nodes
.values()
.all(|node| node.op_type != "SkipSimplifiedLayerNormalization"),
"the fusion must fire through the full pass list"
);
assert!(graph.validate().is_ok());
}
#[test]
fn gate_leaves_skip_rmsnorm_below_hidden_floor() {
let hidden = RMSNORM_FUSION_MIN_HIDDEN - RMSNORM_FUSION_WARP_HALF4_MULTIPLE;
let mut graph = skip_rms_graph(hidden, hidden + 256);
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.any(|node| node.op_type == "SkipSimplifiedLayerNormalization"),
"a hidden below the size floor must keep the standalone norm"
);
}
#[test]
fn gate_folds_skip_rmsnorm_at_hidden_floor() {
let hidden = RMSNORM_FUSION_MIN_HIDDEN;
let mut graph = skip_rms_graph(hidden, hidden + 256);
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.all(|node| node.op_type != "SkipSimplifiedLayerNormalization"),
"a hidden at the size floor must fuse"
);
assert!(graph.validate().is_ok());
}
fn post_attention_swiglu_graph(hidden: usize, intermediate: usize) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let pre_x = value(&mut graph, "pre_x", DataType::Float16, hidden + 128);
graph.add_input(pre_x);
let pre_out = projection(&mut graph, "pre", pre_x, hidden + 128, hidden);
let residual = value(&mut graph, "residual", DataType::Float16, hidden);
graph.add_input(residual);
let gamma = vec1d(&mut graph, "gamma", DataType::Float16, hidden);
graph.set_initializer(
gamma,
WeightRef::Inline(TensorData::from_raw(
DataType::Float16,
vec![hidden],
vec![0u8; hidden * 2],
)),
);
let normalized = value(&mut graph, "normalized", DataType::Float16, hidden);
let stat_mean = graph.create_value(DataType::Float32, Vec::new());
let stat_inv_std = graph.create_value(DataType::Float32, Vec::new());
let sum = value(&mut graph, "sum", DataType::Float16, hidden);
let mut skip = Node::new(
NodeId(0),
"SkipSimplifiedLayerNormalization",
vec![Some(pre_out), Some(residual), Some(gamma)],
vec![normalized, stat_mean, stat_inv_std, sum],
);
skip.domain = MICROSOFT_DOMAIN.into();
skip.attributes
.insert("epsilon".into(), Attribute::Float(9.999_999e-7));
graph.insert_node(skip);
let gate_out = projection(&mut graph, "gate", normalized, hidden, intermediate);
let up_out = projection(&mut graph, "up", normalized, hidden, intermediate);
let silu_out = value(&mut graph, "silu", DataType::Float16, intermediate);
let out = value(&mut graph, "output", DataType::Float16, intermediate);
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);
let sum_sink = value(&mut graph, "sum_sink", DataType::Float16, hidden);
graph.insert_node(Node::new(
NodeId(0),
"Identity",
vec![Some(sum)],
vec![sum_sink],
));
graph.add_output(sum_sink);
graph
}
#[test]
fn folds_skip_rmsnorm_into_gate_up_swiglu_node() {
let hidden = RMSNORM_FUSION_MIN_HIDDEN;
let intermediate = hidden * 2;
let mut graph = post_attention_swiglu_graph(hidden, intermediate);
for pass in cuda_optimization_passes() {
pass.run(&mut graph, &PassContext::new()).unwrap();
}
assert!(
graph
.nodes
.values()
.all(|node| node.op_type != "SkipSimplifiedLayerNormalization"),
"the standalone norm must be folded away"
);
let swiglu = graph
.nodes
.values()
.find(|node| node.attr(GATE_UP_SWIGLU_FUSION_ATTR).is_some())
.expect("a fused gate/up SwiGLU node must exist");
let pre_out = value_id_by_name(&graph, "pre_out");
let gamma = value_id_by_name(&graph, "gamma");
assert_eq!(
swiglu.inputs[0],
Some(pre_out),
"activation is the preceding residual sum"
);
assert_eq!(swiglu.inputs.len(), 6, "gamma appended at slot 5");
assert_eq!(swiglu.inputs.get(5).copied().flatten(), Some(gamma));
assert_eq!(
swiglu
.attr(MATMUL_NBITS_RMSNORM_PROLOGUE_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert_eq!(
swiglu
.attr(MATMUL_NBITS_RMSNORM_EPSILON_ATTR)
.and_then(Attribute::as_float),
Some(9.999_999e-7)
);
assert!(graph.validate().is_ok());
}
#[test]
fn gate_up_swiglu_node_stays_unfused_below_hidden_floor() {
let hidden = RMSNORM_FUSION_MIN_HIDDEN - RMSNORM_FUSION_WARP_HALF4_MULTIPLE;
let intermediate = hidden * 2;
let mut graph = post_attention_swiglu_graph(hidden, intermediate);
for pass in cuda_optimization_passes() {
pass.run(&mut graph, &PassContext::new()).unwrap();
}
assert!(
graph
.nodes
.values()
.any(|node| node.op_type == "SkipSimplifiedLayerNormalization"),
"below the floor the standalone norm must survive"
);
let swiglu = graph
.nodes
.values()
.find(|node| node.attr(GATE_UP_SWIGLU_FUSION_ATTR).is_some())
.expect("the gate/up pair still fuses on its own");
assert!(
swiglu.attr(MATMUL_NBITS_RMSNORM_PROLOGUE_ATTR).is_none(),
"the SwiGLU node must not carry an RMS prologue below the floor"
);
}
#[test]
fn folds_chained_blocks_sharing_residual_sum() {
let hidden = RMSNORM_FUSION_MIN_HIDDEN;
let following_n = hidden + 256;
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let res_in = value(&mut graph, "res_in", DataType::Float16, hidden);
graph.add_input(res_in);
let gamma0 = vec1d(&mut graph, "gamma0", DataType::Float16, hidden);
let gamma1 = vec1d(&mut graph, "gamma1", DataType::Float16, hidden);
for g in [gamma0, gamma1] {
graph.set_initializer(
g,
WeightRef::Inline(TensorData::from_raw(
DataType::Float16,
vec![hidden],
vec![0u8; hidden * 2],
)),
);
}
let x0 = value(&mut graph, "x0", DataType::Float16, hidden + 128);
graph.add_input(x0);
let pre0 = projection(&mut graph, "pre0", x0, hidden + 128, hidden);
let norm0 = value(&mut graph, "norm0", DataType::Float16, hidden);
let mean0 = graph.create_value(DataType::Float32, Vec::new());
let inv0 = graph.create_value(DataType::Float32, Vec::new());
let sum0 = value(&mut graph, "sum0", DataType::Float16, hidden);
let mut skip0 = Node::new(
NodeId(0),
"SkipSimplifiedLayerNormalization",
vec![Some(pre0), Some(res_in), Some(gamma0)],
vec![norm0, mean0, inv0, sum0],
);
skip0.domain = MICROSOFT_DOMAIN.into();
skip0
.attributes
.insert("epsilon".into(), Attribute::Float(1e-6));
graph.insert_node(skip0);
let post0 = projection(&mut graph, "post0", norm0, hidden, following_n);
graph.add_output(post0);
let x1 = value(&mut graph, "x1", DataType::Float16, hidden + 128);
graph.add_input(x1);
let pre1 = projection(&mut graph, "pre1", x1, hidden + 128, hidden);
let norm1 = value(&mut graph, "norm1", DataType::Float16, hidden);
let mean1 = graph.create_value(DataType::Float32, Vec::new());
let inv1 = graph.create_value(DataType::Float32, Vec::new());
let sum1 = value(&mut graph, "sum1", DataType::Float16, hidden);
let mut skip1 = Node::new(
NodeId(0),
"SkipSimplifiedLayerNormalization",
vec![Some(pre1), Some(sum0), Some(gamma1)],
vec![norm1, mean1, inv1, sum1],
);
skip1.domain = MICROSOFT_DOMAIN.into();
skip1
.attributes
.insert("epsilon".into(), Attribute::Float(1e-6));
graph.insert_node(skip1);
let post1 = projection(&mut graph, "post1", norm1, hidden, following_n);
graph.add_output(post1);
let sum1_sink = value(&mut graph, "sum1_sink", DataType::Float16, hidden);
graph.insert_node(Node::new(
NodeId(0),
"Identity",
vec![Some(sum1)],
vec![sum1_sink],
));
graph.add_output(sum1_sink);
CudaSkipRmsNormMatMulFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.all(|node| node.op_type != "SkipSimplifiedLayerNormalization"),
"both chained norms must fold"
);
assert!(
graph.validate().is_ok(),
"no dangling shared residual value"
);
let pre0_out = value_id_by_name(&graph, "pre0_out");
let pre1_out = value_id_by_name(&graph, "pre1_out");
let pre1_node = node_producing(&graph, pre1_out);
assert_eq!(
pre1_node.inputs.get(5).copied().flatten(),
Some(pre0_out),
"block 1 folds the redirected residual (block 0's preceding output)"
);
}
fn cast_node(
graph: &mut Graph,
name: &str,
input: ValueId,
to: DataType,
width: usize,
) -> ValueId {
let out = value(graph, name, to, width);
let mut node = Node::new(NodeId(0), "Cast", vec![Some(input)], vec![out]);
node.attributes
.insert("to".into(), Attribute::Int(to as i64));
graph.insert_node(node);
out
}
fn cast_wrapped_skip_norm_graph(hidden: usize) -> (Graph, ValueId, ValueId, ValueId) {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let residual = value(&mut graph, "residual", DataType::Float16, hidden);
let mm = value(&mut graph, "mm", DataType::Float16, hidden);
graph.add_input(residual);
graph.add_input(mm);
let gamma = vec1d(&mut graph, "gamma", DataType::Float32, hidden);
graph.set_initializer(
gamma,
WeightRef::Inline(TensorData::from_raw(
DataType::Float32,
vec![hidden],
vec![0u8; hidden * 4],
)),
);
let in0 = cast_node(&mut graph, "in0", residual, DataType::Float32, hidden);
let in1 = cast_node(&mut graph, "in1", mm, DataType::Float32, hidden);
let norm_out = value(&mut graph, "norm_out", DataType::Float32, hidden);
let sum_out = value(&mut graph, "sum_out", DataType::Float32, hidden);
let mut skip = Node::new(
NodeId(0),
"SkipSimplifiedLayerNormalization",
vec![Some(in0), Some(in1), Some(gamma)],
vec![norm_out, sum_out],
);
skip.domain = MICROSOFT_DOMAIN.into();
skip.attributes
.insert("epsilon".into(), Attribute::Float(1e-5));
graph.insert_node(skip);
let normalized = cast_node(
&mut graph,
"normalized",
norm_out,
DataType::Float16,
hidden,
);
let residual_out = cast_node(
&mut graph,
"residual_out",
sum_out,
DataType::Float16,
hidden,
);
graph.add_output(normalized);
graph.add_output(residual_out);
(graph, residual, mm, gamma)
}
#[test]
fn drops_casts_around_fp32_wrapped_skip_norm() {
let hidden = 128;
let (mut graph, residual, mm, gamma) = cast_wrapped_skip_norm_graph(hidden);
CudaDropNormalizationCasts
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph.nodes.values().filter(|n| n.op_type == "Cast").count(),
0,
"all cast wrappers must be removed"
);
let skip = graph
.nodes
.values()
.find(|n| n.op_type == "SkipSimplifiedLayerNormalization")
.expect("norm node retained");
assert_eq!(
skip.inputs,
vec![Some(residual), Some(mm), Some(gamma)],
"activation inputs rewired to fp16 sources; gamma untouched"
);
for &out in &skip.outputs {
assert_eq!(graph.value(out).dtype, DataType::Float16);
}
assert_eq!(graph.value(gamma).dtype, DataType::Float32);
assert_eq!(graph.outputs.len(), 2);
for &out in &graph.outputs {
assert_eq!(graph.value(out).dtype, DataType::Float16);
}
assert!(graph.validate().is_ok());
}
#[test]
fn leaves_native_fp16_skip_norm_untouched() {
let hidden = 128;
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let residual = value(&mut graph, "residual", DataType::Float16, hidden);
let mm = value(&mut graph, "mm", DataType::Float16, hidden);
graph.add_input(residual);
graph.add_input(mm);
let gamma = vec1d(&mut graph, "gamma", DataType::Float16, hidden);
graph.set_initializer(
gamma,
WeightRef::Inline(TensorData::from_raw(
DataType::Float16,
vec![hidden],
vec![0u8; hidden * 2],
)),
);
let norm_out = value(&mut graph, "norm_out", DataType::Float16, hidden);
let mut skip = Node::new(
NodeId(0),
"SkipSimplifiedLayerNormalization",
vec![Some(residual), Some(mm), Some(gamma)],
vec![norm_out],
);
skip.domain = MICROSOFT_DOMAIN.into();
graph.insert_node(skip);
graph.add_output(norm_out);
let before = graph.nodes.len();
CudaDropNormalizationCasts
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.nodes.len(), before, "no nodes added or removed");
let skip = graph
.nodes
.values()
.find(|n| n.op_type == "SkipSimplifiedLayerNormalization")
.expect("norm retained");
assert_eq!(skip.inputs, vec![Some(residual), Some(mm), Some(gamma)]);
assert_eq!(graph.value(norm_out).dtype, DataType::Float16);
}
#[test]
fn norm_cast_fold_fires_through_full_cuda_pass_list() {
let (mut graph, ..) = cast_wrapped_skip_norm_graph(128);
for pass in cuda_optimization_passes() {
pass.run(&mut graph, &PassContext::new()).unwrap();
}
assert_eq!(
graph.nodes.values().filter(|n| n.op_type == "Cast").count(),
0,
"cast wrappers must be gone after the full pass list"
);
assert!(graph.validate().is_ok());
}
#[test]
fn leaves_cast_wrapped_norm_with_fp32_consumer_intact() {
let hidden = 128;
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let residual = value(&mut graph, "residual", DataType::Float16, hidden);
let mm = value(&mut graph, "mm", DataType::Float16, hidden);
graph.add_input(residual);
graph.add_input(mm);
let gamma = vec1d(&mut graph, "gamma", DataType::Float32, hidden);
graph.set_initializer(
gamma,
WeightRef::Inline(TensorData::from_raw(
DataType::Float32,
vec![hidden],
vec![0u8; hidden * 4],
)),
);
let in0 = cast_node(&mut graph, "in0", residual, DataType::Float32, hidden);
let in1 = cast_node(&mut graph, "in1", mm, DataType::Float32, hidden);
let norm_out = value(&mut graph, "norm_out", DataType::Float32, hidden);
let mut skip = Node::new(
NodeId(0),
"SkipSimplifiedLayerNormalization",
vec![Some(in0), Some(in1), Some(gamma)],
vec![norm_out],
);
skip.domain = MICROSOFT_DOMAIN.into();
graph.insert_node(skip);
let normalized = cast_node(
&mut graph,
"normalized",
norm_out,
DataType::Float16,
hidden,
);
graph.add_output(normalized);
let fp32_kept = value(&mut graph, "fp32_kept", DataType::Float32, hidden);
graph.insert_node(Node::new(
NodeId(0),
"Identity",
vec![Some(norm_out)],
vec![fp32_kept],
));
graph.add_output(fp32_kept);
let casts_before = graph.nodes.values().filter(|n| n.op_type == "Cast").count();
CudaDropNormalizationCasts
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph.nodes.values().filter(|n| n.op_type == "Cast").count(),
casts_before,
"a fp32 boundary consumer must block the fold"
);
let skip = graph
.nodes
.values()
.find(|n| n.op_type == "SkipSimplifiedLayerNormalization")
.expect("norm retained");
assert_eq!(skip.inputs, vec![Some(in0), Some(in1), Some(gamma)]);
assert_eq!(graph.value(norm_out).dtype, DataType::Float32);
assert!(graph.validate().is_ok());
}
#[test]
fn folds_norms_sharing_an_input_cast() {
let hidden = 128;
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, hidden);
let res_a = value(&mut graph, "res_a", DataType::Float16, hidden);
let res_b = value(&mut graph, "res_b", DataType::Float16, hidden);
graph.add_input(x);
graph.add_input(res_a);
graph.add_input(res_b);
let gamma_a = vec1d(&mut graph, "gamma_a", DataType::Float32, hidden);
let gamma_b = vec1d(&mut graph, "gamma_b", DataType::Float32, hidden);
for g in [gamma_a, gamma_b] {
graph.set_initializer(
g,
WeightRef::Inline(TensorData::from_raw(
DataType::Float32,
vec![hidden],
vec![0u8; hidden * 4],
)),
);
}
let shared = cast_node(&mut graph, "shared", x, DataType::Float32, hidden);
let skip_a_in = cast_node(&mut graph, "skip_a_in", res_a, DataType::Float32, hidden);
let skip_b_in = cast_node(&mut graph, "skip_b_in", res_b, DataType::Float32, hidden);
let norm_a = value(&mut graph, "norm_a", DataType::Float32, hidden);
let mut skip_a = Node::new(
NodeId(0),
"SkipSimplifiedLayerNormalization",
vec![Some(shared), Some(skip_a_in), Some(gamma_a)],
vec![norm_a],
);
skip_a.domain = MICROSOFT_DOMAIN.into();
graph.insert_node(skip_a);
let norm_b = value(&mut graph, "norm_b", DataType::Float32, hidden);
let mut skip_b = Node::new(
NodeId(0),
"SkipSimplifiedLayerNormalization",
vec![Some(shared), Some(skip_b_in), Some(gamma_b)],
vec![norm_b],
);
skip_b.domain = MICROSOFT_DOMAIN.into();
graph.insert_node(skip_b);
let out_a = cast_node(&mut graph, "out_a", norm_a, DataType::Float16, hidden);
let out_b = cast_node(&mut graph, "out_b", norm_b, DataType::Float16, hidden);
graph.add_output(out_a);
graph.add_output(out_b);
CudaDropNormalizationCasts
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph.nodes.values().filter(|n| n.op_type == "Cast").count(),
0,
"the shared input cast and all wrappers must be removed once unused"
);
let norms: Vec<_> = graph
.nodes
.values()
.filter(|n| n.op_type == "SkipSimplifiedLayerNormalization")
.collect();
assert_eq!(norms.len(), 2, "both norms retained");
for norm in norms {
assert_eq!(
norm.inputs[0],
Some(x),
"both norms read the shared pre-cast fp16 source"
);
}
assert_eq!(graph.value(norm_a).dtype, DataType::Float16);
assert_eq!(graph.value(norm_b).dtype, DataType::Float16);
assert!(graph.validate().is_ok());
}
fn fp16_bytes(rows: usize, cols: usize, fill: u16) -> Vec<u8> {
(0..rows * cols).flat_map(|_| fill.to_le_bytes()).collect()
}
const THEN_COS_SENTINEL: u16 = 0x3C00; const THEN_SIN_SENTINEL: u16 = 0x4000; const ELSE_COS_SENTINEL: u16 = 0x4200; const ELSE_SIN_SENTINEL: u16 = 0x4400;
fn constant_branch(outputs: &[(&str, DataType, Vec<usize>, Vec<u8>)]) -> Graph {
let mut branch = Graph::new();
for (name, dtype, dims, bytes) in outputs {
let out = branch.create_named_value(
*name,
*dtype,
dims.iter().map(|&d| Dim::Static(d)).collect::<Vec<_>>(),
);
branch.add_output(out);
let mut node = Node::new(NodeId(0), "Constant", vec![], vec![out]);
node.attributes.insert(
"value".into(),
Attribute::Tensor(TensorData::from_raw(*dtype, dims.clone(), bytes.clone())),
);
branch.insert_node(node);
}
branch
}
fn longrope_if_graph(
then_dims: Vec<usize>,
else_dims: Vec<usize>,
threshold: i64,
cond_from_greater: bool,
) -> (Graph, ValueId, ValueId) {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let cond = if cond_from_greater {
let seq = graph.create_named_value("seq_len", DataType::Int64, Vec::new());
graph.add_input(seq);
let thr = graph.create_named_value("threshold", DataType::Int64, Vec::new());
graph.set_initializer(
thr,
WeightRef::Inline(TensorData::from_raw(
DataType::Int64,
Vec::new(),
threshold.to_le_bytes().to_vec(),
)),
);
let cond = graph.create_named_value("cond", DataType::Bool, Vec::new());
graph.insert_node(Node::new(
NodeId(0),
"Greater",
vec![Some(seq), Some(thr)],
vec![cond],
));
cond
} else {
let cond = graph.create_named_value("cond", DataType::Bool, Vec::new());
graph.add_input(cond);
cond
};
let cos = graph.create_named_value(
"cos_cache",
DataType::Float16,
then_dims
.iter()
.map(|&d| Dim::Static(d))
.collect::<Vec<_>>(),
);
let sin = graph.create_named_value(
"sin_cache",
DataType::Float16,
then_dims
.iter()
.map(|&d| Dim::Static(d))
.collect::<Vec<_>>(),
);
let if_node =
graph.insert_node(Node::new(NodeId(0), "If", vec![Some(cond)], vec![cos, sin]));
graph.add_output(cos);
graph.add_output(sin);
let then_cos = fp16_bytes(then_dims[0], then_dims[1], THEN_COS_SENTINEL);
let then_sin = fp16_bytes(then_dims[0], then_dims[1], THEN_SIN_SENTINEL);
let else_cos = fp16_bytes(else_dims[0], else_dims[1], ELSE_COS_SENTINEL);
let else_sin = fp16_bytes(else_dims[0], else_dims[1], ELSE_SIN_SENTINEL);
graph.subgraphs.insert(
(if_node, "then_branch".into()),
constant_branch(&[
("cos_large", DataType::Float16, then_dims.clone(), then_cos),
("sin_large", DataType::Float16, then_dims.clone(), then_sin),
]),
);
graph.subgraphs.insert(
(if_node, "else_branch".into()),
constant_branch(&[
("cos_small", DataType::Float16, else_dims.clone(), else_cos),
("sin_small", DataType::Float16, else_dims.clone(), else_sin),
]),
);
(graph, cos, sin)
}
fn where_nodes(graph: &Graph) -> Vec<NodeId> {
graph
.nodes
.iter()
.filter_map(|(id, n)| (n.op_type == "Where").then_some(id))
.collect()
}
#[test]
fn lowers_differing_shape_longrope_if_to_padded_where() {
let (mut graph, cos, sin) = longrope_if_graph(vec![4, 2], vec![2, 2], 2, true);
CudaOnDeviceConstantSelect
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph.nodes.values().all(|n| n.op_type != "If"),
"the host If must be gone"
);
assert!(graph.subgraphs.is_empty(), "If subgraphs must be removed");
let wheres = where_nodes(&graph);
assert_eq!(wheres.len(), 2, "one Where per If output");
assert_eq!(graph.value(cos).shape, static_shape([4, 2]));
assert_eq!(graph.value(sin).shape, static_shape([4, 2]));
for &w in &wheres {
let node = graph.node(w);
let x = node.inputs[1].unwrap();
let y = node.inputs[2].unwrap();
let x_bytes = match graph.initializers.get(&x).unwrap() {
WeightRef::Inline(t) => &t.data,
_ => panic!("x must be inline"),
};
let y_bytes = match graph.initializers.get(&y).unwrap() {
WeightRef::Inline(t) => &t.data,
_ => panic!("y must be inline"),
};
let (then_sentinel, else_sentinel) = match node.outputs[0] {
output if output == cos => (THEN_COS_SENTINEL, ELSE_COS_SENTINEL),
output if output == sin => (THEN_SIN_SENTINEL, ELSE_SIN_SENTINEL),
output => panic!("unexpected Where output {output:?}"),
};
assert_eq!(
x_bytes.len(),
4 * 2 * 2,
"then const is the full long table"
);
assert_eq!(
y_bytes.len(),
4 * 2 * 2,
"else const padded to the long shape"
);
assert_eq!(
x_bytes,
&fp16_bytes(4, 2, then_sentinel),
"x must carry the predicate-true branch values"
);
assert_eq!(
&y_bytes[..2 * 2 * 2],
fp16_bytes(2, 2, else_sentinel),
"y must carry the predicate-false branch values"
);
assert!(
y_bytes[2 * 2 * 2..].iter().all(|&b| b == 0),
"appended rows are zero padding"
);
}
assert!(graph.validate().is_ok());
}
#[test]
fn lowers_equal_shape_if_without_threshold() {
let (mut graph, _cos, _sin) = longrope_if_graph(vec![3, 2], vec![3, 2], 0, false);
CudaOnDeviceConstantSelect
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(graph.nodes.values().all(|n| n.op_type != "If"));
let wheres = where_nodes(&graph);
assert_eq!(wheres.len(), 2);
for &w in &wheres {
let node = graph.node(w);
let y = node.inputs[2].unwrap();
let y_bytes = match graph.initializers.get(&y).unwrap() {
WeightRef::Inline(t) => &t.data,
_ => panic!(),
};
assert_eq!(y_bytes.len(), 3 * 2 * 2, "no padding for equal shapes");
}
assert!(graph.validate().is_ok());
}
#[test]
fn skips_if_with_non_constant_branch() {
let (mut graph, _cos, _sin) = longrope_if_graph(vec![4, 2], vec![2, 2], 2, true);
let key = (
graph
.nodes
.iter()
.find_map(|(id, n)| (n.op_type == "If").then_some(id))
.unwrap(),
"then_branch".to_string(),
);
let branch = graph.subgraphs.get_mut(&key).unwrap();
let victim = branch
.nodes
.iter()
.find_map(|(id, n)| (n.op_type == "Constant").then_some(id))
.unwrap();
branch.node_mut(victim).op_type = "Add".into();
CudaOnDeviceConstantSelect
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph.nodes.values().any(|n| n.op_type == "If"),
"a non-constant branch must not be rewritten"
);
assert!(where_nodes(&graph).is_empty());
}
#[test]
fn skips_padded_if_when_threshold_mismatches_short_table() {
let (mut graph, _cos, _sin) = longrope_if_graph(vec![4, 2], vec![2, 2], 3, true);
CudaOnDeviceConstantSelect
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(graph.nodes.values().any(|n| n.op_type == "If"));
assert!(where_nodes(&graph).is_empty());
}
#[test]
fn skips_padded_if_without_greater_predicate() {
let (mut graph, _cos, _sin) = longrope_if_graph(vec![4, 2], vec![2, 2], 2, false);
CudaOnDeviceConstantSelect
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(graph.nodes.values().any(|n| n.op_type == "If"));
assert!(where_nodes(&graph).is_empty());
}
#[test]
fn skips_if_when_true_branch_is_smaller() {
let (mut graph, _cos, _sin) = longrope_if_graph(vec![2, 2], vec![4, 2], 2, true);
CudaOnDeviceConstantSelect
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(graph.nodes.values().any(|n| n.op_type == "If"));
assert!(where_nodes(&graph).is_empty());
}
}