//! ONNX operator code generation. //! //! This module contains code generators for individual ONNX operators. mod activation; mod arithmetic; mod conv; mod matmul; mod normalization; mod pooling; mod reshape; pub use activation::*; pub use arithmetic::*; pub use conv::*; pub use matmul::*; pub use normalization::*; pub use pooling::*; pub use reshape::*; use proc_macro2::TokenStream; use quote::quote; use crate::error::{Error, Result}; use crate::ir::{Node, NodeKind, OnnxGraph}; /// Generate code for a single node. pub fn generate_node_code(node: &Node, graph: &OnnxGraph) -> Result { match &node.kind { // Arithmetic NodeKind::Add => generate_binary_op(node, "add"), NodeKind::Sub => generate_binary_op(node, "sub"), NodeKind::Mul => generate_binary_op(node, "mul"), NodeKind::Div => generate_binary_op(node, "div"), NodeKind::Neg => generate_unary_op(node, "neg"), NodeKind::Abs => generate_unary_op(node, "abs"), NodeKind::Sqrt => generate_unary_op(node, "sqrt"), NodeKind::Exp => generate_unary_op(node, "exp"), NodeKind::Log => generate_unary_op(node, "log"), NodeKind::Pow => generate_pow(node), NodeKind::Floor => generate_unary_op(node, "floor"), NodeKind::Ceil => generate_unary_op(node, "ceil"), NodeKind::Round => generate_unary_op(node, "round"), NodeKind::Clip => generate_clip(node), NodeKind::Min => generate_min_max(node, "min"), NodeKind::Max => generate_min_max(node, "max"), // Activations NodeKind::Relu => generate_unary_op(node, "relu"), NodeKind::Sigmoid => generate_unary_op(node, "sigmoid"), NodeKind::Tanh => generate_unary_op(node, "tanh"), NodeKind::Softmax => generate_softmax(node), NodeKind::LogSoftmax => generate_log_softmax(node), NodeKind::Gelu => generate_unary_op(node, "gelu"), NodeKind::Silu => generate_unary_op(node, "silu"), NodeKind::LeakyRelu => generate_leaky_relu(node), NodeKind::Elu => generate_elu(node), NodeKind::HardSwish => generate_unary_op(node, "hardswish"), NodeKind::HardSigmoid => generate_unary_op(node, "hardsigmoid"), // Matrix operations NodeKind::MatMul => generate_matmul(node), NodeKind::Gemm => generate_gemm(node, graph), // Normalization NodeKind::BatchNormalization => generate_batch_norm(node, graph), NodeKind::LayerNormalization => generate_layer_norm(node, graph), NodeKind::InstanceNormalization => generate_instance_norm(node, graph), NodeKind::GroupNormalization => generate_group_norm(node, graph), // Convolution NodeKind::Conv => generate_conv(node, graph), NodeKind::ConvTranspose => generate_conv_transpose(node, graph), // Pooling NodeKind::MaxPool => generate_max_pool(node), NodeKind::AveragePool => generate_avg_pool(node), NodeKind::GlobalAveragePool => generate_global_avg_pool(node), NodeKind::GlobalMaxPool => generate_global_max_pool(node), // Shape operations NodeKind::Reshape => generate_reshape(node, graph), NodeKind::Transpose => generate_transpose(node), NodeKind::Flatten => generate_flatten(node), NodeKind::Concat => generate_concat(node), NodeKind::Squeeze => generate_squeeze(node), NodeKind::Unsqueeze => generate_unsqueeze(node), NodeKind::Split => generate_split(node), NodeKind::Slice => generate_slice(node, graph), NodeKind::Pad => generate_pad(node, graph), NodeKind::Gather => generate_gather(node), NodeKind::Expand => generate_expand(node), NodeKind::Tile => generate_tile(node), // Reduction operations NodeKind::ReduceSum => generate_reduce(node, "sum"), NodeKind::ReduceMean => generate_reduce(node, "mean"), NodeKind::ReduceMax => generate_reduce(node, "max"), NodeKind::ReduceMin => generate_reduce(node, "min"), NodeKind::ReduceProd => generate_reduce(node, "prod"), // Comparison operations NodeKind::Equal => generate_comparison(node, "eq"), NodeKind::Greater => generate_comparison(node, "gt"), NodeKind::Less => generate_comparison(node, "lt"), NodeKind::GreaterOrEqual => generate_comparison(node, "ge"), NodeKind::LessOrEqual => generate_comparison(node, "le"), NodeKind::Not => generate_unary_op(node, "logical_not"), NodeKind::And => generate_binary_op(node, "logical_and"), NodeKind::Or => generate_binary_op(node, "logical_or"), NodeKind::Where => generate_where(node), // Cast operations NodeKind::Cast => generate_cast(node), NodeKind::CastLike => generate_cast_like(node), // Other NodeKind::Identity => generate_identity(node), NodeKind::Dropout => generate_dropout(node), NodeKind::Constant => generate_constant(node, graph), NodeKind::Shape => generate_shape(node), NodeKind::Size => generate_size(node), // Unsupported NodeKind::Custom(op) => Err(Error::UnsupportedOperator(op.clone())), other => Err(Error::UnsupportedOperator(format!("{:?}", other))), } } /// Generate identity (passthrough) operation. fn generate_identity(node: &Node) -> Result { let input = &node.inputs[0]; let output = &node.outputs[0]; let input_ident = quote::format_ident!("{}", sanitize_name(input)); let output_ident = quote::format_ident!("{}", sanitize_name(output)); Ok(quote! { let #output_ident = #input_ident.clone(); }) } /// Generate dropout (passthrough in inference). fn generate_dropout(node: &Node) -> Result { // In inference mode, dropout is identity generate_identity(node) } /// Generate constant tensor. fn generate_constant(node: &Node, graph: &OnnxGraph) -> Result { let output = &node.outputs[0]; let output_ident = quote::format_ident!("{}", sanitize_name(output)); // The constant value should be in attributes or initializers if let Some(tensor) = graph.get_initializer(output) { let shape: Vec<_> = tensor.shape.static_dims(); Ok(quote! { let #output_ident = self.#output_ident.clone(); }) } else { // TODO: Handle inline constant attributes Ok(quote! { // TODO: Constant generation let #output_ident = todo!("Constant tensor"); }) } } /// Sanitize tensor name to be a valid Rust identifier. pub fn sanitize_name(name: &str) -> String { let mut result = String::new(); for (i, c) in name.chars().enumerate() { if c.is_alphanumeric() { if i == 0 && c.is_numeric() { result.push('_'); } result.push(c); } else { result.push('_'); } } // Handle Rust keywords match result.as_str() { "type" | "fn" | "let" | "mut" | "ref" | "self" | "Self" | "mod" | "pub" | "use" | "struct" | "enum" | "trait" | "impl" | "for" | "while" | "loop" | "if" | "else" | "match" | "return" | "break" | "continue" | "move" | "box" | "where" | "async" | "await" | "dyn" | "abstract" | "become" | "const" | "crate" | "do" | "extern" | "final" | "in" | "macro" | "override" | "priv" | "static" | "super" | "try" | "typeof" | "unsafe" | "unsized" | "virtual" | "yield" => { result.push('_'); } _ => {} } if result.is_empty() { result = "tensor".to_string(); } result } /// Generate clip operation. fn generate_clip(node: &Node) -> Result { let input = &node.inputs[0]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); // Get min/max from inputs (opset 11+) or attributes let min_val = node.get_float("min").unwrap_or(f32::MIN); let max_val = node.get_float("max").unwrap_or(f32::MAX); Ok(quote! { let #out_ident = #in_ident.clamp(#min_val, #max_val)?; }) } /// Generate min/max operation. fn generate_min_max(node: &Node, op: &str) -> Result { let output = &node.outputs[0]; let out_ident = quote::format_ident!("{}", sanitize_name(output)); let input_idents: Vec<_> = node .inputs .iter() .filter(|s| !s.is_empty()) .map(|s| quote::format_ident!("{}", sanitize_name(s))) .collect(); let op_ident = quote::format_ident!("{}", op); if input_idents.len() == 2 { let a = &input_idents[0]; let b = &input_idents[1]; Ok(quote! { let #out_ident = #a.#op_ident(&#b)?; }) } else { // Multiple inputs - chain operations with fold let first = &input_idents[0]; let rest = &input_idents[1..]; Ok(quote! { let #out_ident = { let mut result = #first.clone(); #(result = result.#op_ident(&#rest)?;)* result }; }) } } /// Generate log softmax operation. fn generate_log_softmax(node: &Node) -> Result { let input = &node.inputs[0]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); let axis = node.get_int("axis").unwrap_or(-1); Ok(quote! { let #out_ident = #in_ident.log_softmax(#axis)?; }) } /// Generate ELU operation. fn generate_elu(node: &Node) -> Result { let input = &node.inputs[0]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); let alpha = node.get_float("alpha").unwrap_or(1.0); Ok(quote! { let #out_ident = #in_ident.elu(#alpha)?; }) } /// Generate instance normalization operation. fn generate_instance_norm(node: &Node, _graph: &OnnxGraph) -> Result { let input = &node.inputs[0]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); let epsilon = node.get_float("epsilon").unwrap_or(1e-5); let has_scale = node.inputs.len() > 1 && !node.inputs[1].is_empty(); let has_bias = node.inputs.len() > 2 && !node.inputs[2].is_empty(); let scale_arg = if has_scale { let scale_ident = quote::format_ident!("{}", sanitize_name(&node.inputs[1])); quote! { Some(&self.#scale_ident) } } else { quote! { None } }; let bias_arg = if has_bias { let bias_ident = quote::format_ident!("{}", sanitize_name(&node.inputs[2])); quote! { Some(&self.#bias_ident) } } else { quote! { None } }; Ok(quote! { let #out_ident = rtx_nn::functional::instance_norm( &#in_ident, #scale_arg, #bias_arg, #epsilon, )?; }) } /// Generate group normalization operation. fn generate_group_norm(node: &Node, _graph: &OnnxGraph) -> Result { let input = &node.inputs[0]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); let epsilon = node.get_float("epsilon").unwrap_or(1e-5); let num_groups = node.get_int("num_groups").unwrap_or(32) as usize; let has_scale = node.inputs.len() > 1 && !node.inputs[1].is_empty(); let has_bias = node.inputs.len() > 2 && !node.inputs[2].is_empty(); let scale_arg = if has_scale { let scale_ident = quote::format_ident!("{}", sanitize_name(&node.inputs[1])); quote! { Some(&self.#scale_ident) } } else { quote! { None } }; let bias_arg = if has_bias { let bias_ident = quote::format_ident!("{}", sanitize_name(&node.inputs[2])); quote! { Some(&self.#bias_ident) } } else { quote! { None } }; Ok(quote! { let #out_ident = rtx_nn::functional::group_norm( &#in_ident, #num_groups, #scale_arg, #bias_arg, #epsilon, )?; }) } /// Generate reduce operation. fn generate_reduce(node: &Node, op: &str) -> Result { let input = &node.inputs[0]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); let axes = node.get_ints("axes"); let keepdims = node.get_int("keepdims").unwrap_or(1) != 0; let op_ident = quote::format_ident!("{}", op); if let Some(axes) = axes { let axes_i64: Vec = axes.to_vec(); Ok(quote! { let #out_ident = #in_ident.#op_ident(&[#(#axes_i64),*], #keepdims)?; }) } else { // Reduce all dimensions Ok(quote! { let #out_ident = #in_ident.#op_ident(None, #keepdims)?; }) } } /// Generate comparison operation. fn generate_comparison(node: &Node, op: &str) -> Result { let input_a = &node.inputs[0]; let input_b = &node.inputs[1]; let output = &node.outputs[0]; let a_ident = quote::format_ident!("{}", sanitize_name(input_a)); let b_ident = quote::format_ident!("{}", sanitize_name(input_b)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); let op_ident = quote::format_ident!("{}", op); Ok(quote! { let #out_ident = #a_ident.#op_ident(&#b_ident)?; }) } /// Generate where (conditional select) operation. fn generate_where(node: &Node) -> Result { let condition = &node.inputs[0]; let x = &node.inputs[1]; let y = &node.inputs[2]; let output = &node.outputs[0]; let cond_ident = quote::format_ident!("{}", sanitize_name(condition)); let x_ident = quote::format_ident!("{}", sanitize_name(x)); let y_ident = quote::format_ident!("{}", sanitize_name(y)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); Ok(quote! { let #out_ident = rtx_core::Tensor::where_cond(&#cond_ident, &#x_ident, &#y_ident)?; }) } /// Generate cast operation. fn generate_cast(node: &Node) -> Result { let input = &node.inputs[0]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); // Get target data type let to = node.get_int("to").unwrap_or(1); // 1 = FLOAT let dtype = match to { 1 => quote! { rtx_core::DType::F32 }, 2 => quote! { rtx_core::DType::U8 }, 3 => quote! { rtx_core::DType::I8 }, 5 => quote! { rtx_core::DType::I16 }, 6 => quote! { rtx_core::DType::I32 }, 7 => quote! { rtx_core::DType::I64 }, 9 => quote! { rtx_core::DType::Bool }, 10 => quote! { rtx_core::DType::F16 }, 11 => quote! { rtx_core::DType::F64 }, 16 => quote! { rtx_core::DType::BF16 }, _ => quote! { rtx_core::DType::F32 }, }; Ok(quote! { let #out_ident = #in_ident.to_dtype(#dtype)?; }) } /// Generate cast-like operation. fn generate_cast_like(node: &Node) -> Result { let input = &node.inputs[0]; let target = &node.inputs[1]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let target_ident = quote::format_ident!("{}", sanitize_name(target)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); Ok(quote! { let #out_ident = #in_ident.to_dtype(#target_ident.dtype())?; }) } /// Generate shape operation. fn generate_shape(node: &Node) -> Result { let input = &node.inputs[0]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); Ok(quote! { let #out_ident = rtx_core::Tensor::from_slice( &#in_ident.shape().iter().map(|&d| d as i64).collect::>(), &[#in_ident.ndim()], rtx_core::DType::I64, #in_ident.device(), )?; }) } /// Generate size operation. fn generate_size(node: &Node) -> Result { let input = &node.inputs[0]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); Ok(quote! { let #out_ident = rtx_core::Tensor::scalar( #in_ident.numel() as i64, rtx_core::DType::I64, #in_ident.device(), )?; }) }