use std::{
fmt::Display,
sync::atomic::{AtomicUsize, Ordering},
};
use nove_tensor::{DType, Device, Tensor};
use crate::{Model, ModelError};
static ID: AtomicUsize = AtomicUsize::new(0);
#[derive(Debug, Clone)]
pub struct MaxPool2d {
kernel_size: (usize, usize),
stride: (usize, usize),
id: usize,
}
impl MaxPool2d {
pub fn kernel_size(&self) -> (usize, usize) {
self.kernel_size
}
pub fn stride(&self) -> (usize, usize) {
self.stride
}
}
impl Model for MaxPool2d {
type Input = Tensor;
type Output = Tensor;
fn forward(&mut self, input: Self::Input) -> Result<Self::Output, ModelError> {
let y = input.max_pool2d(self.kernel_size, self.stride)?;
Ok(y)
}
fn require_grad(&mut self, _grad_enabled: bool) -> Result<(), ModelError> {
Ok(())
}
fn to_device(&mut self, _device: &Device) -> Result<(), ModelError> {
Ok(())
}
fn to_dtype(&mut self, _dtype: &DType) -> Result<(), ModelError> {
Ok(())
}
fn parameters(&self) -> Result<Vec<Tensor>, ModelError> {
Ok(vec![])
}
fn named_parameters(&self) -> Result<std::collections::HashMap<String, Tensor>, ModelError> {
Ok(std::collections::HashMap::new())
}
}
impl Display for MaxPool2d {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"maxpool2d.{}(kernel_size={:?}, stride={:?})",
self.id, self.kernel_size, self.stride,
)
}
}
pub struct MaxPool2dBuilder {
kernel_size: Option<(usize, usize)>,
stride: Option<(usize, usize)>,
}
impl Default for MaxPool2dBuilder {
fn default() -> Self {
Self {
kernel_size: None,
stride: None,
}
}
}
impl MaxPool2dBuilder {
pub fn kernel_size(&mut self, kernel_size: (usize, usize)) -> &mut Self {
self.kernel_size = Some(kernel_size);
self
}
pub fn stride(&mut self, stride: (usize, usize)) -> &mut Self {
self.stride = Some(stride);
self
}
pub fn build(&self) -> Result<MaxPool2d, ModelError> {
let kernel_size = self.kernel_size.ok_or(ModelError::MissingArgument(
"kernel_size in MaxPool2dBuilder".to_string(),
))?;
if kernel_size.0 == 0 || kernel_size.1 == 0 {
return Err(ModelError::InvalidArgument(
"kernel_size in MaxPool2dBuilder must be greater than 0".to_string(),
));
}
let stride = self.stride.unwrap_or(kernel_size);
if stride.0 == 0 || stride.1 == 0 {
return Err(ModelError::InvalidArgument(
"stride in MaxPool2dBuilder must be greater than 0".to_string(),
));
}
let id = ID.fetch_add(1, Ordering::Relaxed);
Ok(MaxPool2d {
kernel_size,
stride,
id,
})
}
}