// maxpool2d.mbt 鈥?2D max-pooling primitive (forward only).
//
// Layout: NCHW row-major flat `Array[Float]`, matching conv2d.mbt.
//
// Conventions:
//   - Input:  [n, c, h,  w]   length = n * c * h * w
//   - Output: [n, c, ho, wo]  length = n * c * ho * wo
//
//   ho = (h + 2*pad - kh) / stride + 1
//   wo = (w + 2*pad - kw) / stride + 1
//
// Out-of-bounds positions are treated as -infinity (so they never win
// the max). Padding is zero-padded input (negative infinity at padding
// positions).

///|
/// MaxPool2d parameter container.
pub struct MaxPool2dParam {
  kh : Int
  kw : Int
  stride : Int
  pad : Int
}

///|
/// Build a MaxPool2dParam. `stride` defaults to `kh` (non-overlapping).
pub fn MaxPool2dParam::new(
  kh : Int,
  kw : Int,
  stride? : Int,
  pad? : Int = 0,
) -> MaxPool2dParam {
  let s = if stride is Some(v) { v } else { kh }
  { kh, kw, stride: s, pad }
}

///|
/// Forward pass for a 2D max pool.
///
/// `input` : length = n * c * h * w
/// returns : length = n * c * ho * wo
pub fn maxpool2d_forward(
  input : Array[Float],
  n : Int,
  c : Int,
  h : Int,
  w : Int,
  param : MaxPool2dParam,
) -> Array[Float] {
  let kh = param.kh
  let kw = param.kw
  let stride = param.stride
  let pad = param.pad
  let ho = (h + 2 * pad - kh) / stride + 1
  let wo = (w + 2 * pad - kw) / stride + 1
  let out : Array[Float] = Array::make(n * c * ho * wo, 0.0F)
  // -inf constant (Float32 lowest finite ~ -3.4e38)
  let neg_inf : Float = -3.4e38F
  let input_n_stride = c * h * w
  let input_c_stride = h * w
  let out_n_stride = c * ho * wo
  let out_c_stride = ho * wo
  for batch in 0..= 0 && ih < h {
              let mut kw_i = 0
              while kw_i < kw {
                let iw = iw_base + kw_i
                if iw >= 0 && iw < w {
                  let v = input[in_c_off + ih * w + iw]
                  if v > best {
                    best = v
                  }
                }
                kw_i = kw_i + 1
              }
            }
            kh_i = kh_i + 1
          }
          out[out_c_off + out_h * wo + out_w] = best
          out_w = out_w + 1
        }
        out_h = out_h + 1
      }
      ch = ch + 1
    }
  }
  out
}