1use 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, ¶ms).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(¶ms);
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, ¶ms)?;
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 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}