use onnx_runtime_ir::{
Attribute, DataType, Dim, Graph, Node, NodeId, TensorData, ValueId, WeightRef, static_shape,
};
use onnx_runtime_optimizer::{
OptimizationPass, OptimizerError, PassContext, Result as OptimizerResult,
};
use crate::kernels::linear_attention::{
FUSE_BETA_SIGMOID_ATTR, FUSE_DECAY_SOFTPLUS_ATTR, FUSE_NEG_EXP_ATTR,
};
use crate::runtime::CudaDeviceCapabilities;
pub(crate) const SILU_MUL_FUSION_ATTR: &str = "_cuda_silu_mul";
pub(crate) const DECOMPOSED_SILU_ATTR: &str = "_cuda_decomposed_silu";
pub(crate) const CUDA_RSQRT_ATTR: &str = "_cuda_rsqrt";
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;
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaSiluFusion;
pub(crate) fn cuda_optimization_passes(
device: Option<CudaDeviceCapabilities>,
) -> Vec<Box<dyn OptimizationPass>> {
vec![
Box::new(CudaSiluFusion),
Box::new(CudaRsqrtFusion),
Box::new(CudaL2NormalizeFusion),
Box::new(CudaLinearAttentionGatingFusion),
Box::new(CudaFoldConstantTranspose),
Box::new(CudaFoldConstantCast),
Box::new(CudaQkvProjectionFusion),
Box::new(CudaDropNormalizationCasts),
Box::new(CudaSkipRmsNormFusion),
Box::new(CudaMatMulNBitsBiasFusion),
Box::new(CudaSwiGluFusion),
Box::new(CudaGateUpSwiGluFusion),
Box::new(CudaSkipRmsNormMatMulFusion::for_device(device)),
Box::new(CudaOnDeviceConstantSelect),
Box::new(CudaDropIdentityCast),
]
}
impl OptimizationPass for CudaSiluFusion {
fn name(&self) -> &str {
"CudaSiluFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
let sigmoid_ids: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "Sigmoid"
&& node.is_default_domain()
&& node.inputs.len() == 1
&& node.outputs.len() == 1)
.then_some(id)
})
.collect();
let mut changed = false;
for sigmoid_id in sigmoid_ids {
let Some(sigmoid) = graph.try_node(sigmoid_id) else {
continue;
};
let Some(x) = sigmoid.inputs[0] else {
continue;
};
if !matches!(graph.value(x).dtype, DataType::Float16 | DataType::BFloat16) {
continue;
}
let sigmoid_output = sigmoid.outputs[0];
if graph.outputs.contains(&sigmoid_output) {
continue;
}
let consumers = graph.consumers(sigmoid_output);
if consumers.len() != 1 {
continue;
}
let mul_id = consumers[0];
let mul = graph.node(mul_id);
if mul.op_type != "Mul"
|| !mul.is_default_domain()
|| mul.inputs.len() != 2
|| mul.outputs.len() != 1
|| !((mul.inputs[0] == Some(x) && mul.inputs[1] == Some(sigmoid_output))
|| (mul.inputs[1] == Some(x) && mul.inputs[0] == Some(sigmoid_output)))
{
continue;
}
let mut silu = mul.clone();
silu.op_type = "Silu".to_string();
silu.domain = MICROSOFT_DOMAIN.to_string();
silu.version = None;
silu.inputs = vec![Some(x)];
silu.attributes.clear();
silu.attributes
.insert(DECOMPOSED_SILU_ATTR.into(), Attribute::Int(1));
graph.replace_node(mul_id, silu);
graph.remove_node(sigmoid_id);
graph
.opset_imports
.entry(MICROSOFT_DOMAIN.to_string())
.or_insert(1);
changed = true;
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaRsqrtFusion;
impl OptimizationPass for CudaRsqrtFusion {
fn name(&self) -> &str {
"CudaRsqrtFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
let sqrt_ids: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "Sqrt"
&& node.is_default_domain()
&& node.inputs.len() == 1
&& node.outputs.len() == 1)
.then_some(id)
})
.collect();
let mut changed = false;
for sqrt_id in sqrt_ids {
let Some(sqrt) = graph.try_node(sqrt_id) else {
continue;
};
let Some(x) = sqrt.inputs[0] else {
continue;
};
if !matches!(
graph.value(x).dtype,
DataType::Float32 | DataType::Float16 | DataType::BFloat16
) {
continue;
}
let sqrt_output = sqrt.outputs[0];
if graph.outputs.contains(&sqrt_output) {
continue;
}
let consumers = graph.consumers(sqrt_output);
if consumers.len() != 1 {
continue;
}
let recip_id = consumers[0];
let recip = graph.node(recip_id);
if recip.op_type != "Reciprocal"
|| !recip.is_default_domain()
|| recip.inputs.len() != 1
|| recip.outputs.len() != 1
|| recip.inputs[0] != Some(sqrt_output)
|| recip.attr(CUDA_RSQRT_ATTR).is_some()
{
continue;
}
let mut rsqrt = recip.clone();
rsqrt.inputs = vec![Some(x)];
rsqrt
.attributes
.insert(CUDA_RSQRT_ATTR.into(), Attribute::Int(1));
graph.replace_node(recip_id, rsqrt);
graph.remove_node(sqrt_id);
changed = true;
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
fn l2_normalize_fusion_disabled() -> bool {
std::env::var_os("ONNX_GENAI_CUDA_L2NORM_FUSION").is_some_and(|v| v == "0")
}
pub(crate) struct CudaL2NormalizeFusion;
impl CudaL2NormalizeFusion {
fn reduce_axis(graph: &Graph, rss: &Node) -> Option<i64> {
if rss
.attr("keepdims")
.and_then(Attribute::as_int)
.unwrap_or(1)
!= 1
{
return None;
}
if let Some(axes) = rss.attr("axes").and_then(Attribute::as_ints) {
return (axes.len() == 1).then(|| axes[0]);
}
if let Some(Some(axes_val)) = rss.inputs.get(1) {
if let Some(WeightRef::Inline(t)) = graph.initializers.get(axes_val)
&& t.dtype == DataType::Int64
&& t.numel() == 1
&& t.data.len() >= 8
{
let mut bytes = [0u8; 8];
bytes.copy_from_slice(&t.data[..8]);
return Some(i64::from_le_bytes(bytes));
}
return None;
}
let x = rss.inputs.first().copied().flatten()?;
let in_shape = &graph.value(x).shape;
let out_shape = &graph.value(rss.outputs[0]).shape;
if in_shape.len() != out_shape.len() {
return None;
}
let mut axis = None;
for (index, (a, b)) in in_shape.iter().zip(out_shape.iter()).enumerate() {
let (Some(a), Some(b)) = (a.as_static(), b.as_static()) else {
return None;
};
if a != b {
if b != 1 || axis.is_some() {
return None;
}
axis = Some(index as i64);
}
}
axis
}
}
impl OptimizationPass for CudaL2NormalizeFusion {
fn name(&self) -> &str {
"CudaL2NormalizeFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
if l2_normalize_fusion_disabled() {
return Ok(());
}
let div_ids: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "Div"
&& node.is_default_domain()
&& node.inputs.len() == 2
&& node.outputs.len() == 1)
.then_some(id)
})
.collect();
let mut changed = false;
for div_id in div_ids {
let Some(div) = graph.try_node(div_id) else {
continue;
};
let (Some(x), Some(nrm)) = (div.inputs[0], div.inputs[1]) else {
continue;
};
if !matches!(
graph.value(x).dtype,
DataType::Float32 | DataType::Float16 | DataType::BFloat16
) {
continue;
}
let Some(sqrt_id) = graph.value(nrm).producer else {
continue;
};
let sqrt = graph.node(sqrt_id);
if sqrt.op_type != "Sqrt"
|| !sqrt.is_default_domain()
|| sqrt.inputs.len() != 1
|| sqrt.outputs.len() != 1
|| graph.value(nrm).is_graph_output
|| graph.consumers(nrm).len() != 1
{
continue;
}
let Some(sq) = sqrt.inputs[0] else {
continue;
};
let Some(rss_id) = graph.value(sq).producer else {
continue;
};
let rss = graph.node(rss_id);
if rss.op_type != "ReduceSumSquare"
|| !rss.is_default_domain()
|| rss.inputs.first().copied().flatten() != Some(x)
|| rss.outputs.len() != 1
|| graph.value(sq).is_graph_output
|| graph.consumers(sq).len() != 1
{
continue;
}
let Some(axis) = Self::reduce_axis(graph, rss) else {
continue;
};
let mut lpnorm = div.clone();
lpnorm.op_type = "LpNormalization".to_string();
lpnorm.domain = String::new();
lpnorm.version = None;
lpnorm.inputs = vec![Some(x)];
lpnorm.attributes.clear();
lpnorm.attributes.insert("p".into(), Attribute::Int(2));
lpnorm
.attributes
.insert("axis".into(), Attribute::Int(axis));
lpnorm
.attributes
.insert("fused_reduce_chain".into(), Attribute::Int(1));
graph.replace_node(div_id, lpnorm);
graph.remove_node(sqrt_id);
graph.remove_node(rss_id);
changed = true;
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaLinearAttentionGatingFusion;
const LINEAR_ATTENTION_GATING_DISABLE_ENV: &str = "ONNX_GENAI_CUDA_DISABLE_LINATTN_GATING_FUSION";
fn linear_attention_gating_disabled() -> bool {
std::env::var_os(LINEAR_ATTENTION_GATING_DISABLE_ENV)
.is_some_and(|value| value != "0" && !value.is_empty())
}
struct BetaSigmoidFold {
sigmoid_id: NodeId,
raw: ValueId,
}
struct DecaySoftplusFold {
a: ValueId,
dt_bias: ValueId,
neg_exp_a: ValueId,
fuse_neg_exp: bool,
dead: Vec<NodeId>,
}
impl CudaLinearAttentionGatingFusion {
fn is_gated_delta_la(node: &Node) -> bool {
node.op_type == "LinearAttention"
&& node.domain == MICROSOFT_DOMAIN
&& node.inputs.len() >= 6
&& node.outputs.len() == 2
}
fn sole_producer(graph: &Graph, value: ValueId, consumer_id: NodeId) -> Option<NodeId> {
if graph.outputs.contains(&value) {
return None;
}
let consumers = graph.consumers(value);
if consumers.len() != 1 || consumers[0] != consumer_id {
return None;
}
graph.value(value).producer
}
fn match_beta_sigmoid(
graph: &Graph,
beta: ValueId,
la_id: NodeId,
io_dtype: DataType,
) -> Option<BetaSigmoidFold> {
let sigmoid_id = Self::sole_producer(graph, beta, la_id)?;
let sigmoid = graph.node(sigmoid_id);
if sigmoid.op_type != "Sigmoid"
|| !sigmoid.is_default_domain()
|| sigmoid.inputs.len() != 1
|| sigmoid.outputs.len() != 1
{
return None;
}
let raw = sigmoid.inputs[0]?;
if graph.value(raw).dtype != io_dtype {
return None;
}
Some(BetaSigmoidFold { sigmoid_id, raw })
}
fn match_neg_exp_chain(
graph: &Graph,
neg_exp_a: ValueId,
consumer_id: NodeId,
io_dtype: DataType,
) -> Option<(ValueId, [NodeId; 2])> {
let neg_id = Self::sole_producer(graph, neg_exp_a, consumer_id)?;
let neg = graph.node(neg_id);
if neg.op_type != "Neg"
|| !neg.is_default_domain()
|| neg.inputs.len() != 1
|| neg.outputs.len() != 1
{
return None;
}
let exp_out = neg.inputs[0]?;
let exp_id = Self::sole_producer(graph, exp_out, neg_id)?;
let exp = graph.node(exp_id);
if exp.op_type != "Exp"
|| !exp.is_default_domain()
|| exp.inputs.len() != 1
|| exp.outputs.len() != 1
{
return None;
}
let a_log = exp.inputs[0]?;
if !graph.initializers.contains_key(&a_log) || graph.value(a_log).dtype != io_dtype {
return None;
}
Some((a_log, [neg_id, exp_id]))
}
fn match_decay_softplus(
graph: &Graph,
decay: ValueId,
la_id: NodeId,
io_dtype: DataType,
) -> Option<DecaySoftplusFold> {
let mut dead = Vec::new();
let (mul_out, mul_consumer) = {
let producer = Self::sole_producer(graph, decay, la_id)?;
let node = graph.node(producer);
if node.op_type == "Cast"
&& node.is_default_domain()
&& node.inputs.len() == 1
&& node.outputs.len() == 1
&& graph.value(node.inputs[0]?).dtype == io_dtype
{
dead.push(producer);
(node.inputs[0]?, producer)
} else {
(decay, la_id)
}
};
let mul_id = Self::sole_producer(graph, mul_out, mul_consumer)?;
let mul = graph.node(mul_id);
if mul.op_type != "Mul"
|| !mul.is_default_domain()
|| mul.inputs.len() != 2
|| mul.outputs.len() != 1
{
return None;
}
let (a0, a1) = (mul.inputs[0]?, mul.inputs[1]?);
dead.push(mul_id);
let (neg_exp_a, softplus_out, fuse_neg_exp) = if graph.initializers.contains_key(&a0) {
(a0, a1, false)
} else if graph.initializers.contains_key(&a1) {
(a1, a0, false)
} else if let Some((a_log, chain)) = Self::match_neg_exp_chain(graph, a0, mul_id, io_dtype)
{
dead.extend(chain);
(a_log, a1, true)
} else if let Some((a_log, chain)) = Self::match_neg_exp_chain(graph, a1, mul_id, io_dtype)
{
dead.extend(chain);
(a_log, a0, true)
} else {
return None;
};
if !fuse_neg_exp && graph.value(neg_exp_a).dtype != io_dtype {
return None;
}
let softplus_id = Self::sole_producer(graph, softplus_out, mul_id)?;
let softplus = graph.node(softplus_id);
if softplus.op_type != "Softplus"
|| !softplus.is_default_domain()
|| softplus.inputs.len() != 1
|| softplus.outputs.len() != 1
{
return None;
}
let add_out = softplus.inputs[0]?;
dead.push(softplus_id);
let add_id = Self::sole_producer(graph, add_out, softplus_id)?;
let add = graph.node(add_id);
if add.op_type != "Add"
|| !add.is_default_domain()
|| add.inputs.len() != 2
|| add.outputs.len() != 1
{
return None;
}
let (b0, b1) = (add.inputs[0]?, add.inputs[1]?);
let (dt_bias, a) = if graph.initializers.contains_key(&b1) {
(b1, b0)
} else if graph.initializers.contains_key(&b0) {
(b0, b1)
} else {
return None;
};
if graph.value(dt_bias).dtype != io_dtype || graph.value(a).dtype != io_dtype {
return None;
}
dead.push(add_id);
Some(DecaySoftplusFold {
a,
dt_bias,
neg_exp_a,
fuse_neg_exp,
dead,
})
}
}
impl OptimizationPass for CudaLinearAttentionGatingFusion {
fn name(&self) -> &str {
"CudaLinearAttentionGatingFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
if linear_attention_gating_disabled() {
return Ok(());
}
let la_ids: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| Self::is_gated_delta_la(node).then_some(id))
.collect();
let mut changed = false;
for la_id in la_ids {
let Some(la) = graph.try_node(la_id) else {
continue;
};
let Some(io_dtype) = la.inputs[0].map(|v| graph.value(v).dtype) else {
continue;
};
let beta_slot = la.inputs.get(5).copied().flatten();
let decay_slot = la.inputs.get(4).copied().flatten();
let decay_foldable = la.inputs.len() == 6;
let beta_fold =
beta_slot.and_then(|beta| Self::match_beta_sigmoid(graph, beta, la_id, io_dtype));
let decay_fold = if decay_foldable {
decay_slot
.and_then(|decay| Self::match_decay_softplus(graph, decay, la_id, io_dtype))
} else {
None
};
if beta_fold.is_none() && decay_fold.is_none() {
continue;
}
let mut new_la = la.clone();
let mut dead = Vec::new();
if let Some(fold) = beta_fold {
new_la.inputs[5] = Some(fold.raw);
new_la
.attributes
.insert(FUSE_BETA_SIGMOID_ATTR.into(), Attribute::Int(1));
dead.push(fold.sigmoid_id);
}
if let Some(fold) = decay_fold {
new_la.inputs[4] = Some(fold.a);
new_la.inputs.push(Some(fold.dt_bias));
new_la.inputs.push(Some(fold.neg_exp_a));
new_la
.attributes
.insert(FUSE_DECAY_SOFTPLUS_ATTR.into(), Attribute::Int(1));
if fold.fuse_neg_exp {
new_la
.attributes
.insert(FUSE_NEG_EXP_ATTR.into(), Attribute::Int(1));
}
dead.extend(fold.dead);
}
graph.replace_node(la_id, new_la);
graph.remove_nodes(&dead);
changed = true;
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
#[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,
narrow_dtype: DataType,
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;
if node.op_type == "RMSNormalization" && matches!(node.domain.as_str(), "" | "ai.onnx") {
node.op_type = "SimplifiedLayerNormalization".into();
}
graph.replace_node(plan.node_id, node);
for output in plan.retyped_outputs {
graph.value_mut(output).dtype = plan.narrow_dtype;
}
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]),
("RMSNormalization", "" | "ai.onnx") => Some(&[0]),
_ => None,
}
}
fn fp32_cast_from_narrow(
&self,
graph: &Graph,
value: ValueId,
) -> Option<(NodeId, ValueId, DataType)> {
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()?;
let source_dtype = graph.try_value(source)?.dtype;
if source_dtype != DataType::Float16 && source_dtype != DataType::BFloat16 {
return None;
}
Some((producer, source, source_dtype))
}
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;
}
if node.op_type == "RMSNormalization" && node.outputs.len() != 1 {
return None;
}
let mut new_inputs = node.inputs.clone();
let mut dead_input_casts = Vec::new();
let mut narrow_dtype: Option<DataType> = None;
for &index in activation_indices {
let value = node.inputs.get(index).copied().flatten()?;
let (cast_id, source, source_dtype) = self.fp32_cast_from_narrow(graph, value)?;
if *narrow_dtype.get_or_insert(source_dtype) != source_dtype {
return None;
}
new_inputs[index] = Some(source);
dead_input_casts.push(cast_id);
}
let narrow_dtype = narrow_dtype?;
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)? != narrow_dtype {
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,
narrow_dtype,
new_inputs,
retyped_outputs,
output_cast_bypass,
dead_input_casts,
})
}
}
const SKIP_RMSNORM_FUSION_ENABLE_ENV: &str = "ONNX_GENAI_CUDA_ENABLE_SKIP_RMSNORM_FUSION";
fn skip_rmsnorm_fusion_enabled() -> bool {
std::env::var_os(SKIP_RMSNORM_FUSION_ENABLE_ENV)
.is_some_and(|value| value != "0" && !value.is_empty())
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaSkipRmsNormFusion;
struct SkipRmsNormFoldPlan {
add_id: NodeId,
norm_id: NodeId,
a: ValueId,
b: ValueId,
gamma: ValueId,
normalized_out: ValueId,
residual_sum: ValueId,
epsilon: f32,
stat_shape: onnx_runtime_ir::Shape,
}
impl OptimizationPass for CudaSkipRmsNormFusion {
fn name(&self) -> &str {
"CudaSkipRmsNormFusion"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
if !skip_rmsnorm_fusion_enabled() {
return Ok(());
}
let norm_ids: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "SimplifiedLayerNormalization" && node.is_default_domain())
.then_some(id)
})
.collect();
let mut plans: Vec<SkipRmsNormFoldPlan> = Vec::new();
let mut used_adds: std::collections::HashSet<NodeId> = std::collections::HashSet::new();
for norm_id in norm_ids {
if let Some(plan) = self.plan_fold(graph, norm_id) {
if !used_adds.insert(plan.add_id) {
continue;
}
plans.push(plan);
}
}
let changed = !plans.is_empty();
for plan in plans {
graph.remove_node(plan.add_id);
let mean_v = graph.create_value(DataType::Float32, plan.stat_shape.clone());
let invstd_v = graph.create_value(DataType::Float32, plan.stat_shape.clone());
let mut skip = Node::new(
NodeId(0),
"SkipSimplifiedLayerNormalization",
vec![Some(plan.a), Some(plan.b), Some(plan.gamma)],
vec![plan.normalized_out, mean_v, invstd_v, plan.residual_sum],
);
skip.domain = MICROSOFT_DOMAIN.to_string();
skip.attributes
.insert("epsilon".into(), Attribute::Float(plan.epsilon));
graph.replace_node(plan.norm_id, skip);
graph
.opset_imports
.entry(MICROSOFT_DOMAIN.to_string())
.or_insert(1);
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
impl CudaSkipRmsNormFusion {
fn plan_fold(&self, graph: &Graph, norm_id: NodeId) -> Option<SkipRmsNormFoldPlan> {
let norm = graph.try_node(norm_id)?;
if norm.outputs.len() != 1 || norm.inputs.len() < 2 {
return None;
}
let norm_input = norm.inputs[0]?;
let gamma = norm.inputs[1]?;
if norm.inputs.get(2).copied().flatten().is_some() {
return None; }
let normalized_out = *norm.outputs.first()?;
if graph.value(normalized_out).is_graph_output {
return None;
}
if graph.value(norm_input).dtype != DataType::BFloat16 {
return None;
}
let gamma_dtype = graph.value(gamma).dtype;
if gamma_dtype != DataType::BFloat16 && gamma_dtype != DataType::Float32 {
return None;
}
let add_id = graph.try_value(norm_input)?.producer?;
let add = graph.try_node(add_id)?;
if add.op_type != "Add"
|| !add.is_default_domain()
|| add.inputs.len() != 2
|| add.outputs.len() != 1
{
return None;
}
let a = add.inputs[0]?;
let b = add.inputs[1]?;
let residual_sum = *add.outputs.first()?;
if residual_sum != norm_input || graph.value(residual_sum).is_graph_output {
return None;
}
let a_meta = graph.value(a);
let b_meta = graph.value(b);
let sum_meta = graph.value(residual_sum);
if a_meta.dtype != DataType::BFloat16 || b_meta.dtype != DataType::BFloat16 {
return None;
}
if a_meta.shape != sum_meta.shape || b_meta.shape != sum_meta.shape {
return None;
}
if sum_meta.shape.is_empty() {
return None;
}
let stat_shape: onnx_runtime_ir::Shape =
sum_meta.shape[..sum_meta.shape.len() - 1].to_vec();
let epsilon = norm
.attr("epsilon")
.and_then(Attribute::as_float)
.unwrap_or(1e-5);
Some(SkipRmsNormFoldPlan {
add_id,
norm_id,
a,
b,
gamma,
normalized_out,
residual_sum,
epsilon,
stat_shape,
})
}
}
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 CudaFoldConstantCast;
const CONST_CAST_FOLD_DISABLE_ENV: &str = "ONNX_GENAI_CUDA_DISABLE_CONST_CAST_FOLD";
fn const_cast_fold_disabled() -> bool {
std::env::var_os(CONST_CAST_FOLD_DISABLE_ENV)
.is_some_and(|value| value != "0" && !value.is_empty())
}
struct CastFoldPlan {
node: NodeId,
output: ValueId,
dtype: DataType,
dims: Vec<usize>,
bytes: Vec<u8>,
}
impl OptimizationPass for CudaFoldConstantCast {
fn name(&self) -> &str {
"CudaFoldConstantCast"
}
fn run(&self, graph: &mut Graph, ctx: &PassContext) -> OptimizerResult<()> {
if const_cast_fold_disabled() {
return Ok(());
}
let candidates: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "Cast"
&& matches!(node.domain.as_str(), "" | "ai.onnx")
&& node.inputs.len() == 1
&& node.outputs.len() == 1)
.then_some(id)
})
.collect();
let mut plans: Vec<CastFoldPlan> = 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.dims.clone());
let tensor = TensorData::from_raw(plan.dtype, plan.dims, plan.bytes);
graph.set_initializer(plan.output, WeightRef::Inline(tensor));
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
impl CudaFoldConstantCast {
fn plan_fold(&self, graph: &Graph, ctx: &PassContext, node_id: NodeId) -> Option<CastFoldPlan> {
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 src_dtype = weight.dtype();
let dst_dtype = DataType::from_onnx(node.attr("to").and_then(Attribute::as_int)? as i32)?;
if src_dtype.byte_size() == 0 || src_dtype.is_sub_byte() {
return None;
}
let dims = weight.dims().to_vec();
let src = ctx.initializer_bytes(weight)?;
let expected = dims
.iter()
.product::<usize>()
.checked_mul(src_dtype.byte_size())?;
if src.len() != expected {
return None;
}
if src_dtype == dst_dtype {
return None;
}
let bytes = convert_float_bytes(src, src_dtype, dst_dtype)?;
Some(CastFoldPlan {
node: node_id,
output,
dtype: dst_dtype,
dims,
bytes,
})
}
}
fn convert_float_bytes(src: &[u8], from: DataType, to: DataType) -> Option<Vec<u8>> {
let values: Vec<f64> = match from {
DataType::Float32 => src
.chunks_exact(4)
.map(|c| f32::from_ne_bytes([c[0], c[1], c[2], c[3]]) as f64)
.collect(),
DataType::Float64 => src
.chunks_exact(8)
.map(|c| f64::from_ne_bytes([c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]]))
.collect(),
DataType::Float16 => src
.chunks_exact(2)
.map(|c| half::f16::from_bits(u16::from_ne_bytes([c[0], c[1]])).to_f32() as f64)
.collect(),
DataType::BFloat16 => src
.chunks_exact(2)
.map(|c| half::bf16::from_bits(u16::from_ne_bytes([c[0], c[1]])).to_f32() as f64)
.collect(),
_ => return None,
};
let mut out = Vec::with_capacity(values.len() * to.byte_size().max(1));
match to {
DataType::Float32 => {
for v in values {
out.extend_from_slice(&(v as f32).to_ne_bytes());
}
}
DataType::Float64 => {
for v in values {
out.extend_from_slice(&v.to_ne_bytes());
}
}
DataType::Float16 => {
for v in values {
out.extend_from_slice(&half::f16::from_f32(v as f32).to_bits().to_ne_bytes());
}
}
DataType::BFloat16 => {
for v in values {
out.extend_from_slice(&half::bf16::from_f32(v as f32).to_bits().to_ne_bytes());
}
}
_ => return None,
}
Some(out)
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaDropIdentityCast;
const IDENTITY_CAST_FOLD_DISABLE_ENV: &str = "ONNX_GENAI_CUDA_DISABLE_IDENTITY_CAST_FOLD";
fn identity_cast_fold_disabled() -> bool {
std::env::var_os(IDENTITY_CAST_FOLD_DISABLE_ENV)
.is_some_and(|value| value != "0" && !value.is_empty())
}
impl OptimizationPass for CudaDropIdentityCast {
fn name(&self) -> &str {
"CudaDropIdentityCast"
}
fn run(&self, graph: &mut Graph, _ctx: &PassContext) -> OptimizerResult<()> {
if identity_cast_fold_disabled() {
return Ok(());
}
let candidates: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
((node.op_type == "Cast" || node.op_type == "CastLike")
&& matches!(node.domain.as_str(), "" | "ai.onnx")
&& !node.inputs.is_empty()
&& node.outputs.len() == 1)
.then_some(id)
})
.collect();
let mut changed = false;
for node_id in candidates {
let Some(node) = graph.try_node(node_id) else {
continue;
};
let Some(input) = node.inputs[0] else {
continue;
};
let output = node.outputs[0];
if input == output {
continue;
}
if !graph.value_type_is_known(input) || !graph.value_type_is_known(output) {
continue;
}
if graph.value(input).dtype != graph.value(output).dtype {
continue;
}
if graph.outputs.contains(&output) {
continue;
}
graph.replace_all_uses(output, input);
graph.remove_node(node_id);
graph.gc_value_if_orphan(output);
changed = true;
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
}
#[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],
plan.matmul_inputs[3],
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>; 4],
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);
const MATMUL_NBITS_ZERO_POINTS_INPUT: usize = 3;
const MATMUL_NBITS_GROUP_INDEX_INPUT: usize = 4;
if matmul.inputs.len() > MATMUL_NBITS_GROUP_INDEX_INPUT
&& matmul
.inputs
.iter()
.skip(MATMUL_NBITS_GROUP_INDEX_INPUT)
.any(Option::is_some)
{
return None;
}
if matmul.inputs.first().copied().flatten().is_none()
|| matmul.inputs.get(1).copied().flatten().is_none()
|| matmul.inputs.get(2).copied().flatten().is_none()
{
return None;
}
let zero_points = matmul
.inputs
.get(MATMUL_NBITS_ZERO_POINTS_INPUT)
.copied()
.flatten();
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],
zero_points,
],
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_SIZES: [i64; 1] = [32];
fn rmsnorm_fusion_supports_block_size(block_size: i64) -> bool {
RMSNORM_FUSION_SUPPORTED_BLOCK_SIZES.contains(&block_size)
}
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_ANCHOR_SM_COUNT: u32 = 132;
fn derived_min_hidden(sm_count: u32) -> usize {
let sm_count = sm_count.max(1);
let anchor_chunks = (RMSNORM_FUSION_MIN_HIDDEN / RMSNORM_FUSION_WARP_HALF4_MULTIPLE) as u64;
let anchor_sm = u64::from(RMSNORM_FUSION_ANCHOR_SM_COUNT);
let chunks = (anchor_chunks * u64::from(sm_count) + anchor_sm / 2) / anchor_sm;
(chunks.max(1) as usize) * 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,
device: Option<CudaDeviceCapabilities>,
) -> bool {
let derived = device
.map(|caps| derived_min_hidden(caps.multiprocessor_count()))
.unwrap_or(RMSNORM_FUSION_MIN_HIDDEN);
let floor = env_usize(RMSNORM_FUSION_MIN_HIDDEN_ENV, derived);
norm_size >= floor
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaSkipRmsNormMatMulFusion {
device: Option<CudaDeviceCapabilities>,
}
impl CudaSkipRmsNormMatMulFusion {
pub(crate) fn for_device(device: Option<CudaDeviceCapabilities>) -> Self {
Self { device }
}
}
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 input_probe = self.preceding_gemv_attributed(graph, input_value, norm_size);
let skip_probe = self.preceding_gemv_attributed(graph, skip_value, norm_size);
if std::env::var_os("ONNX_GENAI_CUDA_DEBUG_RMSNORM_FOLD").is_some() {
let describe = |r: &Result<NodeId, &'static str>| match r {
Ok(_) => "OK".to_string(),
Err(reason) => (*reason).to_string(),
};
eprintln!(
"rmsnorm_fold_probe: norm={} input={} skip={}",
graph
.try_node(skip_id)
.map(|n| n.name.as_str())
.unwrap_or("<unnamed>"),
describe(&input_probe),
describe(&skip_probe),
);
}
let (preceding_out, residual) = match (input_probe.is_ok(), skip_probe.is_ok()) {
(true, true) => return None,
(true, false) => (input_value, skip_value),
(false, true) => (skip_value, input_value),
(false, false) => 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) {
if std::env::var_os("ONNX_GENAI_CUDA_DEBUG_RMSNORM_FOLD").is_some() {
let n = graph.node(*following_id);
eprintln!(
"rmsnorm_fold_decline: following GEMV not prologue-capable: op={} name={} inputs={} K={:?} N={:?}",
n.op_type,
n.name,
n.input_values().count(),
n.attr("K").and_then(Attribute::as_int),
n.attr("N").and_then(Attribute::as_int),
);
}
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, self.device)
{
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> {
self.preceding_gemv_attributed(graph, value, norm_size).ok()
}
fn preceding_gemv_attributed(
&self,
graph: &Graph,
value: ValueId,
norm_size: usize,
) -> Result<NodeId, &'static str> {
let producer = graph
.try_value(value)
.and_then(|v| v.producer)
.ok_or("no producer")?;
let node = graph.try_node(producer).ok_or("producer missing")?;
if node.op_type != "MatMulNBits" || node.domain != MICROSOFT_DOMAIN {
return Err("producer is not a com.microsoft MatMulNBits");
}
if node.attr(GATE_UP_SWIGLU_FUSION_ATTR).is_some() {
return Err("producer already carries the gate/up SwiGLU fusion");
}
if node.attr(MATMUL_NBITS_RMSNORM_PROLOGUE_ATTR).is_some() {
return Err("producer already carries an RMSNorm prologue");
}
let value_count = node.input_values().count();
if !(value_count == 3 || value_count == 4) {
return Err("producer is not the A/B/scales(+zero-points) form");
}
if node.inputs.iter().skip(4).any(Option::is_some) {
return Err("producer has a group index or pre-existing bias");
}
if let Some(zero_points) = node.inputs.get(3).copied().flatten()
&& graph.try_value(zero_points).map(|value| value.dtype) != Some(DataType::Uint8)
{
return Err("producer zero-points are not uint8");
}
if !self.is_fusable_bits_fp16_matmul(graph, node) {
return Err("producer is not a fusable int4/int8 fp16 GEMV");
}
if node.attr("N").and_then(Attribute::as_int).ok_or("no N")? as usize != norm_size {
return Err("producer N does not equal the hidden size");
}
if graph.consumers(value).len() != 1 {
return Err("producer output has more than one consumer");
}
if graph.value(value).is_graph_output {
return Err("producer output is a graph output");
}
Ok(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)
.is_some_and(rmsnorm_fusion_supports_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 CudaSwiGluFusion {
fn is_silu(node: &Node) -> bool {
if node.inputs.len() != 1 || node.outputs.len() != 1 {
return false;
}
match node.op_type.as_str() {
"Silu" => node.domain == MICROSOFT_DOMAIN,
"Swish" => {
matches!(node.domain.as_str(), "" | "ai.onnx")
&& node.attributes.keys().all(|name| name.as_str() == "alpha")
&& node
.attr("alpha")
.and_then(Attribute::as_float)
.unwrap_or(1.0)
== 1.0
}
_ => false,
}
}
}
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)| Self::is_silu(node).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));
if silu.attr(DECOMPOSED_SILU_ATTR).and_then(Attribute::as_int) == Some(1) {
fused
.attributes
.insert(DECOMPOSED_SILU_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,
decomposed_silu: bool,
}
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));
if plan.decomposed_silu {
fused
.attributes
.insert(DECOMPOSED_SILU_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,
decomposed_silu: mul.attr(DECOMPOSED_SILU_ATTR).and_then(Attribute::as_int) == Some(1),
})
}
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]?;
let activation_dtype = graph.value(activation).dtype;
if !matches!(activation_dtype, DataType::Float16 | DataType::BFloat16)
|| graph.value(matmul.outputs[0]).dtype != activation_dtype
|| graph.value(scales).dtype != activation_dtype
{
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,
}
const QKV_FUSION_ENABLE_ENV: &str = "ONNX_GENAI_CUDA_ENABLE_QKV_FUSION";
fn qkv_fusion_enabled() -> bool {
std::env::var_os(QKV_FUSION_ENABLE_ENV).is_some_and(|value| value != "0" && !value.is_empty())
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct CudaQkvProjectionFusion;
struct QkvProj {
matmul: NodeId,
out: ValueId,
weight: ValueId,
scales: ValueId,
zero_points: Option<ValueId>,
activation: ValueId,
n: usize,
k: usize,
block_size: i64,
bits: i64,
}
struct QkvFusionPlan {
activation: ValueId,
q: QkvProj,
k: QkvProj,
v: QkvProj,
}
impl OptimizationPass for CudaQkvProjectionFusion {
fn name(&self) -> &str {
"CudaQkvProjectionFusion"
}
fn run(&self, graph: &mut Graph, ctx: &PassContext) -> OptimizerResult<()> {
if !qkv_fusion_enabled() {
return Ok(());
}
self.fuse_all(graph, ctx)
}
}
impl CudaQkvProjectionFusion {
fn fuse_all(&self, graph: &mut Graph, ctx: &PassContext) -> OptimizerResult<()> {
let gqa_nodes: Vec<NodeId> = graph
.nodes
.iter()
.filter_map(|(id, node)| {
(node.op_type == "GroupQueryAttention" && node.domain == MICROSOFT_DOMAIN)
.then_some(id)
})
.collect();
let mut plans: Vec<QkvFusionPlan> = Vec::new();
for gqa_id in gqa_nodes {
if let Some(plan) = self.plan_fuse(graph, gqa_id) {
plans.push(plan);
}
}
let changed = !plans.is_empty();
for plan in plans {
self.apply_fuse(graph, ctx, plan)?;
}
if changed {
graph.validate().map_err(OptimizerError::from)?;
}
Ok(())
}
fn plan_fuse(&self, graph: &Graph, gqa_id: NodeId) -> Option<QkvFusionPlan> {
let gqa = graph.try_node(gqa_id)?;
if gqa.inputs.len() < 3 {
return None;
}
let q_mm = self.trace_back_to_matmul(graph, gqa.inputs[0]?)?;
let k_mm = self.trace_back_to_matmul(graph, gqa.inputs[1]?)?;
let v_mm = self.trace_back_to_matmul(graph, gqa.inputs[2]?)?;
if q_mm == k_mm || q_mm == v_mm || k_mm == v_mm {
return None;
}
let q = self.eligible_projection(graph, q_mm)?;
let k = self.eligible_projection(graph, k_mm)?;
let v = self.eligible_projection(graph, v_mm)?;
if q.activation != k.activation || q.activation != v.activation {
return None;
}
if q.k != k.k || q.k != v.k {
return None;
}
if q.block_size != k.block_size || q.block_size != v.block_size {
return None;
}
if q.bits != k.bits || q.bits != v.bits {
return None;
}
if q.zero_points.is_some() != k.zero_points.is_some()
|| q.zero_points.is_some() != v.zero_points.is_some()
{
return None;
}
if graph.value(q.out).dtype != graph.value(k.out).dtype
|| graph.value(q.out).dtype != graph.value(v.out).dtype
{
return None;
}
if graph.initializers.get(&q.scales)?.dtype() != graph.initializers.get(&k.scales)?.dtype()
|| graph.initializers.get(&q.scales)?.dtype()
!= graph.initializers.get(&v.scales)?.dtype()
{
return None;
}
Some(QkvFusionPlan {
activation: q.activation,
q,
k,
v,
})
}
fn trace_back_to_matmul(&self, graph: &Graph, mut value: ValueId) -> Option<NodeId> {
for _ in 0..12 {
let producer = graph.try_value(value)?.producer?;
let node = graph.try_node(producer)?;
if node.op_type == "MatMulNBits" && node.domain == MICROSOFT_DOMAIN {
return Some(producer);
}
value = node.inputs.first().copied().flatten()?;
}
None
}
fn eligible_projection(&self, graph: &Graph, matmul_id: NodeId) -> Option<QkvProj> {
let node = graph.try_node(matmul_id)?;
let present: Vec<ValueId> = node.input_values().collect();
if !(present.len() == 3 || present.len() == 4)
|| node.inputs.iter().skip(4).any(Option::is_some)
{
return None;
}
let zero_points = node.inputs.get(3).copied().flatten();
let n = node.attr("N").and_then(Attribute::as_int)? as usize;
let k = node.attr("K").and_then(Attribute::as_int)? as usize;
let block_size = node.attr("block_size").and_then(Attribute::as_int)?;
let bits = node.attr("bits").and_then(Attribute::as_int).unwrap_or(4);
if block_size != GATE_UP_SWIGLU_SUPPORTED_BLOCK_SIZE as i64
|| bits != GATE_UP_SWIGLU_SUPPORTED_BITS
{
return None;
}
let activation = node.inputs[0]?;
let weight = node.inputs[1]?;
let scales = node.inputs[2]?;
if !graph.initializers.contains_key(&weight) || !graph.initializers.contains_key(&scales) {
return None;
}
if let Some(zp) = zero_points
&& !graph.initializers.contains_key(&zp)
{
return None;
}
let out = node.outputs[0];
if graph.consumers(out).len() != 1 || graph.value(out).is_graph_output {
return None;
}
Some(QkvProj {
matmul: matmul_id,
out,
weight,
scales,
zero_points,
activation,
n,
k,
block_size,
bits,
})
}
fn apply_fuse(
&self,
graph: &mut Graph,
ctx: &PassContext,
plan: QkvFusionPlan,
) -> OptimizerResult<()> {
let QkvFusionPlan {
activation,
q,
k,
v,
} = plan;
let n_total = q.n + k.n + v.n;
let weight_bytes = self.concat_initializers(graph, ctx, &[q.weight, k.weight, v.weight])?;
let scale_bytes = self.concat_initializers(graph, ctx, &[q.scales, k.scales, v.scales])?;
let zero_bytes = match (q.zero_points, k.zero_points, v.zero_points) {
(Some(qz), Some(kz), Some(vz)) => {
Some(self.concat_initializers(graph, ctx, &[qz, kz, vz])?)
}
_ => None,
};
let q_weight = graph.initializers.get(&q.weight).ok_or_else(|| {
OptimizerError::Fusion("qkv fusion: missing q weight initializer".into())
})?;
let mut weight_dims = q_weight.dims().to_vec();
let weight_dtype = q_weight.dtype();
if weight_dims.is_empty() {
return Err(OptimizerError::Fusion(
"qkv fusion: scalar weight initializer".into(),
));
}
weight_dims[0] = n_total;
let scale_dtype = graph
.initializers
.get(&q.scales)
.ok_or_else(|| OptimizerError::Fusion("qkv fusion: missing q scales".into()))?
.dtype();
let fused_weight = graph.create_value(weight_dtype, static_shape(weight_dims.clone()));
graph.set_initializer(
fused_weight,
WeightRef::Inline(TensorData::from_raw(
weight_dtype,
weight_dims,
weight_bytes,
)),
);
let scale_len = scale_bytes.len() / scale_dtype.byte_size().max(1);
let fused_scales = graph.create_value(scale_dtype, static_shape(vec![scale_len]));
graph.set_initializer(
fused_scales,
WeightRef::Inline(TensorData::from_raw(
scale_dtype,
vec![scale_len],
scale_bytes,
)),
);
let fused_zero = zero_bytes.map(|bytes| {
let value = graph.create_value(DataType::Uint8, static_shape(vec![bytes.len()]));
graph.set_initializer(
value,
WeightRef::Inline(TensorData::from_raw(
DataType::Uint8,
vec![bytes.len()],
bytes,
)),
);
value
});
let out_dtype = graph.value(q.out).dtype;
let mut fused_shape = graph.value(q.out).shape.clone();
match fused_shape.last_mut() {
Some(dim) => *dim = Dim::Static(n_total),
None => {
return Err(OptimizerError::Fusion(
"qkv fusion: scalar projection output".into(),
));
}
}
let fused_out = graph.create_value(out_dtype, fused_shape);
let mut attributes = graph.node(q.matmul).attributes.clone();
attributes.insert("N".into(), Attribute::Int(n_total as i64));
let version = graph.node(q.matmul).version;
let mut fused_inputs = vec![Some(activation), Some(fused_weight), Some(fused_scales)];
if let Some(zero) = fused_zero {
fused_inputs.push(Some(zero));
}
let mut fused_matmul = Node::new(NodeId(0), "MatMulNBits", fused_inputs, vec![fused_out]);
fused_matmul.domain = MICROSOFT_DOMAIN.into();
fused_matmul.version = version;
fused_matmul.attributes = attributes;
graph.insert_node(fused_matmul);
let axis = graph.value(q.out).shape.len() as i64 - 1;
for proj in [&q, &k, &v] {
graph.remove_node(proj.matmul);
}
for proj in [&q, &k, &v] {
for value in [Some(proj.weight), Some(proj.scales), proj.zero_points]
.into_iter()
.flatten()
{
graph.initializers.remove(&value);
graph.gc_value_if_orphan(value);
}
}
let mut split = Node::new(
NodeId(0),
"Split",
vec![Some(fused_out)],
vec![q.out, k.out, v.out],
);
split.attributes.insert("axis".into(), Attribute::Int(axis));
split.attributes.insert(
"split".into(),
Attribute::Ints(vec![q.n as i64, k.n as i64, v.n as i64]),
);
graph.insert_node(split);
Ok(())
}
fn concat_initializers(
&self,
graph: &Graph,
ctx: &PassContext,
values: &[ValueId],
) -> OptimizerResult<Vec<u8>> {
let mut out = Vec::new();
for &value in values {
let weight = graph.initializers.get(&value).ok_or_else(|| {
OptimizerError::Fusion("qkv fusion: missing initializer for concat".into())
})?;
let bytes = ctx.initializer_bytes(weight).ok_or_else(|| {
OptimizerError::Fusion("qkv fusion: could not resolve initializer bytes".into())
})?;
out.extend_from_slice(bytes);
}
Ok(out)
}
}
#[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 crate::test_support::EnvVarGuard;
use onnx_runtime_ir::{Dim, Node, NodeId, ValueId};
#[test]
fn optimizer_source_routes_all_env_mutation_through_guard() {
let src = include_str!("optimizer.rs");
let set_needle = format!("env::{}", ["set", "_var"].concat());
let remove_needle = format!("env::{}", ["remove", "_var"].concat());
assert!(
!src.contains(&set_needle),
"found a direct `env::{}` in optimizer.rs; route it through \
EnvVarGuard (test_support) so it cannot race parallel tests",
["set", "_var"].concat()
);
assert!(
!src.contains(&remove_needle),
"found a direct `env::{}` in optimizer.rs; route it through \
EnvVarGuard (test_support) so it cannot race parallel tests",
["remove", "_var"].concat()
);
}
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 bf16_skip_seam_graph(hidden: usize, gamma_dtype: DataType) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let a = value(&mut graph, "a", DataType::BFloat16, hidden);
let b = value(&mut graph, "b", DataType::BFloat16, hidden);
graph.add_input(a);
graph.add_input(b);
let sum = value(&mut graph, "sum", DataType::BFloat16, hidden);
graph.insert_node(Node::new(
NodeId(0),
"Add",
vec![Some(a), Some(b)],
vec![sum],
));
let gamma = vec1d(&mut graph, "gamma", gamma_dtype, hidden);
graph.set_initializer(
gamma,
WeightRef::Inline(TensorData::from_raw(
gamma_dtype,
vec![hidden],
vec![0u8; hidden * gamma_dtype.byte_size()],
)),
);
let normalized = value(&mut graph, "normalized", DataType::BFloat16, hidden);
let mut norm = Node::new(
NodeId(0),
"SimplifiedLayerNormalization",
vec![Some(sum), Some(gamma)],
vec![normalized],
);
norm.attributes
.insert("epsilon".into(), Attribute::Float(9.999_999e-7));
graph.insert_node(norm);
let out = value(&mut graph, "out", DataType::BFloat16, hidden);
graph.insert_node(Node::new(
NodeId(0),
"Identity",
vec![Some(normalized)],
vec![out],
));
graph.add_output(out);
let next_res = value(&mut graph, "next_res", DataType::BFloat16, hidden);
let next_sum = value(&mut graph, "next_sum", DataType::BFloat16, hidden);
graph.insert_node(Node::new(
NodeId(0),
"Add",
vec![Some(sum), Some(next_res)],
vec![next_sum],
));
graph.add_input(next_res);
graph.add_output(next_sum);
graph
}
#[test]
fn folds_bf16_residual_add_norm_into_skip_node() {
let mut env = EnvVarGuard::acquire();
env.unset(SKIP_RMSNORM_FUSION_ENABLE_ENV);
for gamma_dtype in [DataType::BFloat16, DataType::Float32] {
let mut graph = bf16_skip_seam_graph(6656, gamma_dtype);
let a = value_id_by_name(&graph, "a");
let b = value_id_by_name(&graph, "b");
let gamma = value_id_by_name(&graph, "gamma");
let sum = value_id_by_name(&graph, "sum");
let normalized = value_id_by_name(&graph, "normalized");
CudaSkipRmsNormFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph
.nodes
.values()
.any(|n| n.op_type == "SimplifiedLayerNormalization"),
"default (flag unset) leaves the standalone norm (gamma={gamma_dtype:?})"
);
env.set(SKIP_RMSNORM_FUSION_ENABLE_ENV, "1");
let result = CudaSkipRmsNormFusion.run(&mut graph, &PassContext::new());
env.unset(SKIP_RMSNORM_FUSION_ENABLE_ENV);
result.unwrap();
assert!(
graph
.nodes
.values()
.all(|n| n.op_type != "SimplifiedLayerNormalization"),
"the standalone norm must be deleted (gamma={gamma_dtype:?})"
);
let skip_nodes: Vec<&Node> = graph
.nodes
.values()
.filter(|n| n.op_type == "SkipSimplifiedLayerNormalization")
.collect();
assert_eq!(
skip_nodes.len(),
1,
"exactly one skip node (gamma={gamma_dtype:?})"
);
let skip = skip_nodes[0];
assert_eq!(skip.domain, MICROSOFT_DOMAIN);
assert_eq!(
skip.inputs,
vec![Some(a), Some(b), Some(gamma)],
"skip inputs = [a, b, gamma] (gamma={gamma_dtype:?})"
);
assert_eq!(skip.outputs[0], normalized);
assert_eq!(skip.outputs[3], sum);
assert_eq!(skip.outputs.len(), 4);
assert_eq!(graph.value(sum).producer, Some(skip.id));
let sum_consumers = graph.consumers(sum);
assert_eq!(sum_consumers.len(), 1);
assert_eq!(graph.node(sum_consumers[0]).op_type, "Add");
assert_eq!(
graph.nodes.values().filter(|n| n.op_type == "Add").count(),
1,
"only the next-block residual Add remains (gamma={gamma_dtype:?})"
);
}
}
#[test]
fn skip_rmsnorm_fusion_skips_fp16_seam() {
let mut graph = bf16_skip_seam_graph(6656, DataType::Float16);
for name in ["a", "b", "sum", "normalized", "out", "next_res", "next_sum"] {
let id = value_id_by_name(&graph, name);
graph.value_mut(id).dtype = DataType::Float16;
}
let _env = EnvVarGuard::with_var(SKIP_RMSNORM_FUSION_ENABLE_ENV, "1");
let result = CudaSkipRmsNormFusion.run(&mut graph, &PassContext::new());
result.unwrap();
assert!(
graph
.nodes
.values()
.any(|n| n.op_type == "SimplifiedLayerNormalization"),
"fp16 seam must be left for CudaSkipRmsNormMatMulFusion, not folded here"
);
assert!(
graph
.nodes
.values()
.all(|n| n.op_type != "SkipSimplifiedLayerNormalization"),
);
}
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
}
fn swish_swiglu_graph(alpha: Option<f32>) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 24);
let gate = value(&mut graph, "gate", DataType::BFloat16, 7);
let up = value(&mut graph, "up", DataType::BFloat16, 7);
let swish_output = value(&mut graph, "swish", DataType::BFloat16, 7);
let output = value(&mut graph, "output", DataType::BFloat16, 7);
graph.add_input(gate);
graph.add_input(up);
let mut swish = Node::new(NodeId(0), "Swish", vec![Some(gate)], vec![swish_output]);
if let Some(alpha) = alpha {
swish
.attributes
.insert("alpha".into(), Attribute::Float(alpha));
}
graph.insert_node(swish);
graph.insert_node(Node::new(
NodeId(0),
"Mul",
vec![Some(swish_output), Some(up)],
vec![output],
));
graph.add_output(output);
graph
}
#[test]
fn fuses_standard_domain_swish_as_the_swiglu_gate() {
for alpha in [None, Some(1.0)] {
let mut graph = swish_swiglu_graph(alpha);
CudaSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.num_nodes(), 1, "alpha={alpha:?}");
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),
"alpha={alpha:?}"
);
assert!(graph.validate().is_ok());
}
}
#[test]
fn leaves_non_unit_alpha_swish_unfused() {
let mut graph = swish_swiglu_graph(Some(1.702));
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())
);
}
#[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())
);
}
fn decomposed_swiglu_glue_graph(dtype: DataType) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let x = value(&mut graph, "x", dtype, 7);
let up = value(&mut graph, "up", dtype, 7);
let sigmoid_out = value(&mut graph, "sigmoid", dtype, 7);
let silu_out = value(&mut graph, "silu", dtype, 7);
let output = value(&mut graph, "output", dtype, 7);
graph.add_input(x);
graph.add_input(up);
graph.insert_node(Node::new(
NodeId(0),
"Sigmoid",
vec![Some(x)],
vec![sigmoid_out],
));
graph.insert_node(Node::new(
NodeId(0),
"Mul",
vec![Some(x), Some(sigmoid_out)],
vec![silu_out],
));
graph.insert_node(Node::new(
NodeId(0),
"Mul",
vec![Some(silu_out), Some(up)],
vec![output],
));
graph.add_output(output);
graph
}
#[test]
fn collapses_bf16_decomposed_swiglu_glue() {
let mut graph = decomposed_swiglu_glue_graph(DataType::BFloat16);
CudaSiluFusion.run(&mut graph, &PassContext::new()).unwrap();
CudaSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.num_nodes(), 1, "Sigmoid + inner Mul must be removed");
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.attr(DECOMPOSED_SILU_ATTR).and_then(Attribute::as_int),
Some(1),
"the bf16 fused Mul must carry the decomposed marker so the runtime \
selects the byte-exact decomposed_silu_mul_bf16 kernel"
);
assert!(graph.validate().is_ok());
}
#[test]
fn leaves_fp32_decomposed_swiglu_glue_separate() {
let mut graph = decomposed_swiglu_glue_graph(DataType::Float32);
CudaSiluFusion.run(&mut graph, &PassContext::new()).unwrap();
assert_eq!(graph.num_nodes(), 3);
assert!(
graph.nodes.values().all(|node| node.op_type != "Silu"),
"fp32 has no byte-exact decomposed SiLU kernel; must stay unfused"
);
}
fn rsqrt_glue_graph(dtype: DataType) -> (Graph, ValueId, ValueId) {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let sumsq = value(&mut graph, "sumsq", dtype, 5);
let eps = value(&mut graph, "eps", dtype, 5);
let x = value(&mut graph, "x", dtype, 5);
let denom = value(&mut graph, "denom", dtype, 5);
let root = value(&mut graph, "root", dtype, 5);
let scale = value(&mut graph, "scale", dtype, 5);
let output = value(&mut graph, "output", dtype, 5);
graph.add_input(sumsq);
graph.add_input(eps);
graph.add_input(x);
graph.insert_node(Node::new(
NodeId(0),
"Add",
vec![Some(sumsq), Some(eps)],
vec![denom],
));
graph.insert_node(Node::new(NodeId(0), "Sqrt", vec![Some(denom)], vec![root]));
graph.insert_node(Node::new(
NodeId(0),
"Reciprocal",
vec![Some(root)],
vec![scale],
));
graph.insert_node(Node::new(
NodeId(0),
"Mul",
vec![Some(x), Some(scale)],
vec![output],
));
graph.add_output(output);
(graph, denom, root)
}
#[test]
fn fuses_fp16_reciprocal_of_sqrt() {
let (mut graph, denom, _root) = rsqrt_glue_graph(DataType::Float16);
CudaRsqrtFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.num_nodes(), 3, "Sqrt must be removed");
assert!(
graph.nodes.values().all(|n| n.op_type != "Sqrt"),
"Sqrt node must be gone"
);
let recip = graph
.nodes
.values()
.find(|n| n.op_type == "Reciprocal")
.expect("Reciprocal must remain");
assert!(
recip.is_default_domain(),
"node stays a standard Reciprocal"
);
assert_eq!(
recip.attr(CUDA_RSQRT_ATTR).and_then(Attribute::as_int),
Some(1),
"fused Reciprocal must carry the rsqrt marker"
);
assert_eq!(recip.inputs[0], Some(denom));
assert!(graph.validate().is_ok());
}
#[test]
fn does_not_fuse_sqrt_with_extra_consumer() {
let (mut graph, _denom, root) = rsqrt_glue_graph(DataType::Float16);
let sink = value(&mut graph, "sink", DataType::Float16, 5);
graph.insert_node(Node::new(NodeId(0), "Neg", vec![Some(root)], vec![sink]));
graph.add_output(sink);
CudaRsqrtFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph.nodes.values().any(|n| n.op_type == "Sqrt"),
"Sqrt with a second consumer must be preserved"
);
assert!(
graph
.nodes
.values()
.filter(|n| n.op_type == "Reciprocal")
.all(|n| n.attr(CUDA_RSQRT_ATTR).is_none()),
"Reciprocal must not be retagged when Sqrt escapes"
);
}
fn l2_norm_glue_graph(dtype: DataType, width: usize) -> (Graph, ValueId, ValueId) {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 13);
let x = value(&mut graph, "x", dtype, width);
let sq = graph.create_named_value("sq", dtype, vec![Dim::Static(1), Dim::Static(1)]);
let nrm = graph.create_named_value("nrm", dtype, vec![Dim::Static(1), Dim::Static(1)]);
let output = value(&mut graph, "output", dtype, width);
graph.add_input(x);
let mut rss = Node::new(NodeId(0), "ReduceSumSquare", vec![Some(x)], vec![sq]);
rss.attributes
.insert("axes".into(), Attribute::Ints(vec![-1]));
rss.attributes.insert("keepdims".into(), Attribute::Int(1));
graph.insert_node(rss);
graph.insert_node(Node::new(NodeId(0), "Sqrt", vec![Some(sq)], vec![nrm]));
graph.insert_node(Node::new(
NodeId(0),
"Div",
vec![Some(x), Some(nrm)],
vec![output],
));
graph.add_output(output);
(graph, sq, nrm)
}
#[test]
fn fuses_l2_normalize_into_lpnormalization() {
let (mut graph, _sq, _nrm) = l2_norm_glue_graph(DataType::BFloat16, 8);
let x = value_id_by_name(&graph, "x");
let output = value_id_by_name(&graph, "output");
CudaL2NormalizeFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph.num_nodes(),
1,
"only the fused LpNormalization remains"
);
assert!(
graph
.nodes
.values()
.all(|n| n.op_type != "ReduceSumSquare" && n.op_type != "Sqrt"),
"ReduceSumSquare and Sqrt must be removed"
);
let lp = graph
.nodes
.values()
.find(|n| n.op_type == "LpNormalization")
.expect("Div must be rewritten to LpNormalization");
assert!(
lp.is_default_domain(),
"LpNormalization stays default domain"
);
assert_eq!(lp.inputs, vec![Some(x)], "reads the pre-norm activation x");
assert_eq!(lp.outputs, vec![output], "keeps the Div output value");
assert_eq!(
lp.attr("p").and_then(Attribute::as_int),
Some(2),
"L2 norm ⇒ p = 2"
);
assert_eq!(
lp.attr("axis").and_then(Attribute::as_int),
Some(-1),
"axis carried over from the ReduceSumSquare axes"
);
assert_eq!(
lp.attr("fused_reduce_chain").and_then(Attribute::as_int),
Some(1),
"fused node is marked to run the byte-faithful L2-normalize kernel"
);
assert!(graph.validate().is_ok());
}
#[test]
fn l2_normalize_derives_axis_from_shapes_without_axes_attr() {
let (mut graph, _sq, _nrm) = l2_norm_glue_graph(DataType::BFloat16, 8);
let rss_id = graph
.nodes
.iter()
.find(|(_, n)| n.op_type == "ReduceSumSquare")
.map(|(id, _)| id)
.unwrap();
let mut rss = graph.node(rss_id).clone();
rss.attributes.remove("axes");
graph.replace_node(rss_id, rss);
CudaL2NormalizeFusion
.run(&mut graph, &PassContext::new())
.unwrap();
let lp = graph
.nodes
.values()
.find(|n| n.op_type == "LpNormalization")
.expect("shape-derived axis must still fuse");
assert_eq!(
lp.attr("axis").and_then(Attribute::as_int),
Some(1),
"reduced axis (dim 1 → 1) derived from shapes"
);
}
#[test]
fn does_not_fuse_l2_when_norm_escapes() {
let (mut graph, _sq, nrm) = l2_norm_glue_graph(DataType::BFloat16, 8);
let sink = value(&mut graph, "sink", DataType::BFloat16, 1);
graph.insert_node(Node::new(NodeId(0), "Neg", vec![Some(nrm)], vec![sink]));
graph.add_output(sink);
CudaL2NormalizeFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph.nodes.values().any(|n| n.op_type == "ReduceSumSquare")
&& graph.nodes.values().any(|n| n.op_type == "Sqrt")
&& graph.nodes.values().any(|n| n.op_type == "Div"),
"an escaping norm must leave the ReduceSumSquare/Sqrt/Div chain intact"
);
assert!(
graph.nodes.values().all(|n| n.op_type != "LpNormalization"),
"no fusion when the norm escapes"
);
}
#[test]
fn does_not_fuse_l2_with_mismatched_div_numerator() {
let (mut graph, _sq, _nrm) = l2_norm_glue_graph(DataType::BFloat16, 8);
let other = value(&mut graph, "other", DataType::BFloat16, 8);
graph.add_input(other);
let div_id = graph
.nodes
.iter()
.find(|(_, n)| n.op_type == "Div")
.map(|(id, _)| id)
.unwrap();
let mut div = graph.node(div_id).clone();
div.inputs[0] = Some(other); graph.replace_node(div_id, div);
CudaL2NormalizeFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph.nodes.values().all(|n| n.op_type != "LpNormalization"),
"Div numerator ≠ ReduceSumSquare input ⇒ not an L2 normalize"
);
}
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]);
}
fn const_cast_graph(
n: usize,
src_dtype: DataType,
src_bytes: Vec<u8>,
dst_dtype: DataType,
) -> (Graph, ValueId) {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let weight = graph.create_named_value("weight", src_dtype, vec![Dim::Static(n)]);
graph.set_initializer(
weight,
WeightRef::Inline(TensorData::from_raw(src_dtype, vec![n], src_bytes)),
);
let cast_out = graph.create_named_value("cast_out", dst_dtype, vec![Dim::Static(n)]);
let mut node = Node::new(NodeId(0), "Cast", vec![Some(weight)], vec![cast_out]);
node.attributes
.insert("to".into(), Attribute::Int(dst_dtype.to_onnx() as i64));
graph.insert_node(node);
let out = graph.create_named_value("out", dst_dtype, vec![Dim::Static(n)]);
graph.insert_node(Node::new(
NodeId(0),
"Identity",
vec![Some(cast_out)],
vec![out],
));
graph.add_output(out);
(graph, cast_out)
}
#[test]
fn folds_constant_cast_bf16_to_f32_byte_identical() {
let _env = EnvVarGuard::without_var(CONST_CAST_FOLD_DISABLE_ENV);
let n = 6usize;
let mut bytes = Vec::with_capacity(n * 2);
let mut expected = Vec::with_capacity(n);
for i in 0..n {
let v = half::bf16::from_f32(i as f32 * 1.5);
bytes.extend_from_slice(&v.to_le_bytes());
expected.push(v.to_f32());
}
let (mut graph, cast_out) =
const_cast_graph(n, DataType::BFloat16, bytes, DataType::Float32);
assert!(graph.value(cast_out).producer.is_some());
CudaFoldConstantCast
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(graph.nodes.values().all(|node| node.op_type != "Cast"));
assert!(graph.value(cast_out).producer.is_none());
assert_eq!(graph.value(cast_out).dtype, DataType::Float32);
let WeightRef::Inline(tensor) = graph.initializers.get(&cast_out).unwrap() else {
panic!("expected inline initializer");
};
assert_eq!(tensor.dims, vec![n]);
for (i, want) in expected.iter().enumerate() {
let got = f32::from_le_bytes([
tensor.data[i * 4],
tensor.data[i * 4 + 1],
tensor.data[i * 4 + 2],
tensor.data[i * 4 + 3],
]);
assert_eq!(got, *want, "element {i}");
}
assert!(graph.validate().is_ok());
}
#[test]
fn folds_constant_cast_f32_to_bf16_round_to_nearest_even() {
let _env = EnvVarGuard::without_var(CONST_CAST_FOLD_DISABLE_ENV);
let values = [0.0f32, 1.0, 1.5, -2.75, std::f32::consts::PI, 65_504.0];
let mut bytes = Vec::new();
for v in values {
bytes.extend_from_slice(&v.to_le_bytes());
}
let (mut graph, cast_out) =
const_cast_graph(values.len(), DataType::Float32, bytes, DataType::BFloat16);
CudaFoldConstantCast
.run(&mut graph, &PassContext::new())
.unwrap();
let WeightRef::Inline(tensor) = graph.initializers.get(&cast_out).unwrap() else {
panic!("expected inline initializer");
};
for (i, v) in values.iter().enumerate() {
let got = half::bf16::from_le_bytes([tensor.data[i * 2], tensor.data[i * 2 + 1]]);
assert_eq!(got, half::bf16::from_f32(*v), "element {i}");
}
assert!(graph.validate().is_ok());
}
#[test]
fn leaves_cast_of_non_constant() {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let x = graph.create_named_value("x", DataType::BFloat16, vec![Dim::Static(4)]);
graph.add_input(x);
let cast_out =
graph.create_named_value("cast_out", DataType::Float32, vec![Dim::Static(4)]);
let mut node = Node::new(NodeId(0), "Cast", vec![Some(x)], vec![cast_out]);
node.attributes.insert(
"to".into(),
Attribute::Int(DataType::Float32.to_onnx() as i64),
);
graph.insert_node(node);
graph.add_output(cast_out);
CudaFoldConstantCast
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph
.nodes
.values()
.filter(|node| node.op_type == "Cast")
.count(),
1
);
}
#[test]
fn leaves_sub_byte_constant_cast() {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let weight = graph.create_named_value("weight", DataType::Int4, vec![Dim::Static(4)]);
graph.set_initializer(
weight,
WeightRef::Inline(TensorData::from_raw(
DataType::Int4,
vec![4],
vec![0x21, 0x43],
)),
);
let cast_out =
graph.create_named_value("cast_out", DataType::Float32, vec![Dim::Static(4)]);
let mut node = Node::new(NodeId(0), "Cast", vec![Some(weight)], vec![cast_out]);
node.attributes.insert(
"to".into(),
Attribute::Int(DataType::Float32.to_onnx() as i64),
);
graph.insert_node(node);
graph.add_output(cast_out);
CudaFoldConstantCast
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph
.nodes
.values()
.filter(|node| node.op_type == "Cast")
.count(),
1
);
}
#[test]
fn const_cast_fold_respects_disable_env() {
let n = 4usize;
let mut bytes = Vec::new();
for i in 0..n {
bytes.extend_from_slice(&half::bf16::from_f32(i as f32).to_le_bytes());
}
let (mut graph, _cast_out) =
const_cast_graph(n, DataType::BFloat16, bytes, DataType::Float32);
let _env = EnvVarGuard::with_var(CONST_CAST_FOLD_DISABLE_ENV, "1");
let result = CudaFoldConstantCast.run(&mut graph, &PassContext::new());
result.unwrap();
assert_eq!(
graph
.nodes
.values()
.filter(|node| node.op_type == "Cast")
.count(),
1
);
}
#[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
}
fn qkv_bias_graph_with_extra_input(dtype: DataType, n: usize, slot: usize) -> (Graph, ValueId) {
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 is_group_index = slot == 4;
let extra_dtype = if is_group_index {
DataType::Int32
} else {
DataType::Uint8
};
let extra_elems = if is_group_index { k } else { n * (k / 32) };
let extra_bytes = extra_elems * if is_group_index { 4 } else { 1 };
let extra = vec1d(&mut graph, "extra", extra_dtype, extra_elems);
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],
)),
);
graph.set_initializer(
extra,
WeightRef::Inline(TensorData::from_raw(
extra_dtype,
vec![extra_elems],
vec![0u8; extra_bytes],
)),
);
graph.set_initializer(
bias,
WeightRef::Inline(TensorData::from_raw(dtype, vec![n], vec![0u8; n * 2])),
);
let mut inputs = vec![Some(x), Some(packed), Some(scales)];
while inputs.len() < slot {
inputs.push(None);
}
inputs.push(Some(extra));
graph.insert_node(matmul_nbits(inputs, 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, extra)
}
#[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_for_bfloat16_projections() {
let mut graph = gate_up_graph_dtype(QWEN_GATE_UP_K, QWEN_GATE_UP_N, DataType::BFloat16);
CudaGateUpSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph.num_nodes(),
1,
"a BFloat16 gate/up pair must collapse into one node just like fp16"
);
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),
"the BFloat16 pair must carry the paired-kernel marker"
);
}
#[test]
fn does_not_fuse_gate_up_with_mismatched_activation_and_scale_dtypes() {
let mut graph = gate_up_graph_dtype(QWEN_GATE_UP_K, QWEN_GATE_UP_N, DataType::BFloat16);
let scales = graph
.nodes
.values()
.find(|node| node.op_type == "MatMulNBits")
.and_then(|node| node.inputs[2])
.expect("the projection carries a scales input");
graph.value_mut(scales).dtype = DataType::Float16;
CudaGateUpSwiGluFusion
.run(&mut graph, &PassContext::new())
.unwrap();
asserts_not_fused(&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);
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 sigmoid_out = value(&mut graph, "sigmoid", DataType::Float16, 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);
graph.insert_node(Node::new(
NodeId(0),
"Sigmoid",
vec![Some(gate_out)],
vec![sigmoid_out],
));
graph.insert_node(Node::new(
NodeId(0),
"Mul",
vec![Some(gate_out), Some(sigmoid_out)],
vec![silu_out],
));
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(None) {
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_eq!(
fused.attr(DECOMPOSED_SILU_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::default()
.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::default()
.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::default()
.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::default()
.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::default()
.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::default()
.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::default()
.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::default()
.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(None) {
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::default()
.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::default()
.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());
}
#[test]
fn derived_min_hidden_reproduces_h200_anchor_and_scales_with_sm_count() {
assert_eq!(
derived_min_hidden(RMSNORM_FUSION_ANCHOR_SM_COUNT),
RMSNORM_FUSION_MIN_HIDDEN,
"the H200 anchor (132 SM) must map back to its calibrated 1280 floor"
);
assert_eq!(derived_min_hidden(24), 256);
assert!(derived_min_hidden(1) >= RMSNORM_FUSION_WARP_HALF4_MULTIPLE);
assert_eq!(derived_min_hidden(0), RMSNORM_FUSION_WARP_HALF4_MULTIPLE);
assert!(derived_min_hidden(264) > derived_min_hidden(132));
for sm in [1u32, 8, 24, 60, 132, 200, 264] {
assert_eq!(
derived_min_hidden(sm) % RMSNORM_FUSION_WARP_HALF4_MULTIPLE,
0
);
}
}
#[test]
fn fusion_benefit_gate_is_device_derived() {
let _env = EnvVarGuard::without_var(RMSNORM_FUSION_MIN_HIDDEN_ENV);
let h200 = CudaDeviceCapabilities::for_test((9, 0), 132, 0);
let rtx4060 = CudaDeviceCapabilities::for_test((8, 9), 24, 0);
assert!(!fusion_benefit_is_positive(896, 1, 896, Some(h200)));
assert!(fusion_benefit_is_positive(896, 1, 896, Some(rtx4060)));
assert!(!fusion_benefit_is_positive(896, 1, 896, None));
assert!(fusion_benefit_is_positive(
RMSNORM_FUSION_MIN_HIDDEN,
1,
0,
None
));
}
#[test]
fn device_aware_gate_folds_small_hidden_on_consumer_gpu() {
let _env = EnvVarGuard::without_var(RMSNORM_FUSION_MIN_HIDDEN_ENV);
let hidden = 896;
let rtx4060 = CudaDeviceCapabilities::for_test((8, 9), 24, 0);
let mut folded = skip_rms_graph(hidden, hidden + 256);
CudaSkipRmsNormMatMulFusion::for_device(Some(rtx4060))
.run(&mut folded, &PassContext::new())
.unwrap();
assert!(
folded
.nodes
.values()
.all(|node| node.op_type != "SkipSimplifiedLayerNormalization"),
"a 24-SM device must fold hidden 896 (floor 256)"
);
assert!(folded.validate().is_ok());
let mut kept = skip_rms_graph(hidden, hidden + 256);
CudaSkipRmsNormMatMulFusion::default()
.run(&mut kept, &PassContext::new())
.unwrap();
assert!(
kept.nodes
.values()
.any(|node| node.op_type == "SkipSimplifiedLayerNormalization"),
"device-unknown (H200 fallback) must keep the standalone norm at 896"
);
}
#[test]
fn rmsnorm_min_hidden_env_override_wins_over_device_derivation() {
const PROBE: &str = "ONNX_GENAI_TEST_RMSNORM_OVERRIDE_PROBE";
let mut env = EnvVarGuard::without_var(PROBE);
let derived = derived_min_hidden(24); assert_eq!(
env_usize(PROBE, derived),
derived,
"unset -> derived default"
);
env.set(PROBE, "4096");
assert_eq!(env_usize(PROBE, derived), 4096, "set -> override wins");
env.unset(PROBE);
assert_eq!(
env_usize(PROBE, derived),
derived,
"removed -> back to derived"
);
}
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(None) {
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(None) {
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::default()
.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(None) {
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 cast_wrapped_rms_norm_graph_bf16(hidden: usize) -> (Graph, ValueId, ValueId) {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let x_bf16 = value(&mut graph, "x_bf16", DataType::BFloat16, hidden);
graph.add_input(x_bf16);
let scale = vec1d(&mut graph, "scale", DataType::Float32, hidden);
graph.set_initializer(
scale,
WeightRef::Inline(TensorData::from_raw(
DataType::Float32,
vec![hidden],
vec![0u8; hidden * 4],
)),
);
let x_f32 = cast_node(&mut graph, "x_f32", x_bf16, DataType::Float32, hidden);
let norm_out = value(&mut graph, "norm_out", DataType::Float32, hidden);
let mut rms = Node::new(
NodeId(0),
"RMSNormalization",
vec![Some(x_f32), Some(scale)],
vec![norm_out],
);
rms.attributes
.insert("epsilon".into(), Attribute::Float(1e-6));
graph.insert_node(rms);
let normalized = cast_node(
&mut graph,
"normalized",
norm_out,
DataType::BFloat16,
hidden,
);
graph.add_output(normalized);
(graph, x_bf16, scale)
}
#[test]
fn drops_casts_around_fp32_wrapped_bf16_rms_norm() {
let hidden = 128;
let (mut graph, x_bf16, scale) = cast_wrapped_rms_norm_graph_bf16(hidden);
CudaDropNormalizationCasts
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(
graph.nodes.values().filter(|n| n.op_type == "Cast").count(),
0,
"all bf16 cast wrappers must be removed"
);
assert!(
graph
.nodes
.values()
.all(|n| n.op_type != "RMSNormalization"),
"narrowed RMSNormalization must be converted to SimplifiedLayerNormalization"
);
let rms = graph
.nodes
.values()
.find(|n| n.op_type == "SimplifiedLayerNormalization")
.expect("converted norm node retained");
assert_eq!(rms.domain, "", "converted norm stays in the ai.onnx domain");
assert_eq!(
rms.inputs,
vec![Some(x_bf16), Some(scale)],
"activation input rewired to bf16 source; scale untouched"
);
assert_eq!(graph.value(rms.outputs[0]).dtype, DataType::BFloat16);
assert_eq!(graph.value(scale).dtype, DataType::Float32);
assert_eq!(graph.outputs.len(), 1);
assert_eq!(graph.value(graph.outputs[0]).dtype, DataType::BFloat16);
assert!(graph.validate().is_ok());
let registry = onnx_runtime_shape_inference::InferenceRegistry::default_registry();
let opsets = graph.opset_imports.clone();
registry
.infer_graph(
&mut graph,
&opsets,
onnx_runtime_shape_inference::MergePolicy::Permissive,
)
.unwrap();
let rms_out = graph
.nodes
.values()
.find(|n| n.op_type == "SimplifiedLayerNormalization")
.expect("converted norm retained after inference")
.outputs[0];
assert_eq!(
graph.value(rms_out).dtype,
DataType::BFloat16,
"re-inference must keep the narrow (bf16) output, not the f32 scale dtype"
);
}
#[test]
fn leaves_native_bf16_rms_norm_untouched() {
let hidden = 128;
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
let x = value(&mut graph, "x", DataType::BFloat16, hidden);
graph.add_input(x);
let scale = vec1d(&mut graph, "scale", DataType::BFloat16, hidden);
graph.set_initializer(
scale,
WeightRef::Inline(TensorData::from_raw(
DataType::BFloat16,
vec![hidden],
vec![0u8; hidden * 2],
)),
);
let norm_out = value(&mut graph, "norm_out", DataType::BFloat16, hidden);
let mut rms = Node::new(
NodeId(0),
"RMSNormalization",
vec![Some(x), Some(scale)],
vec![norm_out],
);
rms.attributes
.insert("epsilon".into(), Attribute::Float(1e-6));
graph.insert_node(rms);
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 rms = graph
.nodes
.values()
.find(|n| n.op_type == "RMSNormalization")
.expect("norm retained");
assert_eq!(rms.inputs, vec![Some(x), Some(scale)]);
assert_eq!(graph.value(norm_out).dtype, DataType::BFloat16);
}
#[test]
fn leaves_mixed_narrow_dtype_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 a = value(&mut graph, "a", DataType::Float16, hidden);
let b = value(&mut graph, "b", DataType::BFloat16, hidden);
graph.add_input(a);
graph.add_input(b);
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", a, DataType::Float32, hidden);
let in1 = cast_node(&mut graph, "in1", b, 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 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 mixed fp16/bf16 activation norm must not be folded"
);
assert_eq!(graph.value(norm_out).dtype, DataType::Float32);
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());
}
struct QkvGraph {
graph: Graph,
q_weight: ValueId,
k_weight: ValueId,
v_weight: ValueId,
q_out: ValueId,
k_out: ValueId,
v_out: ValueId,
}
fn qkv_graph(k: usize, nq: usize, nk: usize, nv: usize, with_zp: bool) -> QkvGraph {
assert!(k.is_multiple_of(32));
let n_blocks = k / 32;
let blob = 16; let dt = DataType::BFloat16;
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let x = graph.create_named_value("x", dt, vec![Dim::Static(1), Dim::Static(k)]);
graph.add_input(x);
let make_proj = |graph: &mut Graph, tag: &str, n: usize, fill: u8| -> (ValueId, ValueId) {
let w = graph.create_named_value(
format!("{tag}.weight"),
DataType::Uint8,
vec![Dim::Static(n), Dim::Static(n_blocks), Dim::Static(blob)],
);
graph.set_initializer(
w,
WeightRef::Inline(TensorData::from_raw(
DataType::Uint8,
vec![n, n_blocks, blob],
vec![fill; n * n_blocks * blob],
)),
);
let s = graph.create_named_value(
format!("{tag}.scales"),
dt,
vec![Dim::Static(n * n_blocks)],
);
graph.set_initializer(
s,
WeightRef::Inline(TensorData::from_raw(
dt,
vec![n * n_blocks],
vec![fill.wrapping_add(1); n * n_blocks * 2],
)),
);
let mut inputs = vec![Some(x), Some(w), Some(s)];
if with_zp {
let zp_bytes = n * n_blocks / 2;
let z = graph.create_named_value(
format!("{tag}.zp"),
DataType::Uint8,
vec![Dim::Static(zp_bytes)],
);
graph.set_initializer(
z,
WeightRef::Inline(TensorData::from_raw(
DataType::Uint8,
vec![zp_bytes],
vec![fill.wrapping_add(2); zp_bytes],
)),
);
inputs.push(Some(z));
}
let out = graph.create_named_value(
format!("{tag}.out"),
dt,
vec![Dim::Static(1), Dim::Static(n)],
);
let mut mm = Node::new(NodeId(0), "MatMulNBits", inputs, vec![out]);
mm.domain = MICROSOFT_DOMAIN.into();
mm.attributes.insert("K".into(), Attribute::Int(k as i64));
mm.attributes.insert("N".into(), Attribute::Int(n as i64));
mm.attributes
.insert("block_size".into(), Attribute::Int(32));
mm.attributes.insert("bits".into(), Attribute::Int(4));
graph.insert_node(mm);
(w, out)
};
let (q_weight, q_out) = make_proj(&mut graph, "q", nq, 0x11);
let (k_weight, k_out) = make_proj(&mut graph, "k", nk, 0x22);
let (v_weight, v_out) = make_proj(&mut graph, "v", nv, 0x33);
let q_res = graph.create_named_value("q_res", dt, vec![Dim::Static(1), Dim::Static(nq)]);
graph.insert_node(Node::new(
NodeId(0),
"Reshape",
vec![Some(q_out)],
vec![q_res],
));
let k_res = graph.create_named_value("k_res", dt, vec![Dim::Static(1), Dim::Static(nk)]);
graph.insert_node(Node::new(
NodeId(0),
"Reshape",
vec![Some(k_out)],
vec![k_res],
));
let attn = graph.create_named_value("attn", dt, vec![Dim::Static(1), Dim::Static(nq)]);
graph.add_output(attn);
let mut gqa = Node::new(
NodeId(0),
"GroupQueryAttention",
vec![Some(q_res), Some(k_res), Some(v_out)],
vec![attn],
);
gqa.domain = MICROSOFT_DOMAIN.into();
graph.insert_node(gqa);
QkvGraph {
graph,
q_weight,
k_weight,
v_weight,
q_out,
k_out,
v_out,
}
}
fn inline_bytes(graph: &Graph, value: ValueId) -> &[u8] {
match graph.initializers.get(&value).unwrap() {
WeightRef::Inline(t) => &t.data,
WeightRef::External { .. } => panic!("expected inline"),
}
}
#[test]
fn fuses_qkv_projections_into_one_matmul_and_split() {
let mut g = qkv_graph(64, 8, 4, 4, true);
let (q_out, k_out, v_out) = (g.q_out, g.k_out, g.v_out);
CudaQkvProjectionFusion
.fuse_all(&mut g.graph, &PassContext::new())
.unwrap();
let matmuls: Vec<_> = g
.graph
.nodes
.values()
.filter(|n| n.op_type == "MatMulNBits")
.collect();
assert_eq!(matmuls.len(), 1, "three projections collapse to one GEMV");
let fused = matmuls[0];
assert_eq!(fused.attr("N").and_then(Attribute::as_int), Some(16));
assert_eq!(fused.attr("K").and_then(Attribute::as_int), Some(64));
assert_eq!(
fused.input_values().count(),
4,
"activation+weight+scales+zp"
);
let splits: Vec<_> = g
.graph
.nodes
.values()
.filter(|n| n.op_type == "Split")
.collect();
assert_eq!(splits.len(), 1);
let split = splits[0];
assert_eq!(
split.attr("split").and_then(Attribute::as_ints),
Some([8i64, 4, 4].as_slice())
);
assert_eq!(split.attr("axis").and_then(Attribute::as_int), Some(1));
assert_eq!(split.outputs, vec![q_out, k_out, v_out]);
let fused_weight = fused.inputs[1].unwrap();
let bytes = inline_bytes(&g.graph, fused_weight);
let n_blocks = 2usize;
let blob = 16usize;
assert_eq!(bytes.len(), 16 * n_blocks * blob);
assert!(bytes[..8 * n_blocks * blob].iter().all(|&b| b == 0x11));
assert!(
bytes[8 * n_blocks * blob..12 * n_blocks * blob]
.iter()
.all(|&b| b == 0x22)
);
assert!(bytes[12 * n_blocks * blob..].iter().all(|&b| b == 0x33));
for old in [g.q_weight, g.k_weight, g.v_weight] {
assert!(!g.graph.initializers.contains_key(&old));
}
g.graph.validate().unwrap();
}
#[test]
fn fuses_symmetric_qkv_without_zero_points() {
let mut g = qkv_graph(32, 6, 2, 2, false);
CudaQkvProjectionFusion
.fuse_all(&mut g.graph, &PassContext::new())
.unwrap();
let fused = g
.graph
.nodes
.values()
.find(|n| n.op_type == "MatMulNBits")
.unwrap();
assert_eq!(fused.attr("N").and_then(Attribute::as_int), Some(10));
assert_eq!(fused.input_values().count(), 3, "no zero-point slot");
g.graph.validate().unwrap();
}
#[test]
fn does_not_fuse_when_activations_differ() {
let mut g = qkv_graph(64, 8, 4, 4, true);
let other = g.graph.create_named_value(
"other",
DataType::BFloat16,
vec![Dim::Static(1), Dim::Static(64)],
);
g.graph.add_input(other);
let k_mm = g.graph.value(g.k_out).producer.unwrap();
g.graph.replace_input(k_mm, 0, Some(other));
CudaQkvProjectionFusion
.fuse_all(&mut g.graph, &PassContext::new())
.unwrap();
let matmuls = g
.graph
.nodes
.values()
.filter(|n| n.op_type == "MatMulNBits")
.count();
assert_eq!(matmuls, 3, "mismatched activation leaves projections split");
}
#[test]
fn qkv_fusion_is_opt_in_and_disabled_by_default() {
let mut env = EnvVarGuard::acquire();
let mut g = qkv_graph(64, 8, 4, 4, true);
env.unset(QKV_FUSION_ENABLE_ENV);
CudaQkvProjectionFusion
.run(&mut g.graph, &PassContext::new())
.unwrap();
let unfused = g
.graph
.nodes
.values()
.filter(|n| n.op_type == "MatMulNBits")
.count();
assert_eq!(
unfused, 3,
"default (flag unset) leaves projections unfused"
);
env.set(QKV_FUSION_ENABLE_ENV, "1");
let result = CudaQkvProjectionFusion.run(&mut g.graph, &PassContext::new());
result.unwrap();
let fused = g
.graph
.nodes
.values()
.filter(|n| n.op_type == "MatMulNBits")
.count();
assert_eq!(fused, 1, "opt-in flag enables the fusion");
}
#[test]
fn drops_identity_cast_and_rewires_consumer() {
let _env = EnvVarGuard::without_var(IDENTITY_CAST_FOLD_DISABLE_ENV);
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = value(&mut g, "x", DataType::Float32, 4);
g.add_input(x);
let y = cast_node(&mut g, "y", x, DataType::Float32, 4);
let out = value(&mut g, "out", DataType::Float32, 4);
g.insert_node(Node::new(NodeId(0), "Relu", vec![Some(y)], vec![out]));
g.add_output(out);
assert_eq!(g.nodes.values().filter(|n| n.op_type == "Cast").count(), 1);
CudaDropIdentityCast
.run(&mut g, &PassContext::new())
.unwrap();
assert_eq!(
g.nodes.values().filter(|n| n.op_type == "Cast").count(),
0,
"identity cast removed"
);
let relu = g.nodes.values().find(|n| n.op_type == "Relu").unwrap();
assert!(
relu.input_values().any(|v| v == x),
"consumer rewired onto the pre-cast value"
);
assert!(g.validate().is_ok());
}
#[test]
fn keeps_narrowing_cast() {
let _env = EnvVarGuard::without_var(IDENTITY_CAST_FOLD_DISABLE_ENV);
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = value(&mut g, "x", DataType::Float32, 4);
g.add_input(x);
let y = cast_node(&mut g, "y", x, DataType::Float16, 4);
g.add_output(y);
CudaDropIdentityCast
.run(&mut g, &PassContext::new())
.unwrap();
assert_eq!(
g.nodes.values().filter(|n| n.op_type == "Cast").count(),
1,
"narrowing cast preserved"
);
}
#[test]
fn keeps_identity_cast_feeding_graph_output() {
let _env = EnvVarGuard::without_var(IDENTITY_CAST_FOLD_DISABLE_ENV);
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = value(&mut g, "x", DataType::Float32, 4);
g.add_input(x);
let y = cast_node(&mut g, "y", x, DataType::Float32, 4);
g.add_output(y);
CudaDropIdentityCast
.run(&mut g, &PassContext::new())
.unwrap();
assert_eq!(
g.nodes.values().filter(|n| n.op_type == "Cast").count(),
1,
"graph-output cast preserved"
);
}
#[test]
fn opt_out_env_preserves_identity_cast() {
let _env = EnvVarGuard::with_var(IDENTITY_CAST_FOLD_DISABLE_ENV, "1");
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = value(&mut g, "x", DataType::Float32, 4);
g.add_input(x);
let y = cast_node(&mut g, "y", x, DataType::Float32, 4);
let out = value(&mut g, "out", DataType::Float32, 4);
g.insert_node(Node::new(NodeId(0), "Relu", vec![Some(y)], vec![out]));
g.add_output(out);
let result = CudaDropIdentityCast.run(&mut g, &PassContext::new());
result.unwrap();
assert_eq!(
g.nodes.values().filter(|n| n.op_type == "Cast").count(),
1,
"opt-out preserves the identity cast"
);
}
fn gated_delta_la_graph(heads: usize, trailing_cast: bool) -> Graph {
gated_delta_la_graph_ext(heads, trailing_cast, false)
}
fn gated_delta_la_graph_ext(
heads: usize,
trailing_cast: bool,
neg_exp_from_a_log: bool,
) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), 17);
graph.opset_imports.insert(MICROSOFT_DOMAIN.into(), 1);
let q = value(&mut graph, "q", DataType::Float32, heads);
let k = value(&mut graph, "k", DataType::Float32, heads);
let v = value(&mut graph, "v", DataType::Float32, heads);
let past = value(&mut graph, "past", DataType::Float32, heads);
let a = value(&mut graph, "a", DataType::Float32, heads);
let raw = value(&mut graph, "raw", DataType::Float32, heads);
for input in [q, k, v, past, a, raw] {
graph.add_input(input);
}
let dt_bias = vec1d(&mut graph, "dt_bias", DataType::Float32, heads);
let coeff_init_name = if neg_exp_from_a_log {
"A_log"
} else {
"neg_exp_A"
};
let coeff_init = vec1d(&mut graph, coeff_init_name, DataType::Float32, heads);
for init in [dt_bias, coeff_init] {
graph.set_initializer(
init,
WeightRef::Inline(TensorData::from_raw(
DataType::Float32,
vec![heads],
vec![0u8; heads * DataType::Float32.byte_size()],
)),
);
}
let neg_exp_a = if neg_exp_from_a_log {
let exp_out = value(&mut graph, "exp_out", DataType::Float32, heads);
graph.insert_node(Node::new(
NodeId(0),
"Exp",
vec![Some(coeff_init)],
vec![exp_out],
));
let neg_out = value(&mut graph, "neg_exp_A", DataType::Float32, heads);
graph.insert_node(Node::new(
NodeId(0),
"Neg",
vec![Some(exp_out)],
vec![neg_out],
));
neg_out
} else {
coeff_init
};
let add_out = value(&mut graph, "add_out", DataType::Float32, heads);
graph.insert_node(Node::new(
NodeId(0),
"Add",
vec![Some(a), Some(dt_bias)],
vec![add_out],
));
let sp_out = value(&mut graph, "sp_out", DataType::Float32, heads);
graph.insert_node(Node::new(
NodeId(0),
"Softplus",
vec![Some(add_out)],
vec![sp_out],
));
let mul_out = value(&mut graph, "mul_out", DataType::Float32, heads);
graph.insert_node(Node::new(
NodeId(0),
"Mul",
vec![Some(neg_exp_a), Some(sp_out)],
vec![mul_out],
));
let decay = if trailing_cast {
let cast_out = value(&mut graph, "decay", DataType::Float32, heads);
let mut cast = Node::new(NodeId(0), "Cast", vec![Some(mul_out)], vec![cast_out]);
cast.attributes.insert(
"to".into(),
Attribute::Int(DataType::Float32.to_onnx() as i64),
);
graph.insert_node(cast);
cast_out
} else {
mul_out
};
let beta = value(&mut graph, "beta", DataType::Float32, heads);
graph.insert_node(Node::new(NodeId(0), "Sigmoid", vec![Some(raw)], vec![beta]));
let out = value(&mut graph, "out", DataType::Float32, heads);
let present = value(&mut graph, "present", DataType::Float32, heads);
let mut la = Node::new(
NodeId(0),
"LinearAttention",
vec![
Some(q),
Some(k),
Some(v),
Some(past),
Some(decay),
Some(beta),
],
vec![out, present],
);
la.domain = MICROSOFT_DOMAIN.into();
graph.insert_node(la);
graph.add_output(out);
graph.add_output(present);
graph
}
#[test]
fn folds_beta_sigmoid_and_decay_softplus_into_linear_attention() {
let _env = EnvVarGuard::without_var(LINEAR_ATTENTION_GATING_DISABLE_ENV);
for trailing_cast in [false, true] {
let mut graph = gated_delta_la_graph(4, trailing_cast);
let a = value_id_by_name(&graph, "a");
let raw = value_id_by_name(&graph, "raw");
let dt_bias = value_id_by_name(&graph, "dt_bias");
let neg_exp_a = value_id_by_name(&graph, "neg_exp_A");
CudaLinearAttentionGatingFusion
.run(&mut graph, &PassContext::new())
.unwrap();
for op in ["Sigmoid", "Softplus", "Add", "Mul", "Cast"] {
assert!(
graph.nodes.values().all(|n| n.op_type != op),
"{op} must be folded away (trailing_cast={trailing_cast})"
);
}
let la = graph
.nodes
.values()
.find(|n| n.op_type == "LinearAttention")
.unwrap();
assert_eq!(
la.attr(FUSE_BETA_SIGMOID_ATTR).and_then(Attribute::as_int),
Some(1)
);
assert_eq!(
la.attr(FUSE_DECAY_SOFTPLUS_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert_eq!(la.inputs[5], Some(raw));
assert_eq!(la.inputs[4], Some(a));
assert_eq!(la.inputs.len(), 8);
assert_eq!(la.inputs[6], Some(dt_bias));
assert_eq!(la.inputs[7], Some(neg_exp_a));
}
}
#[test]
fn folds_neg_exp_a_log_chain_into_linear_attention() {
let _env = EnvVarGuard::without_var(LINEAR_ATTENTION_GATING_DISABLE_ENV);
for trailing_cast in [false, true] {
let mut graph = gated_delta_la_graph_ext(4, trailing_cast, true);
let a = value_id_by_name(&graph, "a");
let dt_bias = value_id_by_name(&graph, "dt_bias");
let a_log = value_id_by_name(&graph, "A_log");
CudaLinearAttentionGatingFusion
.run(&mut graph, &PassContext::new())
.unwrap();
for op in ["Sigmoid", "Softplus", "Add", "Mul", "Cast", "Exp", "Neg"] {
assert!(
graph.nodes.values().all(|n| n.op_type != op),
"{op} must be folded away (trailing_cast={trailing_cast})"
);
}
let la = graph
.nodes
.values()
.find(|n| n.op_type == "LinearAttention")
.unwrap();
assert_eq!(
la.attr(FUSE_DECAY_SOFTPLUS_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert_eq!(
la.attr(FUSE_NEG_EXP_ATTR).and_then(Attribute::as_int),
Some(1)
);
assert_eq!(la.inputs[4], Some(a));
assert_eq!(la.inputs.len(), 8);
assert_eq!(la.inputs[6], Some(dt_bias));
assert_eq!(la.inputs[7], Some(a_log));
}
}
#[test]
fn precomputed_neg_exp_a_initializer_does_not_set_neg_exp_marker() {
let _env = EnvVarGuard::without_var(LINEAR_ATTENTION_GATING_DISABLE_ENV);
let mut graph = gated_delta_la_graph(4, false);
CudaLinearAttentionGatingFusion
.run(&mut graph, &PassContext::new())
.unwrap();
let la = graph
.nodes
.values()
.find(|n| n.op_type == "LinearAttention")
.unwrap();
assert_eq!(
la.attr(FUSE_DECAY_SOFTPLUS_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert!(la.attr(FUSE_NEG_EXP_ATTR).is_none());
}
#[test]
fn leaves_beta_gate_when_it_escapes_to_a_second_consumer() {
let _env = EnvVarGuard::without_var(LINEAR_ATTENTION_GATING_DISABLE_ENV);
let mut graph = gated_delta_la_graph(4, false);
let beta = value_id_by_name(&graph, "beta");
let sink = value(&mut graph, "beta_sink", DataType::Float32, 4);
graph.insert_node(Node::new(NodeId(0), "Relu", vec![Some(beta)], vec![sink]));
graph.add_output(sink);
CudaLinearAttentionGatingFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert!(
graph.nodes.values().any(|n| n.op_type == "Sigmoid"),
"the escaping beta Sigmoid is preserved"
);
let la = graph
.nodes
.values()
.find(|n| n.op_type == "LinearAttention")
.unwrap();
assert!(la.attr(FUSE_BETA_SIGMOID_ATTR).is_none());
assert_eq!(
la.attr(FUSE_DECAY_SOFTPLUS_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert_eq!(la.inputs.len(), 8);
}
#[test]
fn opt_out_env_preserves_exported_gate_chains() {
let mut graph = gated_delta_la_graph(4, true);
let _env = EnvVarGuard::with_var(LINEAR_ATTENTION_GATING_DISABLE_ENV, "1");
let result = CudaLinearAttentionGatingFusion.run(&mut graph, &PassContext::new());
result.unwrap();
for op in ["Sigmoid", "Softplus", "Add", "Mul", "Cast"] {
assert!(
graph.nodes.values().any(|n| n.op_type == op),
"opt-out preserves the standalone {op}"
);
}
let la = graph
.nodes
.values()
.find(|n| n.op_type == "LinearAttention")
.unwrap();
assert!(la.attr(FUSE_BETA_SIGMOID_ATTR).is_none());
assert!(la.attr(FUSE_DECAY_SOFTPLUS_ATTR).is_none());
assert_eq!(la.inputs.len(), 6);
}
#[test]
fn folds_bias_for_asymmetrically_quantized_weights() {
let (mut graph, zp) = qkv_bias_graph_with_extra_input(DataType::Float16, 1152, 3);
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().expect("fused node");
assert_eq!(
fused
.attr(MATMUL_NBITS_FOLDED_BIAS_ATTR)
.and_then(Attribute::as_int),
Some(1)
);
assert_eq!(
fused.inputs[3],
Some(zp),
"zero-points must survive the fold"
);
assert!(fused.inputs[5].is_some(), "bias must be wired at index 5");
assert!(graph.validate().is_ok());
}
#[test]
fn does_not_fold_bias_when_a_group_index_is_present() {
let (mut graph, _gidx) = qkv_bias_graph_with_extra_input(DataType::Float16, 1152, 4);
CudaMatMulNBitsBiasFusion
.run(&mut graph, &PassContext::new())
.unwrap();
assert_eq!(graph.num_nodes(), 2, "a group index must block the fold");
}
}