maxPool2d function

Tensor maxPool2d(
  1. Tensor x, {
  2. required int kernel,
  3. required int stride,
  4. int padding = 0,
  5. bool ceilMode = false,
})

Max-pool [N, C, H, W] with kernel k, stride, padding (all spatially symmetric). Output shape [N, C, Hout, Wout] where Hout = (H + 2*p - k) / s + 1 (or ceil-divided if ceilMode).

Implementation

Tensor maxPool2d(
  Tensor x, {
  required int kernel,
  required int stride,
  int padding = 0,
  bool ceilMode = false,
}) {
  if (x.shape.length != 4) {
    throw ArgumentError('maxPool2d: expected [N, C, H, W]; got ${x.shape}');
  }
  final n = x.shape[0];
  final c = x.shape[1];
  final h = x.shape[2];
  final w = x.shape[3];
  int outDim(int inSize) {
    final num = inSize + 2 * padding - kernel;
    if (num < 0) return 0;
    if (ceilMode) {
      return (num + stride - 1) ~/ stride + 1;
    }
    return num ~/ stride + 1;
  }

  final hOut = outDim(h);
  final wOut = outDim(w);
  if (hOut <= 0 || wOut <= 0) {
    throw ArgumentError('maxPool2d: non-positive output ${[n, c, hOut, wOut]}');
  }
  final data = x.toFloat32List();
  final out = Float32List(n * c * hOut * wOut);
  for (int ni = 0; ni < n; ni++) {
    for (int ci = 0; ci < c; ci++) {
      final srcBase = (ni * c + ci) * h * w;
      final dstBase = (ni * c + ci) * hOut * wOut;
      for (int oy = 0; oy < hOut; oy++) {
        final iy0 = oy * stride - padding;
        for (int ox = 0; ox < wOut; ox++) {
          final ix0 = ox * stride - padding;
          double best = -double.infinity;
          for (int ky = 0; ky < kernel; ky++) {
            final iy = iy0 + ky;
            if (iy < 0 || iy >= h) continue;
            for (int kx = 0; kx < kernel; kx++) {
              final ix = ix0 + kx;
              if (ix < 0 || ix >= w) continue;
              final v = data[srcBase + iy * w + ix];
              if (v > best) best = v;
            }
          }
          if (best == -double.infinity) best = 0.0;
          out[dstBase + oy * wOut + ox] = best;
        }
      }
    }
  }
  return Tensor.fromFloat32List([n, c, hOut, wOut], out, device: x.device);
}