torch_em.model.unet
1from typing import List, Optional, Union 2 3import numpy as np 4import torch 5import torch.nn as nn 6 7 8# 9# Model Internal Post-processing 10# 11# Note: these are mainly for bioimage.io models, where postprocessing has to be done 12# inside of the model unless its defined in the general spec 13 14 15class AccumulateChannels(nn.Module): 16 """@private 17 """ 18 def __init__( 19 self, 20 invariant_channels, 21 accumulate_channels, 22 accumulator 23 ): 24 super().__init__() 25 self.invariant_channels = invariant_channels 26 self.accumulate_channels = accumulate_channels 27 assert accumulator in ("mean", "min", "max") 28 self.accumulator = getattr(torch, accumulator) 29 30 def _accumulate(self, x, c0, c1): 31 res = self.accumulator(x[:, c0:c1], dim=1, keepdim=True) 32 if not torch.is_tensor(res): 33 res = res.values 34 assert torch.is_tensor(res) 35 return res 36 37 def forward(self, x): 38 if self.invariant_channels is None: 39 c0, c1 = self.accumulate_channels 40 return self._accumulate(x, c0, c1) 41 else: 42 i0, i1 = self.invariant_channels 43 c0, c1 = self.accumulate_channels 44 return torch.cat([x[:, i0:i1], self._accumulate(x, c0, c1)], dim=1) 45 46 47def affinities_to_boundaries(aff_channels, accumulator="max"): 48 """@private 49 """ 50 return AccumulateChannels(None, aff_channels, accumulator) 51 52 53def affinities_with_foreground_to_boundaries(aff_channels, fg_channel=(0, 1), accumulator="max"): 54 """@private 55 """ 56 return AccumulateChannels(fg_channel, aff_channels, accumulator) 57 58 59def affinities_to_boundaries2d(): 60 """@private 61 """ 62 return affinities_to_boundaries((0, 2)) 63 64 65def affinities_with_foreground_to_boundaries2d(): 66 """@private 67 """ 68 return affinities_with_foreground_to_boundaries((1, 3)) 69 70 71def affinities_to_boundaries3d(): 72 """@private 73 """ 74 return affinities_to_boundaries((0, 3)) 75 76 77def affinities_with_foreground_to_boundaries3d(): 78 """@private 79 """ 80 return affinities_with_foreground_to_boundaries((1, 4)) 81 82 83def affinities_to_boundaries_anisotropic(): 84 """@private 85 """ 86 return AccumulateChannels(None, (1, 3), "max") 87 88 89POSTPROCESSING = { 90 "affinities_to_boundaries_anisotropic": affinities_to_boundaries_anisotropic, 91 "affinities_to_boundaries2d": affinities_to_boundaries2d, 92 "affinities_with_foreground_to_boundaries2d": affinities_with_foreground_to_boundaries2d, 93 "affinities_to_boundaries3d": affinities_to_boundaries3d, 94 "affinities_with_foreground_to_boundaries3d": affinities_with_foreground_to_boundaries3d, 95} 96"""@private 97""" 98 99 100# 101# Base Implementations 102# 103 104class UNetBase(nn.Module): 105 """Base class for implementing a U-Net. 106 107 Args: 108 encoder: The encoder of the U-Net. 109 base: The base layer of the U-Net. 110 decoder: The decoder of the U-Net. 111 out_conv: The output convolution applied after the last decoder layer. 112 final_activation: The activation applied after the output convolution or last decoder layer. 113 postprocessing: A postprocessing function to apply after the U-Net output. 114 check_shape: Whether to check the input shape to the U-Net forward call. 115 """ 116 def __init__( 117 self, 118 encoder: nn.Module, 119 base: nn.Module, 120 decoder: nn.Module, 121 out_conv: Optional[nn.Module] = None, 122 final_activation: Optional[Union[nn.Module, str]] = None, 123 postprocessing: Optional[Union[nn.Module, str]] = None, 124 check_shape: bool = True, 125 ): 126 super().__init__() 127 if len(encoder) != len(decoder): 128 raise ValueError(f"Incompatible depth of encoder (depth={len(encoder)}) and decoder (depth={len(decoder)})") 129 130 self.encoder = encoder 131 self.base = base 132 self.decoder = decoder 133 134 if out_conv is None: 135 self.return_decoder_outputs = False 136 self._out_channels = self.decoder.out_channels 137 elif isinstance(out_conv, nn.ModuleList): 138 if len(out_conv) != len(self.decoder): 139 raise ValueError(f"Invalid length of out_conv, expected {len(decoder)}, got {len(out_conv)}") 140 self.return_decoder_outputs = True 141 self._out_channels = [None if conv is None else conv.out_channels for conv in out_conv] 142 else: 143 self.return_decoder_outputs = False 144 self._out_channels = out_conv.out_channels 145 self.out_conv = out_conv 146 self.check_shape = check_shape 147 self.final_activation = self._get_activation(final_activation) 148 self.postprocessing = self._get_postprocessing(postprocessing) 149 150 @property 151 def in_channels(self): 152 return self.encoder.in_channels 153 154 @property 155 def out_channels(self): 156 return self._out_channels 157 158 @property 159 def depth(self): 160 return len(self.encoder) 161 162 def _get_activation(self, activation): 163 return_activation = None 164 if activation is None: 165 return None 166 if isinstance(activation, nn.Module): 167 return activation 168 if isinstance(activation, str): 169 return_activation = getattr(nn, activation, None) 170 if return_activation is None: 171 raise ValueError(f"Invalid activation: {activation}") 172 return return_activation() 173 174 def _get_postprocessing(self, postprocessing): 175 if postprocessing is None: 176 return None 177 elif isinstance(postprocessing, nn.Module): 178 return postprocessing 179 elif postprocessing in POSTPROCESSING: 180 return POSTPROCESSING[postprocessing]() 181 else: 182 raise ValueError(f"Invalid postprocessing: {postprocessing}") 183 184 # load encoder / decoder / base states for pretraining 185 def load_encoder_state(self, state): 186 self.encoder.load_state_dict(state) 187 188 def load_decoder_state(self, state): 189 self.decoder.load_state_dict(state) 190 191 def load_base_state(self, state): 192 self.base.load_state_dict(state) 193 194 def _apply_default(self, x): 195 self.encoder.return_outputs = True 196 self.decoder.return_outputs = False 197 198 x, encoder_out = self.encoder(x) 199 x = self.base(x) 200 x = self.decoder(x, encoder_inputs=encoder_out[::-1]) 201 202 if self.out_conv is not None: 203 x = self.out_conv(x) 204 if self.final_activation is not None: 205 x = self.final_activation(x) 206 if self.postprocessing is not None: 207 x = self.postprocessing(x) 208 209 return x 210 211 def _apply_with_side_outputs(self, x): 212 self.encoder.return_outputs = True 213 self.decoder.return_outputs = True 214 215 x, encoder_out = self.encoder(x) 216 x = self.base(x) 217 x = self.decoder(x, encoder_inputs=encoder_out[::-1]) 218 219 x = [x if conv is None else conv(xx) for xx, conv in zip(x, self.out_conv)] 220 if self.final_activation is not None: 221 x = [self.final_activation(xx) for xx in x] 222 223 if self.postprocessing is not None: 224 x = [self.postprocessing(xx) for xx in x] 225 226 # we reverse the list to have the full shape output as first element 227 return x[::-1] 228 229 def _check_shape(self, x): 230 spatial_shape = tuple(x.shape)[2:] 231 depth = len(self.encoder) 232 factor = [2**depth] * len(spatial_shape) 233 if any(sh % fac != 0 for sh, fac in zip(spatial_shape, factor)): 234 msg = f"Invalid shape for U-Net: {spatial_shape} is not divisible by {factor}" 235 raise ValueError(msg) 236 237 def forward(self, x: torch.Tensor) -> torch.Tensor: 238 """Apply U-Net to input data. 239 240 Args: 241 x: The input data. 242 243 Returns: 244 The output of the U-Net. 245 """ 246 # Cast input data to float, hotfix for modelzoo deployment issues, leaving it here for reference. 247 # x = x.float() 248 if getattr(self, "check_shape", True): 249 self._check_shape(x) 250 if self.return_decoder_outputs: 251 return self._apply_with_side_outputs(x) 252 else: 253 return self._apply_default(x) 254 255 256def _update_conv_kwargs(kwargs, scale_factor): 257 # if the scale factor is a scalar or all entries are the same we don"t need to update the kwargs 258 if isinstance(scale_factor, int) or scale_factor.count(scale_factor[0]) == len(scale_factor): 259 return kwargs 260 else: # otherwise set anisotropic kernel 261 kernel_size = kwargs.get("kernel_size", 3) 262 padding = kwargs.get("padding", 1) 263 264 # bail out if kernel size or padding aren"t scalars, because it"s 265 # unclear what to do in this case 266 if not (isinstance(kernel_size, int) and isinstance(padding, int)): 267 return kwargs 268 269 kernel_size = tuple(1 if factor == 1 else kernel_size for factor in scale_factor) 270 padding = tuple(0 if factor == 1 else padding for factor in scale_factor) 271 kwargs.update({"kernel_size": kernel_size, "padding": padding}) 272 return kwargs 273 274 275class Encoder(nn.Module): 276 """@private 277 """ 278 def __init__( 279 self, 280 features, 281 scale_factors, 282 conv_block_impl, 283 pooler_impl, 284 anisotropic_kernel=False, 285 **conv_block_kwargs 286 ): 287 super().__init__() 288 if len(features) != len(scale_factors) + 1: 289 raise ValueError("Incompatible number of features {len(features)} and scale_factors {len(scale_factors)}") 290 291 conv_kwargs = [conv_block_kwargs] * len(scale_factors) 292 if anisotropic_kernel: 293 conv_kwargs = [_update_conv_kwargs(kwargs, scale_factor) 294 for kwargs, scale_factor in zip(conv_kwargs, scale_factors)] 295 296 self.blocks = nn.ModuleList( 297 [conv_block_impl(inc, outc, **kwargs) 298 for inc, outc, kwargs in zip(features[:-1], features[1:], conv_kwargs)] 299 ) 300 self.poolers = nn.ModuleList( 301 [pooler_impl(factor) for factor in scale_factors] 302 ) 303 self.return_outputs = True 304 305 self.in_channels = features[0] 306 self.out_channels = features[-1] 307 308 def __len__(self): 309 return len(self.blocks) 310 311 def forward(self, x): 312 encoder_out = [] 313 for block, pooler in zip(self.blocks, self.poolers): 314 x = block(x) 315 encoder_out.append(x) 316 x = pooler(x) 317 318 if self.return_outputs: 319 return x, encoder_out 320 else: 321 return x 322 323 324class Decoder(nn.Module): 325 """@private 326 """ 327 def __init__( 328 self, 329 features, 330 scale_factors, 331 conv_block_impl, 332 sampler_impl, 333 skip_channels=None, 334 anisotropic_kernel=False, 335 **conv_block_kwargs 336 ): 337 super().__init__() 338 if len(features) != len(scale_factors) + 1: 339 raise ValueError("Incompatible number of features {len(features)} and scale_factors {len(scale_factors)}") 340 if skip_channels is not None and len(skip_channels) != len(scale_factors): 341 raise ValueError("Each decoder level must have a skip-channel count.") 342 343 conv_kwargs = [conv_block_kwargs] * len(scale_factors) 344 if anisotropic_kernel: 345 conv_kwargs = [_update_conv_kwargs(kwargs, scale_factor) 346 for kwargs, scale_factor in zip(conv_kwargs, scale_factors)] 347 348 self.explicit_skip_channels = skip_channels is not None 349 if self.explicit_skip_channels: 350 block_in_channels = [outc + skipc for outc, skipc in zip(features[1:], skip_channels)] 351 else: 352 block_in_channels = features[:-1] 353 self.blocks = nn.ModuleList( 354 [ 355 conv_block_impl(inc, outc, **kwargs) 356 for inc, outc, kwargs in zip(block_in_channels, features[1:], conv_kwargs) 357 ] 358 ) 359 self.samplers = nn.ModuleList( 360 [sampler_impl(factor, inc, outc) for factor, inc, outc 361 in zip(scale_factors, features[:-1], features[1:])] 362 ) 363 self.return_outputs = False 364 365 self.in_channels = features[0] 366 self.out_channels = features[-1] 367 368 def __len__(self): 369 return len(self.blocks) 370 371 # FIXME this prevents traces from being valid for other input sizes, need to find 372 # a solution to traceable cropping 373 def _crop(self, x, shape): 374 start = 2 if self.explicit_skip_channels else 1 375 shape_diff = [(xsh - sh) // 2 for xsh, sh in zip(x.shape[start:], shape[start:])] 376 crop = (slice(None),) * start + tuple( 377 slice(sd, sd + sh) for sd, sh in zip(shape_diff, shape[start:]) 378 ) 379 return x[crop] 380 # # Implementation with torch.narrow, does not fix the tracing warnings! 381 # for dim, (sh, sd) in enumerate(zip(shape, shape_diff)): 382 # x = torch.narrow(x, dim, sd, sh) 383 # return x 384 385 def _concat(self, x1, x2): 386 return torch.cat([x1, self._crop(x2, x1.shape)], dim=1) 387 388 def forward(self, x, encoder_inputs): 389 if len(encoder_inputs) != len(self.blocks): 390 raise ValueError(f"Invalid number of encoder_inputs: expect {len(self.blocks)}, got {len(encoder_inputs)}") 391 392 decoder_out = [] 393 for block, sampler, from_encoder in zip(self.blocks, self.samplers, encoder_inputs): 394 x = sampler(x) 395 x = block(self._concat(x, from_encoder)) 396 decoder_out.append(x) 397 398 if self.return_outputs: 399 return decoder_out + [x] 400 else: 401 return x 402 403 404def get_norm_layer(norm, dim, channels, n_groups=32): 405 """@private 406 """ 407 if norm is None: 408 return None 409 if norm == "InstanceNorm": 410 return nn.InstanceNorm2d(channels) if dim == 2 else nn.InstanceNorm3d(channels) 411 elif norm == "InstanceNormTrackStats": 412 kwargs = {"affine": True, "track_running_stats": True, "momentum": 0.01} 413 return nn.InstanceNorm2d(channels, **kwargs) if dim == 2 else nn.InstanceNorm3d(channels, **kwargs) 414 elif norm == "GroupNorm": 415 return nn.GroupNorm(min(n_groups, channels), channels) 416 elif norm == "BatchNorm": 417 return nn.BatchNorm2d(channels) if dim == 2 else nn.BatchNorm3d(channels) 418 else: 419 raise ValueError(f"Invalid norm: expect one of 'InstanceNorm', 'BatchNorm' or 'GroupNorm', got {norm}") 420 421 422class ConvBlock(nn.Module): 423 """@private 424 """ 425 def __init__(self, in_channels, out_channels, dim, kernel_size=3, padding=1, norm="InstanceNorm"): 426 super().__init__() 427 self.in_channels = in_channels 428 self.out_channels = out_channels 429 430 conv = nn.Conv2d if dim == 2 else nn.Conv3d 431 432 if norm is None: 433 self.block = nn.Sequential( 434 conv(in_channels, out_channels, 435 kernel_size=kernel_size, padding=padding), 436 nn.ReLU(inplace=True), 437 conv(out_channels, out_channels, 438 kernel_size=kernel_size, padding=padding), 439 nn.ReLU(inplace=True) 440 ) 441 else: 442 self.block = nn.Sequential( 443 get_norm_layer(norm, dim, in_channels), 444 conv(in_channels, out_channels, 445 kernel_size=kernel_size, padding=padding), 446 nn.ReLU(inplace=True), 447 get_norm_layer(norm, dim, out_channels), 448 conv(out_channels, out_channels, 449 kernel_size=kernel_size, padding=padding), 450 nn.ReLU(inplace=True) 451 ) 452 453 def forward(self, x): 454 return self.block(x) 455 456 457class Upsampler(nn.Module): 458 """@private 459 """ 460 def __init__(self, scale_factor, in_channels, out_channels, dim, mode): 461 super().__init__() 462 self.mode = mode 463 self.scale_factor = scale_factor 464 465 conv = nn.Conv2d if dim == 2 else nn.Conv3d 466 self.conv = conv(in_channels, out_channels, 1) 467 468 def forward(self, x): 469 x = nn.functional.interpolate(x, scale_factor=self.scale_factor, mode=self.mode, align_corners=False) 470 x = self.conv(x) 471 return x 472 473 474# 475# 2d unet implementations 476# 477 478class ConvBlock2d(ConvBlock): 479 """@private 480 """ 481 def __init__(self, in_channels, out_channels, **kwargs): 482 super().__init__(in_channels, out_channels, dim=2, **kwargs) 483 484 485class Upsampler2d(Upsampler): 486 """@private 487 """ 488 def __init__(self, scale_factor, 489 in_channels, out_channels, 490 mode="bilinear"): 491 super().__init__(scale_factor, in_channels, out_channels, dim=2, mode=mode) 492 493 494class UNet2d(UNetBase): 495 """A 2D U-Net network for segmentation and other image-to-image tasks. 496 497 The number of features for each level of the U-Net are computed as follows: initial_features * gain ** level. 498 The number of levels is determined by the depth argument. By default the U-Net uses two convolutional layers 499 per level, max-pooling for downsampling and linear interpolation for upsampling. 500 These implementations can be changed by providing arguments for `conv_block_impl`, `pooler_impl` 501 and `sampler_impl` respectively. 502 503 Args: 504 in_channels: The number of input image channels. 505 out_channels: The number of output image channels. 506 depth: The number of encoder / decoder levels of the U-Net. 507 initial_features: The initial number of features, corresponding to the features of the first conv block. 508 gain: The gain factor for increasing the features after each level. 509 final_activation: The activation applied after the output convolution or last decoder layer. 510 return_side_outputs: Whether to return the outputs after each decoder level. 511 conv_block_impl: The implementation of the convolutional block. 512 pooler_impl: The implementation of the pooling layer. 513 postprocessing: A postprocessing function to apply after the U-Net output. 514 check_shape: Whether to check the input shape to the U-Net forward call. 515 conv_block_kwargs: The keyword arguments for the convolutional block. 516 """ 517 def __init__( 518 self, 519 in_channels: int, 520 out_channels: int, 521 depth: int = 4, 522 initial_features: int = 32, 523 gain: int = 2, 524 final_activation=None, 525 return_side_outputs: bool = False, 526 conv_block_impl: nn.Module = ConvBlock2d, 527 pooler_impl: nn.Module = nn.MaxPool2d, 528 sampler_impl: nn.Module = Upsampler2d, 529 postprocessing: Optional[Union[nn.Module, str]] = None, 530 check_shape: bool = True, 531 **conv_block_kwargs, 532 ): 533 features_encoder = [in_channels] + [initial_features * gain ** i for i in range(depth)] 534 features_decoder = [initial_features * gain ** i for i in range(depth + 1)][::-1] 535 scale_factors = depth * [2] 536 537 if return_side_outputs: 538 if isinstance(out_channels, int) or out_channels is None: 539 out_channels = [out_channels] * depth 540 if len(out_channels) != depth: 541 raise ValueError() 542 out_conv = nn.ModuleList( 543 [nn.Conv2d(feat, outc, 1) for feat, outc in zip(features_decoder[1:], out_channels)] 544 ) 545 else: 546 out_conv = None if out_channels is None else nn.Conv2d(features_decoder[-1], out_channels, 1) 547 548 super().__init__( 549 encoder=Encoder( 550 features=features_encoder, 551 scale_factors=scale_factors, 552 conv_block_impl=conv_block_impl, 553 pooler_impl=pooler_impl, 554 **conv_block_kwargs 555 ), 556 decoder=Decoder( 557 features=features_decoder, 558 skip_channels=features_encoder[:0:-1], 559 scale_factors=scale_factors[::-1], 560 conv_block_impl=conv_block_impl, 561 sampler_impl=sampler_impl, 562 **conv_block_kwargs 563 ), 564 base=conv_block_impl( 565 features_encoder[-1], features_encoder[-1] * gain, 566 **conv_block_kwargs 567 ), 568 out_conv=out_conv, 569 final_activation=final_activation, 570 postprocessing=postprocessing, 571 check_shape=check_shape, 572 ) 573 self.init_kwargs = {"in_channels": in_channels, "out_channels": out_channels, "depth": depth, 574 "initial_features": initial_features, "gain": gain, 575 "final_activation": final_activation, "return_side_outputs": return_side_outputs, 576 "conv_block_impl": conv_block_impl, "pooler_impl": pooler_impl, 577 "sampler_impl": sampler_impl, "postprocessing": postprocessing, **conv_block_kwargs} 578 579 580# 581# 3d unet implementations 582# 583 584class ConvBlock3d(ConvBlock): 585 """@private 586 """ 587 def __init__(self, in_channels, out_channels, **kwargs): 588 super().__init__(in_channels, out_channels, dim=3, **kwargs) 589 590 591class Upsampler3d(Upsampler): 592 """@private 593 """ 594 def __init__(self, scale_factor, in_channels, out_channels, mode="trilinear"): 595 super().__init__(scale_factor, in_channels, out_channels, dim=3, mode=mode) 596 597 598class AnisotropicUNet(UNetBase): 599 """A 3D U-Net network for segmentation and other image-to-image tasks. 600 601 The number of features for each level of the U-Net are computed as follows: initial_features * gain ** level. 602 The number of levels is determined by the length of the scale_factors argument. 603 The scale factors determine the pooling factors for each level. By specifying [1, 2, 2] the pooling 604 is done in an anisotropic fashion, i.e. only across the xy-plane, 605 by specifying [2, 2, 2] it is done in an isotropic fashion. 606 607 By default the U-Net uses two convolutional layers per level. 608 This can be changed by providing an argument for `conv_block_impl`. 609 610 Args: 611 in_channels: The number of input image channels. 612 out_channels: The number of output image channels. 613 scale_factors: The factors for max pooling for the levels of the U-Net. 614 initial_features: The initial number of features, corresponding to the features of the first conv block. 615 gain: The gain factor for increasing the features after each level. 616 final_activation: The activation applied after the output convolution or last decoder layer. 617 return_side_outputs: Whether to return the outputs after each decoder level. 618 conv_block_impl: The implementation of the convolutional block. 619 anisotropic_kernel: Whether to use an anisotropic kernel in addition to anisotropic scaling factor. 620 postprocessing: A postprocessing function to apply after the U-Net output. 621 check_shape: Whether to check the input shape to the U-Net forward call. 622 conv_block_kwargs: The keyword arguments for the convolutional block. 623 """ 624 def __init__( 625 self, 626 in_channels: int, 627 out_channels: int, 628 scale_factors: List[List[int]], 629 initial_features: int = 32, 630 gain: int = 2, 631 final_activation: Optional[Union[str, nn.Module]] = None, 632 return_side_outputs: bool = False, 633 conv_block_impl: nn.Module = ConvBlock3d, 634 anisotropic_kernel: bool = False, 635 postprocessing: Optional[Union[str, nn.Module]] = None, 636 check_shape: bool = True, 637 **conv_block_kwargs, 638 ): 639 depth = len(scale_factors) 640 features_encoder = [in_channels] + [initial_features * gain ** i for i in range(depth)] 641 features_decoder = [initial_features * gain ** i for i in range(depth + 1)][::-1] 642 643 if return_side_outputs: 644 if isinstance(out_channels, int) or out_channels is None: 645 out_channels = [out_channels] * depth 646 if len(out_channels) != depth: 647 raise ValueError() 648 out_conv = nn.ModuleList( 649 [nn.Conv3d(feat, outc, 1) for feat, outc in zip(features_decoder[1:], out_channels)] 650 ) 651 else: 652 out_conv = None if out_channels is None else nn.Conv3d(features_decoder[-1], out_channels, 1) 653 654 super().__init__( 655 encoder=Encoder( 656 features=features_encoder, 657 scale_factors=scale_factors, 658 conv_block_impl=conv_block_impl, 659 pooler_impl=nn.MaxPool3d, 660 anisotropic_kernel=anisotropic_kernel, 661 **conv_block_kwargs 662 ), 663 decoder=Decoder( 664 features=features_decoder, 665 skip_channels=features_encoder[:0:-1], 666 scale_factors=scale_factors[::-1], 667 conv_block_impl=conv_block_impl, 668 sampler_impl=Upsampler3d, 669 anisotropic_kernel=anisotropic_kernel, 670 **conv_block_kwargs 671 ), 672 base=conv_block_impl( 673 features_encoder[-1], features_encoder[-1] * gain, **conv_block_kwargs 674 ), 675 out_conv=out_conv, 676 final_activation=final_activation, 677 postprocessing=postprocessing, 678 check_shape=check_shape, 679 ) 680 self.init_kwargs = {"in_channels": in_channels, "out_channels": out_channels, "scale_factors": scale_factors, 681 "initial_features": initial_features, "gain": gain, 682 "final_activation": final_activation, "return_side_outputs": return_side_outputs, 683 "conv_block_impl": conv_block_impl, "anisotropic_kernel": anisotropic_kernel, 684 "postprocessing": postprocessing, **conv_block_kwargs} 685 686 def _check_shape(self, x): 687 spatial_shape = tuple(x.shape)[2:] 688 scale_factors = self.init_kwargs.get("scale_factors", [[2, 2, 2]]*len(self.encoder)) 689 factor = [int(np.prod([sf[i] for sf in scale_factors])) for i in range(3)] 690 if len(spatial_shape) != len(factor): 691 msg = f"Invalid shape for U-Net: dimensions don't agree {len(spatial_shape)} != {len(factor)}" 692 raise ValueError(msg) 693 if any(sh % fac != 0 for sh, fac in zip(spatial_shape, factor)): 694 msg = f"Invalid shape for U-Net: {spatial_shape} is not divisible by {factor}" 695 raise ValueError(msg) 696 697 698class UNet3d(AnisotropicUNet): 699 """A 3D U-Net network for segmentation and other image-to-image tasks. 700 701 This class uses the same implementation as `AnisotropicUNet`, with isotropic scaling in each level. 702 703 Args: 704 in_channels: The number of input image channels. 705 out_channels: The number of output image channels. 706 depth: The number of encoder / decoder levels of the U-Net. 707 initial_features: The initial number of features, corresponding to the features of the first conv block. 708 gain: The gain factor for increasing the features after each level. 709 final_activation: The activation applied after the output convolution or last decoder layer. 710 return_side_outputs: Whether to return the outputs after each decoder level. 711 conv_block_impl: The implementation of the convolutional block. 712 postprocessing: A postprocessing function to apply after the U-Net output. 713 check_shape: Whether to check the input shape to the U-Net forward call. 714 conv_block_kwargs: The keyword arguments for the convolutional block. 715 """ 716 def __init__( 717 self, 718 in_channels: int, 719 out_channels: int, 720 depth: int = 4, 721 initial_features: int = 32, 722 gain: int = 2, 723 final_activation: Optional[Union[str, nn.Module]] = None, 724 return_side_outputs: bool = False, 725 conv_block_impl: nn.Module = ConvBlock3d, 726 postprocessing: Optional[Union[str, nn.Module]] = None, 727 check_shape: bool = True, 728 **conv_block_kwargs, 729 ): 730 scale_factors = depth * [2] 731 super().__init__(in_channels, out_channels, scale_factors, 732 initial_features=initial_features, gain=gain, 733 final_activation=final_activation, 734 return_side_outputs=return_side_outputs, 735 anisotropic_kernel=False, 736 postprocessing=postprocessing, 737 conv_block_impl=conv_block_impl, 738 check_shape=check_shape, 739 **conv_block_kwargs) 740 self.init_kwargs = {"in_channels": in_channels, "out_channels": out_channels, "depth": depth, 741 "initial_features": initial_features, "gain": gain, 742 "final_activation": final_activation, "return_side_outputs": return_side_outputs, 743 "conv_block_impl": conv_block_impl, "postprocessing": postprocessing, **conv_block_kwargs}
105class UNetBase(nn.Module): 106 """Base class for implementing a U-Net. 107 108 Args: 109 encoder: The encoder of the U-Net. 110 base: The base layer of the U-Net. 111 decoder: The decoder of the U-Net. 112 out_conv: The output convolution applied after the last decoder layer. 113 final_activation: The activation applied after the output convolution or last decoder layer. 114 postprocessing: A postprocessing function to apply after the U-Net output. 115 check_shape: Whether to check the input shape to the U-Net forward call. 116 """ 117 def __init__( 118 self, 119 encoder: nn.Module, 120 base: nn.Module, 121 decoder: nn.Module, 122 out_conv: Optional[nn.Module] = None, 123 final_activation: Optional[Union[nn.Module, str]] = None, 124 postprocessing: Optional[Union[nn.Module, str]] = None, 125 check_shape: bool = True, 126 ): 127 super().__init__() 128 if len(encoder) != len(decoder): 129 raise ValueError(f"Incompatible depth of encoder (depth={len(encoder)}) and decoder (depth={len(decoder)})") 130 131 self.encoder = encoder 132 self.base = base 133 self.decoder = decoder 134 135 if out_conv is None: 136 self.return_decoder_outputs = False 137 self._out_channels = self.decoder.out_channels 138 elif isinstance(out_conv, nn.ModuleList): 139 if len(out_conv) != len(self.decoder): 140 raise ValueError(f"Invalid length of out_conv, expected {len(decoder)}, got {len(out_conv)}") 141 self.return_decoder_outputs = True 142 self._out_channels = [None if conv is None else conv.out_channels for conv in out_conv] 143 else: 144 self.return_decoder_outputs = False 145 self._out_channels = out_conv.out_channels 146 self.out_conv = out_conv 147 self.check_shape = check_shape 148 self.final_activation = self._get_activation(final_activation) 149 self.postprocessing = self._get_postprocessing(postprocessing) 150 151 @property 152 def in_channels(self): 153 return self.encoder.in_channels 154 155 @property 156 def out_channels(self): 157 return self._out_channels 158 159 @property 160 def depth(self): 161 return len(self.encoder) 162 163 def _get_activation(self, activation): 164 return_activation = None 165 if activation is None: 166 return None 167 if isinstance(activation, nn.Module): 168 return activation 169 if isinstance(activation, str): 170 return_activation = getattr(nn, activation, None) 171 if return_activation is None: 172 raise ValueError(f"Invalid activation: {activation}") 173 return return_activation() 174 175 def _get_postprocessing(self, postprocessing): 176 if postprocessing is None: 177 return None 178 elif isinstance(postprocessing, nn.Module): 179 return postprocessing 180 elif postprocessing in POSTPROCESSING: 181 return POSTPROCESSING[postprocessing]() 182 else: 183 raise ValueError(f"Invalid postprocessing: {postprocessing}") 184 185 # load encoder / decoder / base states for pretraining 186 def load_encoder_state(self, state): 187 self.encoder.load_state_dict(state) 188 189 def load_decoder_state(self, state): 190 self.decoder.load_state_dict(state) 191 192 def load_base_state(self, state): 193 self.base.load_state_dict(state) 194 195 def _apply_default(self, x): 196 self.encoder.return_outputs = True 197 self.decoder.return_outputs = False 198 199 x, encoder_out = self.encoder(x) 200 x = self.base(x) 201 x = self.decoder(x, encoder_inputs=encoder_out[::-1]) 202 203 if self.out_conv is not None: 204 x = self.out_conv(x) 205 if self.final_activation is not None: 206 x = self.final_activation(x) 207 if self.postprocessing is not None: 208 x = self.postprocessing(x) 209 210 return x 211 212 def _apply_with_side_outputs(self, x): 213 self.encoder.return_outputs = True 214 self.decoder.return_outputs = True 215 216 x, encoder_out = self.encoder(x) 217 x = self.base(x) 218 x = self.decoder(x, encoder_inputs=encoder_out[::-1]) 219 220 x = [x if conv is None else conv(xx) for xx, conv in zip(x, self.out_conv)] 221 if self.final_activation is not None: 222 x = [self.final_activation(xx) for xx in x] 223 224 if self.postprocessing is not None: 225 x = [self.postprocessing(xx) for xx in x] 226 227 # we reverse the list to have the full shape output as first element 228 return x[::-1] 229 230 def _check_shape(self, x): 231 spatial_shape = tuple(x.shape)[2:] 232 depth = len(self.encoder) 233 factor = [2**depth] * len(spatial_shape) 234 if any(sh % fac != 0 for sh, fac in zip(spatial_shape, factor)): 235 msg = f"Invalid shape for U-Net: {spatial_shape} is not divisible by {factor}" 236 raise ValueError(msg) 237 238 def forward(self, x: torch.Tensor) -> torch.Tensor: 239 """Apply U-Net to input data. 240 241 Args: 242 x: The input data. 243 244 Returns: 245 The output of the U-Net. 246 """ 247 # Cast input data to float, hotfix for modelzoo deployment issues, leaving it here for reference. 248 # x = x.float() 249 if getattr(self, "check_shape", True): 250 self._check_shape(x) 251 if self.return_decoder_outputs: 252 return self._apply_with_side_outputs(x) 253 else: 254 return self._apply_default(x)
Base class for implementing a U-Net.
Arguments:
- encoder: The encoder of the U-Net.
- base: The base layer of the U-Net.
- decoder: The decoder of the U-Net.
- out_conv: The output convolution applied after the last decoder layer.
- final_activation: The activation applied after the output convolution or last decoder layer.
- postprocessing: A postprocessing function to apply after the U-Net output.
- check_shape: Whether to check the input shape to the U-Net forward call.
117 def __init__( 118 self, 119 encoder: nn.Module, 120 base: nn.Module, 121 decoder: nn.Module, 122 out_conv: Optional[nn.Module] = None, 123 final_activation: Optional[Union[nn.Module, str]] = None, 124 postprocessing: Optional[Union[nn.Module, str]] = None, 125 check_shape: bool = True, 126 ): 127 super().__init__() 128 if len(encoder) != len(decoder): 129 raise ValueError(f"Incompatible depth of encoder (depth={len(encoder)}) and decoder (depth={len(decoder)})") 130 131 self.encoder = encoder 132 self.base = base 133 self.decoder = decoder 134 135 if out_conv is None: 136 self.return_decoder_outputs = False 137 self._out_channels = self.decoder.out_channels 138 elif isinstance(out_conv, nn.ModuleList): 139 if len(out_conv) != len(self.decoder): 140 raise ValueError(f"Invalid length of out_conv, expected {len(decoder)}, got {len(out_conv)}") 141 self.return_decoder_outputs = True 142 self._out_channels = [None if conv is None else conv.out_channels for conv in out_conv] 143 else: 144 self.return_decoder_outputs = False 145 self._out_channels = out_conv.out_channels 146 self.out_conv = out_conv 147 self.check_shape = check_shape 148 self.final_activation = self._get_activation(final_activation) 149 self.postprocessing = self._get_postprocessing(postprocessing)
Initialize internal Module state, shared by both nn.Module and ScriptModule.
238 def forward(self, x: torch.Tensor) -> torch.Tensor: 239 """Apply U-Net to input data. 240 241 Args: 242 x: The input data. 243 244 Returns: 245 The output of the U-Net. 246 """ 247 # Cast input data to float, hotfix for modelzoo deployment issues, leaving it here for reference. 248 # x = x.float() 249 if getattr(self, "check_shape", True): 250 self._check_shape(x) 251 if self.return_decoder_outputs: 252 return self._apply_with_side_outputs(x) 253 else: 254 return self._apply_default(x)
Apply U-Net to input data.
Arguments:
- x: The input data.
Returns:
The output of the U-Net.
495class UNet2d(UNetBase): 496 """A 2D U-Net network for segmentation and other image-to-image tasks. 497 498 The number of features for each level of the U-Net are computed as follows: initial_features * gain ** level. 499 The number of levels is determined by the depth argument. By default the U-Net uses two convolutional layers 500 per level, max-pooling for downsampling and linear interpolation for upsampling. 501 These implementations can be changed by providing arguments for `conv_block_impl`, `pooler_impl` 502 and `sampler_impl` respectively. 503 504 Args: 505 in_channels: The number of input image channels. 506 out_channels: The number of output image channels. 507 depth: The number of encoder / decoder levels of the U-Net. 508 initial_features: The initial number of features, corresponding to the features of the first conv block. 509 gain: The gain factor for increasing the features after each level. 510 final_activation: The activation applied after the output convolution or last decoder layer. 511 return_side_outputs: Whether to return the outputs after each decoder level. 512 conv_block_impl: The implementation of the convolutional block. 513 pooler_impl: The implementation of the pooling layer. 514 postprocessing: A postprocessing function to apply after the U-Net output. 515 check_shape: Whether to check the input shape to the U-Net forward call. 516 conv_block_kwargs: The keyword arguments for the convolutional block. 517 """ 518 def __init__( 519 self, 520 in_channels: int, 521 out_channels: int, 522 depth: int = 4, 523 initial_features: int = 32, 524 gain: int = 2, 525 final_activation=None, 526 return_side_outputs: bool = False, 527 conv_block_impl: nn.Module = ConvBlock2d, 528 pooler_impl: nn.Module = nn.MaxPool2d, 529 sampler_impl: nn.Module = Upsampler2d, 530 postprocessing: Optional[Union[nn.Module, str]] = None, 531 check_shape: bool = True, 532 **conv_block_kwargs, 533 ): 534 features_encoder = [in_channels] + [initial_features * gain ** i for i in range(depth)] 535 features_decoder = [initial_features * gain ** i for i in range(depth + 1)][::-1] 536 scale_factors = depth * [2] 537 538 if return_side_outputs: 539 if isinstance(out_channels, int) or out_channels is None: 540 out_channels = [out_channels] * depth 541 if len(out_channels) != depth: 542 raise ValueError() 543 out_conv = nn.ModuleList( 544 [nn.Conv2d(feat, outc, 1) for feat, outc in zip(features_decoder[1:], out_channels)] 545 ) 546 else: 547 out_conv = None if out_channels is None else nn.Conv2d(features_decoder[-1], out_channels, 1) 548 549 super().__init__( 550 encoder=Encoder( 551 features=features_encoder, 552 scale_factors=scale_factors, 553 conv_block_impl=conv_block_impl, 554 pooler_impl=pooler_impl, 555 **conv_block_kwargs 556 ), 557 decoder=Decoder( 558 features=features_decoder, 559 skip_channels=features_encoder[:0:-1], 560 scale_factors=scale_factors[::-1], 561 conv_block_impl=conv_block_impl, 562 sampler_impl=sampler_impl, 563 **conv_block_kwargs 564 ), 565 base=conv_block_impl( 566 features_encoder[-1], features_encoder[-1] * gain, 567 **conv_block_kwargs 568 ), 569 out_conv=out_conv, 570 final_activation=final_activation, 571 postprocessing=postprocessing, 572 check_shape=check_shape, 573 ) 574 self.init_kwargs = {"in_channels": in_channels, "out_channels": out_channels, "depth": depth, 575 "initial_features": initial_features, "gain": gain, 576 "final_activation": final_activation, "return_side_outputs": return_side_outputs, 577 "conv_block_impl": conv_block_impl, "pooler_impl": pooler_impl, 578 "sampler_impl": sampler_impl, "postprocessing": postprocessing, **conv_block_kwargs}
A 2D U-Net network for segmentation and other image-to-image tasks.
The number of features for each level of the U-Net are computed as follows: initial_features * gain ** level.
The number of levels is determined by the depth argument. By default the U-Net uses two convolutional layers
per level, max-pooling for downsampling and linear interpolation for upsampling.
These implementations can be changed by providing arguments for conv_block_impl, pooler_impl
and sampler_impl respectively.
Arguments:
- in_channels: The number of input image channels.
- out_channels: The number of output image channels.
- depth: The number of encoder / decoder levels of the U-Net.
- initial_features: The initial number of features, corresponding to the features of the first conv block.
- gain: The gain factor for increasing the features after each level.
- final_activation: The activation applied after the output convolution or last decoder layer.
- return_side_outputs: Whether to return the outputs after each decoder level.
- conv_block_impl: The implementation of the convolutional block.
- pooler_impl: The implementation of the pooling layer.
- postprocessing: A postprocessing function to apply after the U-Net output.
- check_shape: Whether to check the input shape to the U-Net forward call.
- conv_block_kwargs: The keyword arguments for the convolutional block.
518 def __init__( 519 self, 520 in_channels: int, 521 out_channels: int, 522 depth: int = 4, 523 initial_features: int = 32, 524 gain: int = 2, 525 final_activation=None, 526 return_side_outputs: bool = False, 527 conv_block_impl: nn.Module = ConvBlock2d, 528 pooler_impl: nn.Module = nn.MaxPool2d, 529 sampler_impl: nn.Module = Upsampler2d, 530 postprocessing: Optional[Union[nn.Module, str]] = None, 531 check_shape: bool = True, 532 **conv_block_kwargs, 533 ): 534 features_encoder = [in_channels] + [initial_features * gain ** i for i in range(depth)] 535 features_decoder = [initial_features * gain ** i for i in range(depth + 1)][::-1] 536 scale_factors = depth * [2] 537 538 if return_side_outputs: 539 if isinstance(out_channels, int) or out_channels is None: 540 out_channels = [out_channels] * depth 541 if len(out_channels) != depth: 542 raise ValueError() 543 out_conv = nn.ModuleList( 544 [nn.Conv2d(feat, outc, 1) for feat, outc in zip(features_decoder[1:], out_channels)] 545 ) 546 else: 547 out_conv = None if out_channels is None else nn.Conv2d(features_decoder[-1], out_channels, 1) 548 549 super().__init__( 550 encoder=Encoder( 551 features=features_encoder, 552 scale_factors=scale_factors, 553 conv_block_impl=conv_block_impl, 554 pooler_impl=pooler_impl, 555 **conv_block_kwargs 556 ), 557 decoder=Decoder( 558 features=features_decoder, 559 skip_channels=features_encoder[:0:-1], 560 scale_factors=scale_factors[::-1], 561 conv_block_impl=conv_block_impl, 562 sampler_impl=sampler_impl, 563 **conv_block_kwargs 564 ), 565 base=conv_block_impl( 566 features_encoder[-1], features_encoder[-1] * gain, 567 **conv_block_kwargs 568 ), 569 out_conv=out_conv, 570 final_activation=final_activation, 571 postprocessing=postprocessing, 572 check_shape=check_shape, 573 ) 574 self.init_kwargs = {"in_channels": in_channels, "out_channels": out_channels, "depth": depth, 575 "initial_features": initial_features, "gain": gain, 576 "final_activation": final_activation, "return_side_outputs": return_side_outputs, 577 "conv_block_impl": conv_block_impl, "pooler_impl": pooler_impl, 578 "sampler_impl": sampler_impl, "postprocessing": postprocessing, **conv_block_kwargs}
Initialize internal Module state, shared by both nn.Module and ScriptModule.
599class AnisotropicUNet(UNetBase): 600 """A 3D U-Net network for segmentation and other image-to-image tasks. 601 602 The number of features for each level of the U-Net are computed as follows: initial_features * gain ** level. 603 The number of levels is determined by the length of the scale_factors argument. 604 The scale factors determine the pooling factors for each level. By specifying [1, 2, 2] the pooling 605 is done in an anisotropic fashion, i.e. only across the xy-plane, 606 by specifying [2, 2, 2] it is done in an isotropic fashion. 607 608 By default the U-Net uses two convolutional layers per level. 609 This can be changed by providing an argument for `conv_block_impl`. 610 611 Args: 612 in_channels: The number of input image channels. 613 out_channels: The number of output image channels. 614 scale_factors: The factors for max pooling for the levels of the U-Net. 615 initial_features: The initial number of features, corresponding to the features of the first conv block. 616 gain: The gain factor for increasing the features after each level. 617 final_activation: The activation applied after the output convolution or last decoder layer. 618 return_side_outputs: Whether to return the outputs after each decoder level. 619 conv_block_impl: The implementation of the convolutional block. 620 anisotropic_kernel: Whether to use an anisotropic kernel in addition to anisotropic scaling factor. 621 postprocessing: A postprocessing function to apply after the U-Net output. 622 check_shape: Whether to check the input shape to the U-Net forward call. 623 conv_block_kwargs: The keyword arguments for the convolutional block. 624 """ 625 def __init__( 626 self, 627 in_channels: int, 628 out_channels: int, 629 scale_factors: List[List[int]], 630 initial_features: int = 32, 631 gain: int = 2, 632 final_activation: Optional[Union[str, nn.Module]] = None, 633 return_side_outputs: bool = False, 634 conv_block_impl: nn.Module = ConvBlock3d, 635 anisotropic_kernel: bool = False, 636 postprocessing: Optional[Union[str, nn.Module]] = None, 637 check_shape: bool = True, 638 **conv_block_kwargs, 639 ): 640 depth = len(scale_factors) 641 features_encoder = [in_channels] + [initial_features * gain ** i for i in range(depth)] 642 features_decoder = [initial_features * gain ** i for i in range(depth + 1)][::-1] 643 644 if return_side_outputs: 645 if isinstance(out_channels, int) or out_channels is None: 646 out_channels = [out_channels] * depth 647 if len(out_channels) != depth: 648 raise ValueError() 649 out_conv = nn.ModuleList( 650 [nn.Conv3d(feat, outc, 1) for feat, outc in zip(features_decoder[1:], out_channels)] 651 ) 652 else: 653 out_conv = None if out_channels is None else nn.Conv3d(features_decoder[-1], out_channels, 1) 654 655 super().__init__( 656 encoder=Encoder( 657 features=features_encoder, 658 scale_factors=scale_factors, 659 conv_block_impl=conv_block_impl, 660 pooler_impl=nn.MaxPool3d, 661 anisotropic_kernel=anisotropic_kernel, 662 **conv_block_kwargs 663 ), 664 decoder=Decoder( 665 features=features_decoder, 666 skip_channels=features_encoder[:0:-1], 667 scale_factors=scale_factors[::-1], 668 conv_block_impl=conv_block_impl, 669 sampler_impl=Upsampler3d, 670 anisotropic_kernel=anisotropic_kernel, 671 **conv_block_kwargs 672 ), 673 base=conv_block_impl( 674 features_encoder[-1], features_encoder[-1] * gain, **conv_block_kwargs 675 ), 676 out_conv=out_conv, 677 final_activation=final_activation, 678 postprocessing=postprocessing, 679 check_shape=check_shape, 680 ) 681 self.init_kwargs = {"in_channels": in_channels, "out_channels": out_channels, "scale_factors": scale_factors, 682 "initial_features": initial_features, "gain": gain, 683 "final_activation": final_activation, "return_side_outputs": return_side_outputs, 684 "conv_block_impl": conv_block_impl, "anisotropic_kernel": anisotropic_kernel, 685 "postprocessing": postprocessing, **conv_block_kwargs} 686 687 def _check_shape(self, x): 688 spatial_shape = tuple(x.shape)[2:] 689 scale_factors = self.init_kwargs.get("scale_factors", [[2, 2, 2]]*len(self.encoder)) 690 factor = [int(np.prod([sf[i] for sf in scale_factors])) for i in range(3)] 691 if len(spatial_shape) != len(factor): 692 msg = f"Invalid shape for U-Net: dimensions don't agree {len(spatial_shape)} != {len(factor)}" 693 raise ValueError(msg) 694 if any(sh % fac != 0 for sh, fac in zip(spatial_shape, factor)): 695 msg = f"Invalid shape for U-Net: {spatial_shape} is not divisible by {factor}" 696 raise ValueError(msg)
A 3D U-Net network for segmentation and other image-to-image tasks.
The number of features for each level of the U-Net are computed as follows: initial_features * gain ** level. The number of levels is determined by the length of the scale_factors argument. The scale factors determine the pooling factors for each level. By specifying [1, 2, 2] the pooling is done in an anisotropic fashion, i.e. only across the xy-plane, by specifying [2, 2, 2] it is done in an isotropic fashion.
By default the U-Net uses two convolutional layers per level.
This can be changed by providing an argument for conv_block_impl.
Arguments:
- in_channels: The number of input image channels.
- out_channels: The number of output image channels.
- scale_factors: The factors for max pooling for the levels of the U-Net.
- initial_features: The initial number of features, corresponding to the features of the first conv block.
- gain: The gain factor for increasing the features after each level.
- final_activation: The activation applied after the output convolution or last decoder layer.
- return_side_outputs: Whether to return the outputs after each decoder level.
- conv_block_impl: The implementation of the convolutional block.
- anisotropic_kernel: Whether to use an anisotropic kernel in addition to anisotropic scaling factor.
- postprocessing: A postprocessing function to apply after the U-Net output.
- check_shape: Whether to check the input shape to the U-Net forward call.
- conv_block_kwargs: The keyword arguments for the convolutional block.
625 def __init__( 626 self, 627 in_channels: int, 628 out_channels: int, 629 scale_factors: List[List[int]], 630 initial_features: int = 32, 631 gain: int = 2, 632 final_activation: Optional[Union[str, nn.Module]] = None, 633 return_side_outputs: bool = False, 634 conv_block_impl: nn.Module = ConvBlock3d, 635 anisotropic_kernel: bool = False, 636 postprocessing: Optional[Union[str, nn.Module]] = None, 637 check_shape: bool = True, 638 **conv_block_kwargs, 639 ): 640 depth = len(scale_factors) 641 features_encoder = [in_channels] + [initial_features * gain ** i for i in range(depth)] 642 features_decoder = [initial_features * gain ** i for i in range(depth + 1)][::-1] 643 644 if return_side_outputs: 645 if isinstance(out_channels, int) or out_channels is None: 646 out_channels = [out_channels] * depth 647 if len(out_channels) != depth: 648 raise ValueError() 649 out_conv = nn.ModuleList( 650 [nn.Conv3d(feat, outc, 1) for feat, outc in zip(features_decoder[1:], out_channels)] 651 ) 652 else: 653 out_conv = None if out_channels is None else nn.Conv3d(features_decoder[-1], out_channels, 1) 654 655 super().__init__( 656 encoder=Encoder( 657 features=features_encoder, 658 scale_factors=scale_factors, 659 conv_block_impl=conv_block_impl, 660 pooler_impl=nn.MaxPool3d, 661 anisotropic_kernel=anisotropic_kernel, 662 **conv_block_kwargs 663 ), 664 decoder=Decoder( 665 features=features_decoder, 666 skip_channels=features_encoder[:0:-1], 667 scale_factors=scale_factors[::-1], 668 conv_block_impl=conv_block_impl, 669 sampler_impl=Upsampler3d, 670 anisotropic_kernel=anisotropic_kernel, 671 **conv_block_kwargs 672 ), 673 base=conv_block_impl( 674 features_encoder[-1], features_encoder[-1] * gain, **conv_block_kwargs 675 ), 676 out_conv=out_conv, 677 final_activation=final_activation, 678 postprocessing=postprocessing, 679 check_shape=check_shape, 680 ) 681 self.init_kwargs = {"in_channels": in_channels, "out_channels": out_channels, "scale_factors": scale_factors, 682 "initial_features": initial_features, "gain": gain, 683 "final_activation": final_activation, "return_side_outputs": return_side_outputs, 684 "conv_block_impl": conv_block_impl, "anisotropic_kernel": anisotropic_kernel, 685 "postprocessing": postprocessing, **conv_block_kwargs}
Initialize internal Module state, shared by both nn.Module and ScriptModule.
699class UNet3d(AnisotropicUNet): 700 """A 3D U-Net network for segmentation and other image-to-image tasks. 701 702 This class uses the same implementation as `AnisotropicUNet`, with isotropic scaling in each level. 703 704 Args: 705 in_channels: The number of input image channels. 706 out_channels: The number of output image channels. 707 depth: The number of encoder / decoder levels of the U-Net. 708 initial_features: The initial number of features, corresponding to the features of the first conv block. 709 gain: The gain factor for increasing the features after each level. 710 final_activation: The activation applied after the output convolution or last decoder layer. 711 return_side_outputs: Whether to return the outputs after each decoder level. 712 conv_block_impl: The implementation of the convolutional block. 713 postprocessing: A postprocessing function to apply after the U-Net output. 714 check_shape: Whether to check the input shape to the U-Net forward call. 715 conv_block_kwargs: The keyword arguments for the convolutional block. 716 """ 717 def __init__( 718 self, 719 in_channels: int, 720 out_channels: int, 721 depth: int = 4, 722 initial_features: int = 32, 723 gain: int = 2, 724 final_activation: Optional[Union[str, nn.Module]] = None, 725 return_side_outputs: bool = False, 726 conv_block_impl: nn.Module = ConvBlock3d, 727 postprocessing: Optional[Union[str, nn.Module]] = None, 728 check_shape: bool = True, 729 **conv_block_kwargs, 730 ): 731 scale_factors = depth * [2] 732 super().__init__(in_channels, out_channels, scale_factors, 733 initial_features=initial_features, gain=gain, 734 final_activation=final_activation, 735 return_side_outputs=return_side_outputs, 736 anisotropic_kernel=False, 737 postprocessing=postprocessing, 738 conv_block_impl=conv_block_impl, 739 check_shape=check_shape, 740 **conv_block_kwargs) 741 self.init_kwargs = {"in_channels": in_channels, "out_channels": out_channels, "depth": depth, 742 "initial_features": initial_features, "gain": gain, 743 "final_activation": final_activation, "return_side_outputs": return_side_outputs, 744 "conv_block_impl": conv_block_impl, "postprocessing": postprocessing, **conv_block_kwargs}
A 3D U-Net network for segmentation and other image-to-image tasks.
This class uses the same implementation as AnisotropicUNet, with isotropic scaling in each level.
Arguments:
- in_channels: The number of input image channels.
- out_channels: The number of output image channels.
- depth: The number of encoder / decoder levels of the U-Net.
- initial_features: The initial number of features, corresponding to the features of the first conv block.
- gain: The gain factor for increasing the features after each level.
- final_activation: The activation applied after the output convolution or last decoder layer.
- return_side_outputs: Whether to return the outputs after each decoder level.
- conv_block_impl: The implementation of the convolutional block.
- postprocessing: A postprocessing function to apply after the U-Net output.
- check_shape: Whether to check the input shape to the U-Net forward call.
- conv_block_kwargs: The keyword arguments for the convolutional block.
717 def __init__( 718 self, 719 in_channels: int, 720 out_channels: int, 721 depth: int = 4, 722 initial_features: int = 32, 723 gain: int = 2, 724 final_activation: Optional[Union[str, nn.Module]] = None, 725 return_side_outputs: bool = False, 726 conv_block_impl: nn.Module = ConvBlock3d, 727 postprocessing: Optional[Union[str, nn.Module]] = None, 728 check_shape: bool = True, 729 **conv_block_kwargs, 730 ): 731 scale_factors = depth * [2] 732 super().__init__(in_channels, out_channels, scale_factors, 733 initial_features=initial_features, gain=gain, 734 final_activation=final_activation, 735 return_side_outputs=return_side_outputs, 736 anisotropic_kernel=False, 737 postprocessing=postprocessing, 738 conv_block_impl=conv_block_impl, 739 check_shape=check_shape, 740 **conv_block_kwargs) 741 self.init_kwargs = {"in_channels": in_channels, "out_channels": out_channels, "depth": depth, 742 "initial_features": initial_features, "gain": gain, 743 "final_activation": final_activation, "return_side_outputs": return_side_outputs, 744 "conv_block_impl": conv_block_impl, "postprocessing": postprocessing, **conv_block_kwargs}
Initialize internal Module state, shared by both nn.Module and ScriptModule.