Create a 1D, 2D, or 3D average pooling module.
(dims, *args, **kwargs)
| 266 | |
| 267 | |
| 268 | def avg_pool_nd(dims, *args, **kwargs): |
| 269 | """ |
| 270 | Create a 1D, 2D, or 3D average pooling module. |
| 271 | """ |
| 272 | if dims == 1: |
| 273 | return nn.AvgPool1d(*args, **kwargs) |
| 274 | elif dims == 2: |
| 275 | return nn.AvgPool2d(*args, **kwargs) |
| 276 | elif dims == 3: |
| 277 | return nn.AvgPool3d(*args, **kwargs) |
| 278 | raise ValueError(f"unsupported dimensions: {dims}") |
| 279 | |
| 280 | |
| 281 | class AlphaBlender(nn.Module): |