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]>
108 lines
3.6 KiB
Rust
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()?;
|
|
})
|
|
}
|