//! 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 { 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 { 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 { 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 { 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()?; }) }