| 159 | ## ----------------------------------------------------------------------------- |
| 160 | |
| 161 | class UNetLarge(nn.Module): |
| 162 | def __init__(self, in_channels, out_channels, xl=False): |
| 163 | super(UNetLarge, self).__init__() |
| 164 | |
| 165 | # Number of channels per layer |
| 166 | ic = in_channels |
| 167 | if xl: |
| 168 | ec1 = 96 |
| 169 | ec2 = 128 |
| 170 | ec3 = 192 |
| 171 | ec4 = 256 |
| 172 | ec5 = 384 |
| 173 | dc4 = 256 |
| 174 | dc3 = 192 |
| 175 | dc2 = 128 |
| 176 | dc1 = 96 |
| 177 | else: |
| 178 | ec1 = 64 |
| 179 | ec2 = 96 |
| 180 | ec3 = 128 |
| 181 | ec4 = 192 |
| 182 | ec5 = 256 |
| 183 | dc4 = 192 |
| 184 | dc3 = 128 |
| 185 | dc2 = 96 |
| 186 | dc1 = 64 |
| 187 | oc = out_channels |
| 188 | |
| 189 | # Convolutions |
| 190 | self.enc_conv1a = Conv(ic, ec1) |
| 191 | self.enc_conv1b = Conv(ec1, ec1) |
| 192 | self.enc_conv2a = Conv(ec1, ec2) |
| 193 | self.enc_conv2b = Conv(ec2, ec2) |
| 194 | self.enc_conv3a = Conv(ec2, ec3) |
| 195 | self.enc_conv3b = Conv(ec3, ec3) |
| 196 | self.enc_conv4a = Conv(ec3, ec4) |
| 197 | self.enc_conv4b = Conv(ec4, ec4) |
| 198 | self.enc_conv5a = Conv(ec4, ec5) |
| 199 | self.enc_conv5b = Conv(ec5, ec5) |
| 200 | self.dec_conv4a = Conv(ec5+ec3, dc4) |
| 201 | self.dec_conv4b = Conv(dc4, dc4) |
| 202 | self.dec_conv3a = Conv(dc4+ec2, dc3) |
| 203 | self.dec_conv3b = Conv(dc3, dc3) |
| 204 | self.dec_conv2a = Conv(dc3+ec1, dc2) |
| 205 | self.dec_conv2b = Conv(dc2, dc2) |
| 206 | self.dec_conv1a = Conv(dc2+ic, dc1) |
| 207 | self.dec_conv1b = Conv(dc1, dc1) |
| 208 | self.dec_conv1c = Conv(dc1, oc) |
| 209 | |
| 210 | # Images must be padded to multiples of the alignment |
| 211 | self.alignment = 16 |
| 212 | |
| 213 | def forward(self, input): |
| 214 | # Encoder |
| 215 | # ------------------------------------------- |
| 216 | |
| 217 | x = relu(self.enc_conv1a(input)) # enc_conv1a |
| 218 | x = relu(self.enc_conv1b(x)) # enc_conv1b |