globalAvgPool2d function

Tensor globalAvgPool2d(
  1. Tensor x
)

Global average pool over the spatial axes: [N, C, H, W] -> [N, C].

Implementation

Tensor globalAvgPool2d(Tensor x) {
  if (x.shape.length != 4) {
    throw ArgumentError(
      'globalAvgPool2d: 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];
  final spatial = h * w;
  final data = x.toFloat32List();
  final out = Float32List(n * c);
  for (int ni = 0; ni < n; ni++) {
    for (int ci = 0; ci < c; ci++) {
      final base = (ni * c + ci) * spatial;
      double s = 0.0;
      for (int k = 0; k < spatial; k++) {
        s += data[base + k];
      }
      out[ni * c + ci] = s / spatial;
    }
  }
  return Tensor.fromFloat32List([n, c], out, device: x.device);
}