globalAvgPool2d function
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);
}