Initial commit
This commit is contained in:
@@ -0,0 +1,496 @@
|
||||
//! ONNX operator code generation.
|
||||
//!
|
||||
//! This module contains code generators for individual ONNX operators.
|
||||
|
||||
mod arithmetic;
|
||||
mod activation;
|
||||
mod conv;
|
||||
mod matmul;
|
||||
mod normalization;
|
||||
mod pooling;
|
||||
mod reshape;
|
||||
|
||||
pub use arithmetic::*;
|
||||
pub use activation::*;
|
||||
pub use conv::*;
|
||||
pub use matmul::*;
|
||||
pub use normalization::*;
|
||||
pub use pooling::*;
|
||||
pub use reshape::*;
|
||||
|
||||
use proc_macro2::TokenStream;
|
||||
use quote::quote;
|
||||
|
||||
use crate::ir::{Node, NodeKind, OnnxGraph};
|
||||
use crate::error::{Error, Result};
|
||||
|
||||
/// Generate code for a single node.
|
||||
pub fn generate_node_code(node: &Node, graph: &OnnxGraph) -> Result<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
// In inference mode, dropout is identity
|
||||
generate_identity(node)
|
||||
}
|
||||
|
||||
/// Generate constant tensor.
|
||||
fn generate_constant(node: &Node, graph: &OnnxGraph) -> Result<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<i64> = 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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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<TokenStream> {
|
||||
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::<Vec<_>>(),
|
||||
&[#in_ident.ndim()],
|
||||
rtx_core::DType::I64,
|
||||
#in_ident.device(),
|
||||
)?;
|
||||
})
|
||||
}
|
||||
|
||||
/// Generate size operation.
|
||||
fn generate_size(node: &Node) -> Result<TokenStream> {
|
||||
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(),
|
||||
)?;
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user