Create a 1D, 2D, or 3D average pooling module.
(dims, *args, **kwargs)
| 239 | |
| 240 | |
| 241 | def avg_pool_nd(dims, *args, **kwargs): |
| 242 | """ |
| 243 | Create a 1D, 2D, or 3D average pooling module. |
| 244 | """ |
| 245 | if dims == 1: |
| 246 | return nn.AvgPool1d(*args, **kwargs) |
| 247 | elif dims == 2: |
| 248 | return nn.AvgPool2d(*args, **kwargs) |
| 249 | elif dims == 3: |
| 250 | return nn.AvgPool3d(*args, **kwargs) |
| 251 | raise ValueError(f"unsupported dimensions: {dims}") |
| 252 | |
| 253 | |
| 254 | class HybridConditioner(nn.Module): |