Skip to main content

gwr_models/processing_element/operators/
maxpool.rs

1// Copyright (c) 2026 Graphcore Ltd. All rights reserved.
2
3//! The MaxPool operator
4//!
5//! See <https://onnx.ai/onnx/operators/onnx__MaxPool.html#l-onnx-doc-maxpool>
6
7use std::rc::Rc;
8
9use gwr_engine::sim_error;
10use gwr_engine::types::{SimError, SimResult};
11use rand::RngExt;
12use serde::{Deserialize, Deserializer, Serialize};
13
14use super::{Operator, Shape, Tensor, TensorPartition};
15use crate::processing_element::operators::dtype::DataType;
16use crate::processing_element::operators::{
17    DimPartition, ExpansionDirection, HasShape, TensorView, apply_dim_partitions,
18    partition_across_dimensions,
19};
20use crate::processing_element::{ComputeCapabilities, MachineOp, MachineOpCounts};
21
22const NAME: &str = "MaxPool";
23const BATCH_DIM: usize = 0;
24const CHANNEL_DIM: usize = 1;
25const FIRST_SPATIAL_DIM: usize = 2;
26
27#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize, PartialEq, Eq)]
28pub enum AutoPad {
29    #[default]
30    #[serde(rename = "NOTSET", alias = "notset")]
31    NotSet,
32    #[serde(rename = "SAME_UPPER", alias = "same_upper")]
33    SameUpper,
34    #[serde(rename = "SAME_LOWER", alias = "same_lower")]
35    SameLower,
36    #[serde(rename = "VALID", alias = "valid")]
37    Valid,
38}
39
40fn deserialize_optional_bool_or_int<'de, D>(deserializer: D) -> Result<Option<bool>, D::Error>
41where
42    D: Deserializer<'de>,
43{
44    #[derive(Deserialize)]
45    #[serde(untagged)]
46    enum BoolOrInt {
47        Bool(bool),
48        Int(u8),
49    }
50
51    match Option::<BoolOrInt>::deserialize(deserializer)? {
52        None => Ok(None),
53        Some(BoolOrInt::Bool(value)) => Ok(Some(value)),
54        Some(BoolOrInt::Int(0)) => Ok(Some(false)),
55        Some(BoolOrInt::Int(1)) => Ok(Some(true)),
56        Some(BoolOrInt::Int(value)) => Err(serde::de::Error::custom(format!(
57            "expected 0 or 1, got {value}"
58        ))),
59    }
60}
61
62#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
63pub struct OperatorMaxPool {
64    #[serde(default, skip_serializing_if = "Option::is_none")]
65    pub auto_pad: Option<AutoPad>,
66    #[serde(
67        default,
68        deserialize_with = "deserialize_optional_bool_or_int",
69        skip_serializing_if = "Option::is_none"
70    )]
71    pub ceil_mode: Option<bool>,
72    #[serde(default, skip_serializing_if = "Option::is_none")]
73    pub dilations: Option<Vec<usize>>,
74    pub kernel_shape: Vec<usize>,
75    #[serde(default, skip_serializing_if = "Option::is_none")]
76    pub pads: Option<Vec<usize>>,
77    #[serde(default, skip_serializing_if = "Option::is_none")]
78    pub storage_order: Option<usize>,
79    #[serde(default, skip_serializing_if = "Option::is_none")]
80    pub strides: Option<Vec<usize>>,
81}
82
83impl OperatorMaxPool {
84    #[must_use]
85    pub fn new(kernel_shape: &[usize]) -> Self {
86        Self {
87            auto_pad: None,
88            ceil_mode: None,
89            dilations: None,
90            kernel_shape: kernel_shape.to_vec(),
91            pads: None,
92            storage_order: None,
93            strides: None,
94        }
95    }
96
97    fn auto_pad(&self) -> AutoPad {
98        self.auto_pad.unwrap_or_default()
99    }
100
101    fn ceil_mode(&self) -> bool {
102        self.ceil_mode.unwrap_or(false)
103    }
104
105    fn normalize_params(&self, spatial_rank: usize) -> Result<PoolParams, SimError> {
106        if spatial_rank == 0 {
107            return sim_error!("{NAME}: input must contain at least one spatial dimension");
108        }
109
110        if self.kernel_shape.len() != spatial_rank {
111            return sim_error!(
112                "{NAME}: kernel_shape rank {} does not match spatial rank {spatial_rank}",
113                self.kernel_shape.len()
114            );
115        }
116
117        if self.kernel_shape.contains(&0) {
118            return sim_error!("{NAME}: kernel_shape entries must be greater than zero");
119        }
120
121        let auto_pad = self.auto_pad();
122        let storage_order = self.storage_order.unwrap_or(0);
123        let pads = self.pads.as_deref().unwrap_or_default();
124
125        if storage_order > 1 {
126            return sim_error!("{NAME}: storage_order must be 0 or 1");
127        }
128
129        if auto_pad != AutoPad::NotSet && !pads.is_empty() {
130            return sim_error!("{NAME}: pads cannot be used with auto_pad");
131        }
132
133        let strides = normalize_axis_values(
134            "strides",
135            self.strides.as_deref().unwrap_or_default(),
136            spatial_rank,
137            1,
138        )?;
139        let dilations = normalize_axis_values(
140            "dilations",
141            self.dilations.as_deref().unwrap_or_default(),
142            spatial_rank,
143            1,
144        )?;
145        let pads = if pads.is_empty() {
146            vec![0; spatial_rank * 2]
147        } else if pads.len() == spatial_rank * 2 {
148            pads.to_vec()
149        } else {
150            return sim_error!(
151                "{NAME}: pads rank {} does not match 2 * spatial rank {}",
152                pads.len(),
153                spatial_rank * 2
154            );
155        };
156
157        if strides.contains(&0) {
158            return sim_error!("{NAME}: strides entries must be greater than zero");
159        }
160        if dilations.contains(&0) {
161            return sim_error!("{NAME}: dilations entries must be greater than zero");
162        }
163
164        Ok(PoolParams {
165            kernel_shape: self.kernel_shape.clone(),
166            strides,
167            dilations,
168            pads_begin: pads[..spatial_rank].to_vec(),
169            pads_end: pads[spatial_rank..].to_vec(),
170        })
171    }
172
173    fn output_shape_and_resolved_params<T: HasShape>(
174        &self,
175        input: &T,
176    ) -> Result<(Shape, PoolParams), SimError> {
177        if input.num_dims() < 3 {
178            return sim_error!("{NAME}: input rank must be at least 3");
179        }
180
181        let input_dims = input.shape().get_dims();
182        let spatial_rank = input.num_dims() - FIRST_SPATIAL_DIM;
183        let mut params = self.normalize_params(spatial_rank)?;
184        let auto_pad = self.auto_pad();
185        let ceil_mode = self.ceil_mode();
186
187        let mut output_dims = vec![input_dims[BATCH_DIM], input_dims[CHANNEL_DIM]];
188        let mut resolved_pads_begin = vec![0; spatial_rank];
189        let mut resolved_pads_end = vec![0; spatial_rank];
190
191        for axis in 0..spatial_rank {
192            let input_dim = input_dims[FIRST_SPATIAL_DIM + axis];
193            if input_dim == 0 {
194                return sim_error!("{NAME}: input spatial dimensions must be greater than zero");
195            }
196            let effective_kernel =
197                effective_kernel(params.kernel_shape[axis], params.dilations[axis]);
198            let stride = params.strides[axis];
199
200            let (output_dim, pad_begin, pad_end) = match auto_pad {
201                AutoPad::NotSet => {
202                    let pad_begin = params.pads_begin[axis];
203                    let pad_end = params.pads_end[axis];
204                    let output_dim = explicit_output_dim(
205                        input_dim,
206                        effective_kernel,
207                        stride,
208                        pad_begin,
209                        pad_end,
210                        ceil_mode,
211                    )?;
212                    (output_dim, pad_begin, pad_end)
213                }
214                AutoPad::Valid => {
215                    let output_dim =
216                        valid_output_dim(input_dim, effective_kernel, stride, ceil_mode)?;
217                    (output_dim, 0, 0)
218                }
219                AutoPad::SameUpper | AutoPad::SameLower => {
220                    let output_dim = input_dim.div_ceil(stride);
221                    let pad_shape =
222                        ((output_dim - 1) * stride + effective_kernel).saturating_sub(input_dim);
223                    let smaller_side = pad_shape / 2;
224                    let larger_side = pad_shape - smaller_side;
225                    match auto_pad {
226                        AutoPad::SameUpper => (output_dim, smaller_side, larger_side),
227                        AutoPad::SameLower => (output_dim, larger_side, smaller_side),
228                        AutoPad::NotSet | AutoPad::Valid => unreachable!(),
229                    }
230                }
231            };
232
233            output_dims.push(output_dim);
234            resolved_pads_begin[axis] = pad_begin;
235            resolved_pads_end[axis] = pad_end;
236        }
237
238        params.pads_begin = resolved_pads_begin;
239        params.pads_end = resolved_pads_end;
240
241        Ok((Shape(output_dims), params))
242    }
243
244    fn infer_input_shape<T: HasShape>(&self, output: &T) -> Result<Shape, SimError> {
245        if output.num_dims() < 3 {
246            return sim_error!("{NAME}: output rank must be at least 3");
247        }
248
249        let output_dims = output.shape().get_dims();
250        let spatial_rank = output.num_dims() - FIRST_SPATIAL_DIM;
251        let params = self.normalize_params(spatial_rank)?;
252        let auto_pad = self.auto_pad();
253
254        let mut input_dims = vec![output_dims[BATCH_DIM], output_dims[CHANNEL_DIM]];
255        for axis in 0..spatial_rank {
256            let output_dim = output_dims[FIRST_SPATIAL_DIM + axis];
257            if output_dim == 0 {
258                return sim_error!("{NAME}: output spatial dimensions must be greater than zero");
259            }
260            let effective_kernel =
261                effective_kernel(params.kernel_shape[axis], params.dilations[axis]);
262            let stride = params.strides[axis];
263            let input_dim = match auto_pad {
264                AutoPad::SameUpper | AutoPad::SameLower => (output_dim - 1) * stride + 1,
265                AutoPad::NotSet | AutoPad::Valid => ((output_dim - 1) * stride + effective_kernel)
266                    .saturating_sub(params.pads_begin[axis] + params.pads_end[axis])
267                    .max(1),
268            };
269            input_dims.push(input_dim);
270        }
271
272        Ok(Shape(input_dims))
273    }
274
275    fn can_partition_spatial(&self, params: &PoolParams) -> bool {
276        self.auto_pad() == AutoPad::NotSet
277            && params.pads_begin.iter().all(|pad| *pad == 0)
278            && params.pads_end.iter().all(|pad| *pad == 0)
279    }
280}
281
282impl Default for OperatorMaxPool {
283    fn default() -> Self {
284        Self::new(&[2, 2])
285    }
286}
287
288pub fn create_maxpool_op<T: HasShape>(
289    tensor: &T,
290    direction: ExpansionDirection,
291    expand_ratio: f64,
292) -> Result<OperatorMaxPool, SimError> {
293    if tensor.num_dims() < FIRST_SPATIAL_DIM + 1 {
294        return sim_error!(
295            "{NAME}: generated MaxPool tensor rank must be at least {}",
296            FIRST_SPATIAL_DIM + 1
297        );
298    }
299
300    let dims = tensor.shape().get_dims();
301    let mut kernel_shape = Vec::with_capacity(dims.len() - FIRST_SPATIAL_DIM);
302    let mut pads_begin = Vec::with_capacity(dims.len() - FIRST_SPATIAL_DIM);
303    let mut pads_end = Vec::with_capacity(dims.len() - FIRST_SPATIAL_DIM);
304
305    for dim in &dims[FIRST_SPATIAL_DIM..] {
306        let (kernel, pad_begin, pad_end) = maxpool_axis_params(*dim, expand_ratio, direction);
307        kernel_shape.push(kernel);
308        pads_begin.push(pad_begin);
309        pads_end.push(pad_end);
310    }
311
312    let mut operator = OperatorMaxPool::new(&kernel_shape);
313    if pads_begin.iter().chain(pads_end.iter()).any(|pad| *pad > 0) {
314        let mut pads = pads_begin;
315        pads.extend(pads_end);
316        operator.pads = Some(pads);
317    }
318
319    Ok(operator)
320}
321
322#[derive(Clone, Debug)]
323struct PoolParams {
324    kernel_shape: Vec<usize>,
325    strides: Vec<usize>,
326    dilations: Vec<usize>,
327    pads_begin: Vec<usize>,
328    pads_end: Vec<usize>,
329}
330
331fn normalize_axis_values(
332    name: &str,
333    values: &[usize],
334    spatial_rank: usize,
335    default: usize,
336) -> Result<Vec<usize>, SimError> {
337    if values.is_empty() {
338        Ok(vec![default; spatial_rank])
339    } else if values.len() == spatial_rank {
340        Ok(values.to_vec())
341    } else {
342        sim_error!(
343            "{NAME}: {name} rank {} does not match spatial rank {spatial_rank}",
344            values.len()
345        )
346    }
347}
348
349fn effective_kernel(kernel: usize, dilation: usize) -> usize {
350    dilation * (kernel - 1) + 1
351}
352
353fn scaled_dim(dim: usize, expand_ratio: f64) -> usize {
354    ((dim as f64) * expand_ratio).round().max(1.0) as usize
355}
356
357fn split_padding(total_padding: usize) -> (usize, usize) {
358    let pad_begin = total_padding / 2;
359    let pad_end = total_padding - pad_begin;
360    (pad_begin, pad_end)
361}
362
363fn maxpool_axis_params(
364    dim_size: usize,
365    expand_ratio: f64,
366    direction: ExpansionDirection,
367) -> (usize, usize, usize) {
368    let target_dim_size = scaled_dim(dim_size, expand_ratio);
369    match direction {
370        ExpansionDirection::Forward if target_dim_size <= dim_size => {
371            (dim_size - target_dim_size + 1, 0, 0)
372        }
373        ExpansionDirection::Forward => {
374            let (pad_begin, pad_end) = split_padding(target_dim_size - dim_size);
375            (1, pad_begin, pad_end)
376        }
377        ExpansionDirection::Backward if target_dim_size >= dim_size => {
378            (target_dim_size - dim_size + 1, 0, 0)
379        }
380        ExpansionDirection::Backward => {
381            let (pad_begin, pad_end) = split_padding(dim_size - target_dim_size);
382            (1, pad_begin, pad_end)
383        }
384    }
385}
386
387fn ceil_div_i128(numerator: i128, denominator: i128) -> i128 {
388    if numerator >= 0 {
389        (numerator + denominator - 1) / denominator
390    } else {
391        numerator / denominator
392    }
393}
394
395fn explicit_output_dim(
396    input_dim: usize,
397    effective_kernel: usize,
398    stride: usize,
399    pad_begin: usize,
400    pad_end: usize,
401    ceil_mode: bool,
402) -> Result<usize, SimError> {
403    let numerator =
404        input_dim as i128 + pad_begin as i128 + pad_end as i128 - effective_kernel as i128;
405    let stride = stride as i128;
406    let mut output_dim = if ceil_mode {
407        ceil_div_i128(numerator, stride) + 1
408    } else {
409        numerator.div_euclid(stride) + 1
410    };
411
412    if ceil_mode
413        && output_dim > 0
414        && (output_dim - 1) * stride >= input_dim as i128 + pad_begin as i128
415    {
416        output_dim -= 1;
417    }
418
419    if output_dim <= 0 {
420        return sim_error!("{NAME}: pooling window produces an empty output dimension");
421    }
422    Ok(output_dim as usize)
423}
424
425fn valid_output_dim(
426    input_dim: usize,
427    effective_kernel: usize,
428    stride: usize,
429    ceil_mode: bool,
430) -> Result<usize, SimError> {
431    let output_dim = if ceil_mode {
432        ceil_div_i128(
433            input_dim as i128 - effective_kernel as i128 + 1,
434            stride as i128,
435        )
436    } else {
437        (input_dim as i128 - effective_kernel as i128).div_euclid(stride as i128) + 1
438    };
439
440    if output_dim <= 0 {
441        return sim_error!("{NAME}: pooling window produces an empty output dimension");
442    }
443    Ok(output_dim as usize)
444}
445
446fn choose_partition_dims<T: HasShape>(output: &T, allow_spatial: bool) -> Vec<usize> {
447    output
448        .shape()
449        .get_dims()
450        .iter()
451        .enumerate()
452        .filter_map(|(dim, size)| {
453            (*size > 1 && (allow_spatial || dim < FIRST_SPATIAL_DIM)).then_some(dim)
454        })
455        .collect()
456}
457
458fn validate_inputs<T: HasShape>(inputs: &[Option<T>]) -> Result<&T, SimError> {
459    if inputs.len() != 1 {
460        return sim_error!("{NAME}: {} inputs found - expected 1", inputs.len());
461    }
462    inputs[0]
463        .as_ref()
464        .ok_or(SimError(format!("{NAME}: missing input 0")))
465}
466
467fn validate_outputs<T: HasShape>(outputs: &[Option<T>]) -> Result<(&T, Option<&T>), SimError> {
468    if !(1..=2).contains(&outputs.len()) {
469        return sim_error!("{NAME}: {} outputs found - expected 1 or 2", outputs.len());
470    }
471
472    let output = outputs[0]
473        .as_ref()
474        .ok_or(SimError(format!("{NAME}: missing output 0")))?;
475    let indices = outputs.get(1).and_then(Option::as_ref);
476
477    if let Some(indices) = indices
478        && indices.shape() != output.shape()
479    {
480        return sim_error!(
481            "{NAME}: Indices shape {:?} must match output shape {:?}",
482            indices.shape(),
483            output.shape()
484        );
485    }
486
487    Ok((output, indices))
488}
489
490fn validate_input_outputs<'a, 'b, T: HasShape>(
491    op: &OperatorMaxPool,
492    inputs: &'a [Option<T>],
493    outputs: &'b [Option<T>],
494) -> Result<(&'a T, &'b T, Option<&'b T>), SimError> {
495    let input = validate_inputs(inputs)?;
496    let (output, indices) = validate_outputs(outputs)?;
497
498    let (expected_shape, _) = op.output_shape_and_resolved_params(input)?;
499    if expected_shape != *output.shape() {
500        return sim_error!(
501            "{NAME}: Invalid output shape - expected {:?}, found {:?}",
502            expected_shape,
503            output.shape()
504        );
505    }
506
507    Ok((input, output, indices))
508}
509
510fn validate_tensor_dtypes(input: &Tensor, output: &Tensor, indices: Option<&Tensor>) -> SimResult {
511    if input.dtype() != output.dtype() {
512        return sim_error!(
513            "{NAME}: output dtype {:?} must match input dtype {:?}",
514            output.dtype(),
515            input.dtype()
516        );
517    }
518
519    if let Some(indices) = indices
520        && *indices.dtype() != DataType::Int64
521    {
522        return sim_error!(
523            "{NAME}: Indices dtype {:?} must be {:?}",
524            indices.dtype(),
525            DataType::Int64
526        );
527    }
528
529    Ok(())
530}
531
532fn should_add_indices_output(rng: &mut impl RngExt, expand_ratio: f64) -> bool {
533    if !expand_ratio.is_finite() || expand_ratio <= 0.0 {
534        false
535    } else if expand_ratio >= 1.0 {
536        true
537    } else {
538        rng.random_bool(expand_ratio)
539    }
540}
541
542pub fn maybe_add_indices_output(
543    outputs: &mut Vec<Option<Tensor>>,
544    expand_ratio: f64,
545    rng: &mut impl RngExt,
546) -> Result<bool, SimError> {
547    if outputs.len() >= 2 || !should_add_indices_output(rng, expand_ratio) {
548        return Ok(false);
549    }
550
551    let output = outputs
552        .first()
553        .and_then(Option::as_ref)
554        .ok_or_else(|| SimError(format!("{NAME}: missing output 0")))?;
555    outputs.push(Some(Tensor {
556        id: None,
557        shape: output.shape().clone(),
558        dtype: DataType::Int64,
559        addr: 0,
560    }));
561    Ok(true)
562}
563
564fn window_valid_element_count(
565    output_coordinate: &[usize],
566    input_spatial_dims: &[usize],
567    params: &PoolParams,
568) -> usize {
569    output_coordinate
570        .iter()
571        .enumerate()
572        .map(|(axis, output_idx)| {
573            let window_start = *output_idx as i128 * params.strides[axis] as i128
574                - params.pads_begin[axis] as i128;
575
576            (0..params.kernel_shape[axis])
577                .filter(|kernel_idx| {
578                    let input_idx = window_start + (*kernel_idx * params.dilations[axis]) as i128;
579                    input_idx >= 0 && input_idx < input_spatial_dims[axis] as i128
580                })
581                .count()
582        })
583        .product()
584}
585
586fn unravel_index(mut linear_idx: usize, dims: &[usize]) -> Vec<usize> {
587    let mut coordinate = vec![0; dims.len()];
588    for (axis, dim) in dims.iter().enumerate().rev() {
589        coordinate[axis] = linear_idx % dim;
590        linear_idx /= dim;
591    }
592    coordinate
593}
594
595fn maxpool_comparisons<T: HasShape>(
596    op: &OperatorMaxPool,
597    inputs: &[Option<T>],
598    outputs: &[Option<T>],
599) -> Result<usize, SimError> {
600    let (input, output, _) = validate_input_outputs(op, inputs, outputs)?;
601    let (_, params) = op.output_shape_and_resolved_params(input)?;
602
603    let input_spatial_dims = &input.shape().get_dims()[FIRST_SPATIAL_DIM..];
604    let output_spatial_dims = &output.shape().get_dims()[FIRST_SPATIAL_DIM..];
605    let num_spatial_outputs = output_spatial_dims.iter().product::<usize>();
606
607    let comparisons_per_batch_channel = (0..num_spatial_outputs)
608        .map(|linear_idx| {
609            let coordinate = unravel_index(linear_idx, output_spatial_dims);
610            window_valid_element_count(&coordinate, input_spatial_dims, &params).saturating_sub(1)
611        })
612        .sum::<usize>();
613
614    Ok(output.get_dim(output.num_dims(), BATCH_DIM)
615        * output.get_dim(output.num_dims(), CHANNEL_DIM)
616        * comparisons_per_batch_channel)
617}
618
619fn input_partition_for_output_partition(
620    input_view: &TensorView,
621    partitions: &[DimPartition],
622    params: &PoolParams,
623) -> Result<TensorView, SimError> {
624    let mut input_shape = input_view.shape().get_dims().clone();
625    let mut input_offsets = input_view.offsets().get_dims().clone();
626
627    for partition in partitions {
628        if input_shape[partition.dim] <= 1 {
629            continue;
630        }
631
632        if partition.dim < FIRST_SPATIAL_DIM {
633            input_offsets[partition.dim] += partition.offset;
634            input_shape[partition.dim] = partition.len;
635            continue;
636        }
637
638        let axis = partition.dim - FIRST_SPATIAL_DIM;
639        let effective_kernel = effective_kernel(params.kernel_shape[axis], params.dilations[axis]);
640        let first_output = partition.offset;
641        let last_output = partition.offset + partition.len - 1;
642        let raw_start =
643            first_output as i128 * params.strides[axis] as i128 - params.pads_begin[axis] as i128;
644        let raw_end = last_output as i128 * params.strides[axis] as i128
645            - params.pads_begin[axis] as i128
646            + effective_kernel as i128;
647
648        let start = raw_start.clamp(0, input_shape[partition.dim] as i128) as usize;
649        let end = raw_end.clamp(0, input_shape[partition.dim] as i128) as usize;
650        if start >= end {
651            return sim_error!("{NAME}: partition produced an empty input view");
652        }
653
654        input_offsets[partition.dim] += start;
655        input_shape[partition.dim] = end - start;
656    }
657
658    Ok(TensorView::new(
659        input_view.tensor().clone(),
660        &input_shape,
661        &input_offsets,
662    ))
663}
664
665impl OperatorMaxPool {
666    pub fn create_outputs(
667        &self,
668        inputs: &[Option<Tensor>],
669        expand_ratio: f64,
670        rng: &mut impl RngExt,
671    ) -> Result<Vec<Option<Tensor>>, SimError> {
672        let input = validate_inputs(inputs)?;
673        let (output_shape, _) = self.output_shape_and_resolved_params(input)?;
674
675        let mut outputs = vec![Some(Tensor {
676            id: None,
677            shape: output_shape,
678            dtype: input.dtype,
679            addr: 0,
680        })];
681        maybe_add_indices_output(&mut outputs, expand_ratio, rng)?;
682        Ok(outputs)
683    }
684
685    pub fn create_inputs(
686        &self,
687        outputs: &[Option<Tensor>],
688        _expand_ratio: f64,
689        _rng: &mut impl RngExt,
690    ) -> Result<Vec<Option<Tensor>>, SimError> {
691        let (output, indices) = validate_outputs(outputs)?;
692        if let Some(indices) = indices
693            && *indices.dtype() != DataType::Int64
694        {
695            return sim_error!(
696                "{NAME}: Indices dtype {:?} must be {:?}",
697                indices.dtype(),
698                DataType::Int64
699            );
700        }
701
702        let input_shape = self.infer_input_shape(output)?;
703        Ok(vec![Some(Tensor {
704            id: None,
705            shape: input_shape,
706            dtype: output.dtype,
707            addr: 0,
708        })])
709    }
710}
711
712impl Operator for OperatorMaxPool {
713    fn validate_tensors(&self, inputs: &[Option<Tensor>], outputs: &[Option<Tensor>]) -> SimResult {
714        let (input, output, indices) = validate_input_outputs(self, inputs, outputs)?;
715        validate_tensor_dtypes(input, output, indices)
716    }
717
718    fn compute_delay_ticks(
719        &self,
720        compute_capabilities: &Rc<ComputeCapabilities>,
721        inputs: &[Option<TensorView>],
722        outputs: &[Option<TensorView>],
723    ) -> Result<usize, SimError> {
724        let comparisons = maxpool_comparisons(self, inputs, outputs)?;
725        compute_capabilities.cycles_for_ops(comparisons, MachineOp::Compare)
726    }
727
728    fn compute_machine_ops(
729        &self,
730        inputs: &[Option<TensorView>],
731        outputs: &[Option<TensorView>],
732    ) -> Result<MachineOpCounts, SimError> {
733        Ok(MachineOpCounts {
734            compares: maxpool_comparisons(self, inputs, outputs)?,
735            ..MachineOpCounts::default()
736        })
737    }
738
739    fn partition_views(
740        &self,
741        input_views: &[Option<TensorView>],
742        output_views: &[Option<TensorView>],
743        num_partitions: usize,
744    ) -> Result<Vec<TensorPartition>, SimError> {
745        let (input_view, output_view, _) = validate_input_outputs(self, input_views, output_views)?;
746        let (_, params) = self.output_shape_and_resolved_params(input_view)?;
747        let allow_spatial = self.can_partition_spatial(&params);
748        let partition_dims = choose_partition_dims(output_view, allow_spatial);
749        let output_view_dims = output_view.shape().get_dims();
750        let partition_specs =
751            partition_across_dimensions(output_view_dims, &partition_dims, num_partitions);
752
753        let mut partitions = Vec::with_capacity(partition_specs.len());
754        for spec in partition_specs {
755            let (output_shape, partition_offsets) = apply_dim_partitions(output_view_dims, &spec);
756            let output_offsets = output_view
757                .offsets()
758                .get_dims()
759                .iter()
760                .zip(partition_offsets.iter())
761                .map(|(base, offset)| base + offset)
762                .collect::<Vec<_>>();
763
764            let input_view = input_partition_for_output_partition(input_view, &spec, &params)?;
765            let outputs = output_views
766                .iter()
767                .map(|maybe_output| {
768                    maybe_output.as_ref().map(|view| {
769                        TensorView::new(view.tensor().clone(), &output_shape, &output_offsets)
770                    })
771                })
772                .collect::<Vec<_>>();
773
774            partitions.push(TensorPartition {
775                inputs: vec![Some(input_view)],
776                outputs,
777            });
778        }
779
780        Ok(partitions)
781    }
782}
783
784#[cfg(test)]
785mod tests {
786    use super::*;
787    use crate::processing_element::operators::dtype::DataType;
788    use crate::processing_element::operators::{Operator, Tensor, partition_tensors};
789
790    fn tensor(shape: &[usize]) -> Option<Tensor> {
791        Some(Tensor::new(shape, &DataType::Bf16, 0))
792    }
793
794    fn indices(shape: &[usize]) -> Option<Tensor> {
795        Some(Tensor::new(shape, &DataType::Int64, 0))
796    }
797
798    fn tensor_view(shape: &[usize]) -> Option<TensorView> {
799        let tensor = Tensor::new(shape, &DataType::Bf16, 0);
800        Some(TensorView::new_full(tensor))
801    }
802
803    fn indices_view(shape: &[usize]) -> Option<TensorView> {
804        let tensor = Tensor::new(shape, &DataType::Int64, 0);
805        Some(TensorView::new_full(tensor))
806    }
807
808    #[test]
809    fn new_leaves_optional_attributes_unspecified_and_uses_onnx_defaults() {
810        let op = OperatorMaxPool::new(&[2, 2]);
811        assert_eq!(op.auto_pad, None);
812        assert_eq!(op.ceil_mode, None);
813        assert_eq!(op.dilations, None);
814        assert_eq!(op.pads, None);
815        assert_eq!(op.storage_order, None);
816        assert_eq!(op.strides, None);
817
818        let mut rng = rand::rng();
819        let outputs = op
820            .create_outputs(&[tensor(&[1, 1, 4, 4])], 1.0, &mut rng)
821            .unwrap();
822
823        assert_eq!(
824            outputs[0].as_ref().unwrap().shape(),
825            &Shape::new(&[1, 1, 3, 3])
826        );
827    }
828
829    #[test]
830    fn generated_forward_op_shrinks_with_larger_kernel_when_expand_ratio_is_below_one() {
831        let input = Tensor::new(&[1, 1, 10, 8], &DataType::Fp32, 0);
832        let op = create_maxpool_op(&input, ExpansionDirection::Forward, 0.5).unwrap();
833
834        assert_eq!(op.kernel_shape, vec![6, 5]);
835        assert_eq!(op.pads, None);
836
837        let mut rng = rand::rng();
838        let outputs = op.create_outputs(&[Some(input)], 1.0, &mut rng).unwrap();
839        assert_eq!(
840            outputs[0].as_ref().unwrap().shape(),
841            &Shape::new(&[1, 1, 5, 4])
842        );
843    }
844
845    #[test]
846    fn generated_forward_op_grows_with_padding_when_expand_ratio_is_above_one() {
847        let input = Tensor::new(&[1, 1, 4, 4], &DataType::Fp32, 0);
848        let op = create_maxpool_op(&input, ExpansionDirection::Forward, 1.5).unwrap();
849
850        assert_eq!(op.kernel_shape, vec![1, 1]);
851        assert_eq!(op.pads, Some(vec![1, 1, 1, 1]));
852
853        let mut rng = rand::rng();
854        let outputs = op.create_outputs(&[Some(input)], 1.0, &mut rng).unwrap();
855        assert_eq!(
856            outputs[0].as_ref().unwrap().shape(),
857            &Shape::new(&[1, 1, 6, 6])
858        );
859    }
860
861    #[test]
862    fn generated_backward_op_uses_expand_ratio_for_input_shape() {
863        let output = Tensor::new(&[1, 1, 4, 4], &DataType::Fp32, 0);
864        let op = create_maxpool_op(&output, ExpansionDirection::Backward, 0.5).unwrap();
865
866        assert_eq!(op.kernel_shape, vec![1, 1]);
867        assert_eq!(op.pads, Some(vec![1, 1, 1, 1]));
868
869        let mut rng = rand::rng();
870        let inputs = op.create_inputs(&[Some(output)], 1.0, &mut rng).unwrap();
871        assert_eq!(
872            inputs[0].as_ref().unwrap().shape(),
873            &Shape::new(&[1, 1, 2, 2])
874        );
875    }
876
877    type OffsetsShapes = (&'static [usize], &'static [usize]);
878    type PartitionOffsetsShapes = (OffsetsShapes, OffsetsShapes, OffsetsShapes);
879
880    fn check_partitions(partitions: &[TensorPartition], expected: &[PartitionOffsetsShapes]) {
881        assert_eq!(partitions.len(), expected.len());
882
883        for (partition, expected_partition) in partitions.iter().zip(expected.iter()) {
884            assert_eq!(partition.inputs.len(), 1);
885            assert_eq!(partition.outputs.len(), 2);
886
887            let input = partition.inputs[0].as_ref().unwrap();
888            let output = partition.outputs[0].as_ref().unwrap();
889            let indices = partition.outputs[1].as_ref().unwrap();
890
891            assert_eq!(
892                input.offsets().get_dims().as_slice(),
893                expected_partition.0.0
894            );
895            assert_eq!(input.shape().get_dims().as_slice(), expected_partition.0.1);
896            assert_eq!(
897                output.offsets().get_dims().as_slice(),
898                expected_partition.1.0
899            );
900            assert_eq!(output.shape().get_dims().as_slice(), expected_partition.1.1);
901            assert_eq!(
902                indices.offsets().get_dims().as_slice(),
903                expected_partition.2.0
904            );
905            assert_eq!(
906                indices.shape().get_dims().as_slice(),
907                expected_partition.2.1
908            );
909        }
910    }
911
912    #[test]
913    fn create_outputs_returns_y_and_indices() {
914        let op = OperatorMaxPool {
915            strides: Some(vec![2, 2]),
916            ..OperatorMaxPool::new(&[2, 2])
917        };
918        let mut rng = rand::rng();
919
920        let outputs = op
921            .create_outputs(&[tensor(&[1, 3, 4, 4])], 1.0, &mut rng)
922            .unwrap();
923
924        assert_eq!(outputs.len(), 2);
925        assert_eq!(
926            outputs[0].as_ref().unwrap().shape(),
927            &Shape::new(&[1, 3, 2, 2])
928        );
929        assert_eq!(outputs[0].as_ref().unwrap().dtype(), &DataType::Bf16);
930        assert_eq!(
931            outputs[1].as_ref().unwrap().shape(),
932            &Shape::new(&[1, 3, 2, 2])
933        );
934        assert_eq!(outputs[1].as_ref().unwrap().dtype(), &DataType::Int64);
935    }
936
937    #[test]
938    fn create_outputs_and_inputs_support_5d_pooling_over_inner_three_dimensions() {
939        let op = OperatorMaxPool {
940            strides: Some(vec![2, 2, 2]),
941            ..OperatorMaxPool::new(&[2, 3, 4])
942        };
943        let mut rng = rand::rng();
944
945        let outputs = op
946            .create_outputs(&[tensor(&[2, 5, 6, 7, 8])], 1.0, &mut rng)
947            .unwrap();
948
949        assert_eq!(outputs.len(), 2);
950        assert_eq!(
951            outputs[0].as_ref().unwrap().shape(),
952            &Shape::new(&[2, 5, 3, 3, 3])
953        );
954        assert_eq!(outputs[0].as_ref().unwrap().dtype(), &DataType::Bf16);
955        assert_eq!(
956            outputs[1].as_ref().unwrap().shape(),
957            &Shape::new(&[2, 5, 3, 3, 3])
958        );
959        assert_eq!(outputs[1].as_ref().unwrap().dtype(), &DataType::Int64);
960        op.validate_tensors(&[tensor(&[2, 5, 6, 7, 8])], &outputs)
961            .unwrap();
962
963        let inputs = op.create_inputs(&outputs, 1.0, &mut rng).unwrap();
964        assert_eq!(
965            inputs[0].as_ref().unwrap().shape(),
966            &Shape::new(&[2, 5, 6, 7, 8])
967        );
968    }
969
970    #[test]
971    fn create_outputs_with_expand_ratio_zero_omits_indices() {
972        let op = OperatorMaxPool {
973            strides: Some(vec![2, 2]),
974            ..OperatorMaxPool::new(&[2, 2])
975        };
976        let mut rng = rand::rng();
977
978        let outputs = op
979            .create_outputs(&[tensor(&[1, 3, 4, 4])], 0.0, &mut rng)
980            .unwrap();
981
982        assert_eq!(outputs.len(), 1);
983        assert_eq!(
984            outputs[0].as_ref().unwrap().shape(),
985            &Shape::new(&[1, 3, 2, 2])
986        );
987        assert_eq!(outputs[0].as_ref().unwrap().dtype(), &DataType::Bf16);
988    }
989
990    #[test]
991    fn validate_accepts_optional_indices_output() {
992        let op = OperatorMaxPool {
993            strides: Some(vec![2, 2]),
994            ..OperatorMaxPool::new(&[2, 2])
995        };
996
997        op.validate_tensors(&[tensor(&[1, 3, 4, 4])], &[tensor(&[1, 3, 2, 2])])
998            .unwrap();
999        op.validate_tensors(
1000            &[tensor(&[1, 3, 4, 4])],
1001            &[tensor(&[1, 3, 2, 2]), indices(&[1, 3, 2, 2])],
1002        )
1003        .unwrap();
1004    }
1005
1006    #[test]
1007    fn validate_rejects_wrong_indices_shape_or_dtype() {
1008        let op = OperatorMaxPool {
1009            strides: Some(vec![2, 2]),
1010            ..OperatorMaxPool::new(&[2, 2])
1011        };
1012
1013        let err = op
1014            .validate_tensors(
1015                &[tensor(&[1, 3, 4, 4])],
1016                &[tensor(&[1, 3, 2, 2]), indices(&[1, 3, 2, 1])],
1017            )
1018            .unwrap_err();
1019        assert!(format!("{err}").contains("Indices shape"));
1020
1021        let err = op
1022            .validate_tensors(
1023                &[tensor(&[1, 3, 4, 4])],
1024                &[tensor(&[1, 3, 2, 2]), tensor(&[1, 3, 2, 2])],
1025            )
1026            .unwrap_err();
1027        assert!(format!("{err}").contains("Indices dtype"));
1028    }
1029
1030    #[test]
1031    fn output_shape_supports_padding_dilation_and_ceil_mode() {
1032        let op = OperatorMaxPool {
1033            ceil_mode: Some(true),
1034            dilations: Some(vec![2, 2]),
1035            pads: Some(vec![1, 1, 1, 1]),
1036            strides: Some(vec![2, 2]),
1037            ..OperatorMaxPool::new(&[3, 3])
1038        };
1039        let mut rng = rand::rng();
1040
1041        let outputs = op
1042            .create_outputs(&[tensor(&[1, 1, 7, 7])], 1.0, &mut rng)
1043            .unwrap();
1044
1045        assert_eq!(
1046            outputs[0].as_ref().unwrap().shape(),
1047            &Shape::new(&[1, 1, 3, 3])
1048        );
1049    }
1050
1051    #[test]
1052    fn output_shape_supports_same_upper_auto_pad() {
1053        let op = OperatorMaxPool {
1054            auto_pad: Some(AutoPad::SameUpper),
1055            strides: Some(vec![2, 2]),
1056            ..OperatorMaxPool::new(&[3, 3])
1057        };
1058        let mut rng = rand::rng();
1059
1060        let outputs = op
1061            .create_outputs(&[tensor(&[1, 1, 5, 6])], 1.0, &mut rng)
1062            .unwrap();
1063
1064        assert_eq!(
1065            outputs[0].as_ref().unwrap().shape(),
1066            &Shape::new(&[1, 1, 3, 3])
1067        );
1068    }
1069
1070    #[test]
1071    fn delay_counts_comparisons_excluding_padding() {
1072        let op = OperatorMaxPool {
1073            pads: Some(vec![1, 1, 1, 1]),
1074            ..OperatorMaxPool::new(&[3, 3])
1075        };
1076        let compute_capabilities = Rc::new(ComputeCapabilities {
1077            adds_per_tick: 200.0,
1078            muls_per_tick: 100.0,
1079            compares_per_tick: 0.5,
1080            sram_bytes: 1024,
1081        });
1082
1083        let delay = op
1084            .compute_delay_ticks(
1085                &compute_capabilities,
1086                &[tensor_view(&[1, 1, 2, 2])],
1087                &[tensor_view(&[1, 1, 2, 2]), indices_view(&[1, 1, 2, 2])],
1088            )
1089            .unwrap();
1090
1091        // Every output window sees the same four real input elements, requiring
1092        // three comparisons per output.
1093        assert_eq!(delay, 24);
1094    }
1095
1096    #[test]
1097    fn partitions_include_both_outputs() {
1098        let op = OperatorMaxPool {
1099            strides: Some(vec![2, 2]),
1100            ..OperatorMaxPool::new(&[2, 2])
1101        };
1102        let inputs = vec![tensor(&[2, 1, 4, 4])];
1103        let outputs = vec![tensor(&[2, 1, 2, 2]), indices(&[2, 1, 2, 2])];
1104
1105        let partitions = partition_tensors(&op, &inputs, &outputs, 2).unwrap();
1106
1107        let expected: &[PartitionOffsetsShapes] = &[
1108            (
1109                (&[0, 0, 0, 0], &[1, 1, 4, 4]),
1110                (&[0, 0, 0, 0], &[1, 1, 2, 2]),
1111                (&[0, 0, 0, 0], &[1, 1, 2, 2]),
1112            ),
1113            (
1114                (&[1, 0, 0, 0], &[1, 1, 4, 4]),
1115                (&[1, 0, 0, 0], &[1, 1, 2, 2]),
1116                (&[1, 0, 0, 0], &[1, 1, 2, 2]),
1117            ),
1118        ];
1119        check_partitions(&partitions, expected);
1120    }
1121}