// avgpool2d.mbt — Average pooling 2D (v0.20.1).
//
// Two pooling modes are supported:
//   - `avgpool2d_forward` / `avgpool2d_backward` — general average
//     pooling with arbitrary kernel/stride/pad.
//   - `global_avg_pool2d_forward` / `global_avg_pool2d_backward` —
//     pools the entire spatial dimension to [n, c, 1, 1] (ResNet's
//     `nn.AdaptiveAvgPool2d((1, 1))` equivalent).
//
// Conventions:
//   - NCHW row-major flat `Array[Float]`.
//   - For padded cells, the divisor is `kh * kw` (NOT adjusted for the
//     actual valid window). This matches PyTorch's
//     `count_include_pad=True` default.
//   - Backward distributes the upstream gradient uniformly across all
//     positions in the kernel window (each position gets `d / (kh*kw)`).

///|
/// AvgPool2d parameter container.
pub struct AvgPool2dParam {
  kh : Int
  kw : Int
  stride : Int
  pad : Int
}

///|
/// Build an AvgPool2dParam. `stride` defaults to `kh`, `pad` to 0.
pub fn AvgPool2dParam::new(
  kh : Int,
  kw : Int,
  stride? : Int = 0,
  pad? : Int = 0,
) -> AvgPool2dParam {
  let s = if stride == 0 { kh } else { stride }
  { kh, kw, stride: s, pad }
}

///|
/// Forward pass.
pub fn avgpool2d_forward(
  input : Array[Float],
  n : Int,
  c : Int,
  h : Int,
  w : Int,
  param : AvgPool2dParam,
) -> 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 pool_size = Float::from_int(kh * kw)
  let out : Array[Float] = Array::make(n * c * ho * wo, 0.0F)
  for ni in 0..= 0 && ih < h {
              for kwi in 0..= 0 && iw < w {
                  let idx = ni * c * h * w + ci * h * w + ih * w + iw
                  sum = sum + input[idx]
                }
              }
            }
          }
          let out_idx = ni * c * ho * wo + ci * ho * wo + ohi * wo + owi
          out[out_idx] = sum / pool_size
        }
      }
    }
  }
  out
}

///|
/// Backward pass. Distributes the upstream gradient uniformly across
/// all positions in each kernel window.
pub fn avgpool2d_backward(
  d_output : Array[Float],
  n : Int,
  c : Int,
  h : Int,
  w : Int,
  param : AvgPool2dParam,
) -> 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 pool_size = 1.0F / Float::from_int(kh * kw)
  let d_input : Array[Float] = Array::make(n * c * h * w, 0.0F)
  for ni in 0..= 0 && ih < h {
              for kwi in 0..= 0 && iw < w {
                  let idx = ni * c * h * w + ci * h * w + ih * w + iw
                  d_input[idx] = d_input[idx] + grad
                }
              }
            }
          }
        }
      }
    }
  }
  d_input
}

// ---------------------------------------------------------------------------
// Global Average Pooling (ResNet's terminal pooling to [n, c, 1, 1]).
// ---------------------------------------------------------------------------

///|
/// Global Average Pooling. Pools the entire spatial dimension to a
/// single value per (n, c). Result shape: `[n, c, 1, 1]`.
pub fn global_avg_pool2d_forward(
  input : Array[Float],
  n : Int,
  c : Int,
  h : Int,
  w : Int,
) -> Array[Float] {
  let pool_size = 1.0F / Float::from_int(h * w)
  let out : Array[Float] = Array::make(n * c, 0.0F)
  for ni in 0.. Array[Float] {
  let pool_size = 1.0F / Float::from_int(h * w)
  let d_input : Array[Float] = Array::make(n * c * h * w, 0.0F)
  for ni in 0..