//! Convolution operator code generation. use proc_macro2::TokenStream; use quote::quote; use super::sanitize_name; use crate::error::Result; use crate::ir::{Node, OnnxGraph}; /// Default strides for 2D convolution. const DEFAULT_STRIDES: &[i64] = &[1, 1]; /// Default padding for 2D convolution. const DEFAULT_PADS: &[i64] = &[0, 0, 0, 0]; /// Default dilations for 2D convolution. const DEFAULT_DILATIONS: &[i64] = &[1, 1]; /// Default output padding for transposed convolution. const DEFAULT_OUTPUT_PADDING: &[i64] = &[0, 0]; /// Generate Conv operation. pub fn generate_conv(node: &Node, _graph: &OnnxGraph) -> Result { let input = &node.inputs[0]; let weight = &node.inputs[1]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let weight_ident = quote::format_ident!("{}", sanitize_name(weight)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); // Get attributes let strides = node.get_ints("strides").unwrap_or(DEFAULT_STRIDES); let pads = node.get_ints("pads").unwrap_or(DEFAULT_PADS); let dilations = node.get_ints("dilations").unwrap_or(DEFAULT_DILATIONS); let group = node.get_int("group").unwrap_or(1) as usize; // Handle optional bias let has_bias = node.inputs.len() > 2 && !node.inputs[2].is_empty(); // Convert pads to padding tuple (assuming symmetric padding) let pad_h = pads.first().copied().unwrap_or(0) as usize; let pad_w = pads.get(1).copied().unwrap_or(0) as usize; let stride_h = strides.first().copied().unwrap_or(1) as usize; let stride_w = strides.get(1).copied().unwrap_or(1) as usize; let dilation_h = dilations.first().copied().unwrap_or(1) as usize; let dilation_w = dilations.get(1).copied().unwrap_or(1) as usize; 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::conv2d( &#in_ident, &self.#weight_ident, #bias_arg, (#stride_h, #stride_w), (#pad_h, #pad_w), (#dilation_h, #dilation_w), #group, )?; }) } /// Generate ConvTranspose operation. pub fn generate_conv_transpose(node: &Node, _graph: &OnnxGraph) -> Result { let input = &node.inputs[0]; let weight = &node.inputs[1]; let output = &node.outputs[0]; let in_ident = quote::format_ident!("{}", sanitize_name(input)); let weight_ident = quote::format_ident!("{}", sanitize_name(weight)); let out_ident = quote::format_ident!("{}", sanitize_name(output)); // Get attributes let strides = node.get_ints("strides").unwrap_or(DEFAULT_STRIDES); let pads = node.get_ints("pads").unwrap_or(DEFAULT_PADS); let output_padding = node .get_ints("output_padding") .unwrap_or(DEFAULT_OUTPUT_PADDING); let dilations = node.get_ints("dilations").unwrap_or(DEFAULT_DILATIONS); let group = node.get_int("group").unwrap_or(1) as usize; // Handle optional bias let has_bias = node.inputs.len() > 2 && !node.inputs[2].is_empty(); let pad_h = pads.first().copied().unwrap_or(0) as usize; let pad_w = pads.get(1).copied().unwrap_or(0) as usize; let stride_h = strides.first().copied().unwrap_or(1) as usize; let stride_w = strides.get(1).copied().unwrap_or(1) as usize; let out_pad_h = output_padding.first().copied().unwrap_or(0) as usize; let out_pad_w = output_padding.get(1).copied().unwrap_or(0) as usize; let dilation_h = dilations.first().copied().unwrap_or(1) as usize; let dilation_w = dilations.get(1).copied().unwrap_or(1) as usize; 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::conv_transpose2d( &#in_ident, &self.#weight_ident, #bias_arg, (#stride_h, #stride_w), (#pad_h, #pad_w), (#out_pad_h, #out_pad_w), (#dilation_h, #dilation_w), #group, )?; }) }