Files
rustytorch/crates/production/rtx-onnx-codegen/src/ops/pooling.rs
T
osobhandClaude Opus 4.6 02d382d5f6 style: apply rustfmt across all crates and demos
Consistent formatting pass: line wrapping, import sorting, trailing
whitespace removal, let-chain indentation, merged derive attributes,
and unsafe block reformatting.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-04-12 07:01:58 -07:00

108 lines
3.6 KiB
Rust

//! Pooling operator code generation.
use proc_macro2::TokenStream;
use quote::quote;
use super::sanitize_name;
use crate::error::Result;
use crate::ir::Node;
/// Default kernel shape for pooling.
const DEFAULT_KERNEL: &[i64] = &[2, 2];
/// Default padding for pooling.
const DEFAULT_PADS: &[i64] = &[0, 0, 0, 0];
/// Generate MaxPool operation.
pub fn generate_max_pool(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 attributes
let kernel_shape = node.get_ints("kernel_shape").unwrap_or(DEFAULT_KERNEL);
let strides = node.get_ints("strides").unwrap_or(kernel_shape);
let pads = node.get_ints("pads").unwrap_or(DEFAULT_PADS);
let ceil_mode = node.get_int("ceil_mode").unwrap_or(0) != 0;
let kernel_h = kernel_shape.first().copied().unwrap_or(2) as usize;
let kernel_w = kernel_shape.get(1).copied().unwrap_or(2) as usize;
let stride_h = strides.first().copied().unwrap_or(kernel_h as i64) as usize;
let stride_w = strides.get(1).copied().unwrap_or(kernel_w as i64) as usize;
let pad_h = pads.first().copied().unwrap_or(0) as usize;
let pad_w = pads.get(1).copied().unwrap_or(0) as usize;
Ok(quote! {
let #out_ident = #in_ident.max_pool2d(
[#kernel_h, #kernel_w],
[#stride_h, #stride_w],
[#pad_h, #pad_w],
#ceil_mode,
)?;
})
}
/// Generate AveragePool operation.
pub fn generate_avg_pool(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 attributes
let kernel_shape = node.get_ints("kernel_shape").unwrap_or(DEFAULT_KERNEL);
let strides = node.get_ints("strides").unwrap_or(kernel_shape);
let pads = node.get_ints("pads").unwrap_or(DEFAULT_PADS);
let count_include_pad = node.get_int("count_include_pad").unwrap_or(0) != 0;
let ceil_mode = node.get_int("ceil_mode").unwrap_or(0) != 0;
let kernel_h = kernel_shape.first().copied().unwrap_or(2) as usize;
let kernel_w = kernel_shape.get(1).copied().unwrap_or(2) as usize;
let stride_h = strides.first().copied().unwrap_or(kernel_h as i64) as usize;
let stride_w = strides.get(1).copied().unwrap_or(kernel_w as i64) as usize;
let pad_h = pads.first().copied().unwrap_or(0) as usize;
let pad_w = pads.get(1).copied().unwrap_or(0) as usize;
Ok(quote! {
let #out_ident = #in_ident.avg_pool2d(
[#kernel_h, #kernel_w],
[#stride_h, #stride_w],
[#pad_h, #pad_w],
#ceil_mode,
#count_include_pad,
)?;
})
}
/// Generate GlobalAveragePool operation.
pub fn generate_global_avg_pool(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 = #in_ident.global_avg_pool2d()?;
})
}
/// Generate GlobalMaxPool operation.
pub fn generate_global_max_pool(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 = #in_ident.global_max_pool2d()?;
})
}