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}
class UNetBase(torch.nn.modules.module.Module):
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.
UNetBase( encoder: torch.nn.modules.module.Module, base: torch.nn.modules.module.Module, decoder: torch.nn.modules.module.Module, out_conv: Optional[torch.nn.modules.module.Module] = None, final_activation: Union[torch.nn.modules.module.Module, str, NoneType] = None, postprocessing: Union[torch.nn.modules.module.Module, str, NoneType] = None, check_shape: bool = True)
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.

encoder
base
decoder
out_conv
check_shape
final_activation
postprocessing
in_channels
151    @property
152    def in_channels(self):
153        return self.encoder.in_channels
out_channels
155    @property
156    def out_channels(self):
157        return self._out_channels
depth
159    @property
160    def depth(self):
161        return len(self.encoder)
def load_encoder_state(self, state):
186    def load_encoder_state(self, state):
187        self.encoder.load_state_dict(state)
def load_decoder_state(self, state):
189    def load_decoder_state(self, state):
190        self.decoder.load_state_dict(state)
def load_base_state(self, state):
192    def load_base_state(self, state):
193        self.base.load_state_dict(state)
def forward(self, x: torch.Tensor) -> torch.Tensor:
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.

class UNet2d(UNetBase):
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.
UNet2d( in_channels: int, out_channels: int, depth: int = 4, initial_features: int = 32, gain: int = 2, final_activation=None, return_side_outputs: bool = False, conv_block_impl: torch.nn.modules.module.Module = <class 'torch_em.model.unet.ConvBlock2d'>, pooler_impl: torch.nn.modules.module.Module = <class 'torch.nn.modules.pooling.MaxPool2d'>, sampler_impl: torch.nn.modules.module.Module = <class 'torch_em.model.unet.Upsampler2d'>, postprocessing: Union[torch.nn.modules.module.Module, str, NoneType] = None, check_shape: bool = True, **conv_block_kwargs)
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.

init_kwargs
class AnisotropicUNet(UNetBase):
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.
AnisotropicUNet( in_channels: int, out_channels: int, scale_factors: List[List[int]], initial_features: int = 32, gain: int = 2, final_activation: Union[torch.nn.modules.module.Module, str, NoneType] = None, return_side_outputs: bool = False, conv_block_impl: torch.nn.modules.module.Module = <class 'torch_em.model.unet.ConvBlock3d'>, anisotropic_kernel: bool = False, postprocessing: Union[torch.nn.modules.module.Module, str, NoneType] = None, check_shape: bool = True, **conv_block_kwargs)
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.

init_kwargs
class UNet3d(AnisotropicUNet):
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.
UNet3d( in_channels: int, out_channels: int, depth: int = 4, initial_features: int = 32, gain: int = 2, final_activation: Union[torch.nn.modules.module.Module, str, NoneType] = None, return_side_outputs: bool = False, conv_block_impl: torch.nn.modules.module.Module = <class 'torch_em.model.unet.ConvBlock3d'>, postprocessing: Union[torch.nn.modules.module.Module, str, NoneType] = None, check_shape: bool = True, **conv_block_kwargs)
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.

init_kwargs