Create a 1D, 2D, or 3D convolution module.
(dims, *args, **kwargs)
| 216 | return super().forward(x.float()).type(x.dtype) |
| 217 | |
| 218 | def conv_nd(dims, *args, **kwargs): |
| 219 | """ |
| 220 | Create a 1D, 2D, or 3D convolution module. |
| 221 | """ |
| 222 | if dims == 1: |
| 223 | return nn.Conv1d(*args, **kwargs) |
| 224 | elif dims == 2: |
| 225 | return nn.Conv2d(*args, **kwargs) |
| 226 | elif dims == 3: |
| 227 | return nn.Conv3d(*args, **kwargs) |
| 228 | raise ValueError(f"unsupported dimensions: {dims}") |
| 229 | |
| 230 | |
| 231 | def linear(*args, **kwargs): |