Create a 1D, 2D, or 3D convolution module.
(dims, *args, **kwargs)
| 139 | return super().forward(x.float()).type(x.dtype) |
| 140 | |
| 141 | def conv_nd(dims, *args, **kwargs): |
| 142 | """ |
| 143 | Create a 1D, 2D, or 3D convolution module. |
| 144 | """ |
| 145 | if dims == 1: |
| 146 | return nn.Conv1d(*args, **kwargs) |
| 147 | elif dims == 2: |
| 148 | return nn.Conv2d(*args, **kwargs) |
| 149 | elif dims == 3: |
| 150 | return nn.Conv3d(*args, **kwargs) |
| 151 | raise ValueError(f"unsupported dimensions: {dims}") |
| 152 | |
| 153 | |
| 154 | def linear(*args, **kwargs): |
nothing calls this directly
no outgoing calls
no test coverage detected