// maxpool2d_backward.mbt 鈥?MaxPool2d backward pass.
//
// MaxPool forward:  out[n, c, ho, wo] = max over window of input
//                   (out-of-bounds positions treated as -inf).
// MaxPool backward: d_input[n, c, ih, iw] = d_output[n, c, ho, wo]
//                   if (ih, iw) was the argmax of window (ho, wo),
//                   else 0.
//
// Requires `argmax_idx` from `maxpool2d_forward_with_idx`. Layout of
// argmax_idx matches output: length n*c*ho*wo, value = flat input
// index (n, c, ih, iw) -> n*(c*h*w) + c*(h*w) + ih*w + iw, or -1 if
// the window was entirely OOB (impossible given -inf default).

///|
/// MaxPool2d forward + argmax recording.
///
/// `argmax_idx` is filled in with the flat input index of the argmax
/// for each output position. If the window is entirely OOB (impossible
/// when stride > 0 since at least one cell is in-bounds), the index is
/// set to -1.
///
/// returns (output, argmax_idx)
pub fn maxpool2d_forward_with_idx(
  input : Array[Float],
  n : Int,
  c : Int,
  h : Int,
  w : Int,
  param : MaxPool2dParam,
) -> (Array[Float], Array[Int]) {
  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)
  let argmax_idx : Array[Int] = Array::make(n * c * ho * wo, -1)
  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 flat_in_idx = in_c_off + ih * w + iw
                  let v = input[flat_in_idx]
                  if v > best {
                    best = v
                    best_idx = flat_in_idx
                  }
                }
                kw_i = kw_i + 1
              }
            }
            kh_i = kh_i + 1
          }
          out[out_c_off + out_h * wo + out_w] = best
          argmax_idx[out_c_off + out_h * wo + out_w] = best_idx
          out_w = out_w + 1
        }
        out_h = out_h + 1
      }
      ch = ch + 1
    }
  }
  (out, argmax_idx)
}

///|
/// MaxPool2d backward pass.
///
/// `d_output` is the upstream gradient of shape [n, c, ho, wo].
/// `argmax_idx` is the recorded argmax from
/// `maxpool2d_forward_with_idx`.
///
/// returns `d_input` of shape [n, c, h, w].
pub fn maxpool2d_backward(
  d_output : Array[Float],
  argmax_idx : Array[Int],
  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 d_input : Array[Float] = Array::make(n * c * h * w, 0.0F)
  let out_n_stride = c * ho * wo
  let out_c_stride = ho * wo
  for batch in 0..= 0 {
            d_input[arg_in] = d_input[arg_in] + d_output[out_idx]
          }
          out_w = out_w + 1
        }
        out_h = out_h + 1
      }
      ch = ch + 1
    }
  }
  d_input
}