Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -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(),
)?;
})
}