Skip to content
30 changes: 30 additions & 0 deletions docs/source/networks.rst
Original file line number Diff line number Diff line change
Expand Up @@ -348,6 +348,11 @@ Layers

.. currentmodule:: monai.networks.layers

`LayerNormNd`
~~~~~~~~~~~~~
.. autoclass:: LayerNormNd
:members:

`ChannelPad`
~~~~~~~~~~~~
.. autoclass:: ChannelPad
Expand Down Expand Up @@ -472,6 +477,31 @@ Nets
.. autoclass:: AHNet
:members:

`ConvNeXt`
~~~~~~~~~~
.. autoclass:: ConvNeXt
:members:

`ConvNeXtTiny`
~~~~~~~~~~~~~~
.. autoclass:: ConvNeXtTiny

`ConvNeXtSmall`
~~~~~~~~~~~~~~~
.. autoclass:: ConvNeXtSmall

`ConvNeXtBase`
~~~~~~~~~~~~~~
.. autoclass:: ConvNeXtBase

`ConvNeXtLarge`
~~~~~~~~~~~~~~~
.. autoclass:: ConvNeXtLarge

`ConvNeXtXLarge`
~~~~~~~~~~~~~~~~
.. autoclass:: ConvNeXtXLarge

`DenseNet`
~~~~~~~~~~
.. autoclass:: DenseNet
Expand Down
1 change: 1 addition & 0 deletions monai/networks/layers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from .factories import Act, Conv, Dropout, LayerFactory, Norm, Pad, Pool, RelPosEmbedding, split_args
from .filtering import BilateralFilter, PHLFilter, TrainableBilateralFilter, TrainableJointBilateralFilter
from .gmm import GaussianMixtureModel
from .layer_norm_nd import LayerNormNd
from .simplelayers import (
LLTM,
ApplyFilter,
Expand Down
58 changes: 58 additions & 0 deletions monai/networks/layers/layer_norm_nd.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
# Copyright (c) MONAI Consortium
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
# http://www.apache.org/licenses/LICENSE-2.0
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

from __future__ import annotations

import torch
import torch.nn as nn


class LayerNormNd(nn.Module):
"""
Layer normalization over the channel dimension of a channels-first tensor.

`torch.nn.LayerNorm` normalizes over the trailing dimensions, so it expects a channels-last layout
such as ``(batch, *spatial, channel)``. Convolutional feature maps in MONAI are channels-first,
``(batch, channel, *spatial)``, and this module normalizes those over the channel dimension only,
for any number of spatial dimensions.

Args:
num_channels: number of channels of the input, i.e. the size of dimension 1.
spatial_dims: number of spatial dimensions of the input image.
eps: value added to the denominator for numerical stability.

Raises:
ValueError: if ``num_channels`` is not positive, or ``spatial_dims`` is negative.
"""

def __init__(self, num_channels: int, spatial_dims: int, eps: float = 1e-6) -> None:
super().__init__()
Comment thread
ericspod marked this conversation as resolved.
if num_channels <= 0:
raise ValueError(f"num_channels must be positive, got {num_channels}.")
if spatial_dims < 0:
raise ValueError(f"spatial_dims must be non-negative, got {spatial_dims}.")
self.eps = eps
self.weight = nn.Parameter(torch.ones(num_channels))
self.bias = nn.Parameter(torch.zeros(num_channels))
# broadcast the affine parameters against (batch, channel, *spatial); precomputed so that the
# module is scriptable without inspecting the rank of the input at runtime.
self.param_shape = [1, num_channels] + [1] * spatial_dims
Comment thread
ericspod marked this conversation as resolved.

def forward(self, x: torch.Tensor) -> torch.Tensor:
mean = x.mean(dim=1, keepdim=True)
xmean = x - mean
var = xmean.pow(2).mean(dim=1, keepdim=True)
# `var` and the sqrt result are intermediates used nowhere else, so mutating them in place is
# safe; `xmean` must stay untouched since its pre-division value is what `pow(2)`'s backward
# needs, so the final division is deliberately out-of-place.
denom = var.add_(self.eps).sqrt_()
x = xmean / denom
return self.weight.view(self.param_shape).mul(x).add_(self.bias.view(self.param_shape))
19 changes: 19 additions & 0 deletions monai/networks/nets/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,25 @@
from .basic_unetplusplus import BasicUNetPlusPlus, BasicUnetPlusPlus, BasicunetPlusPlus, basicunetplusplus
from .classifier import Classifier, Critic, Discriminator
from .controlnet import ControlNet
from .convnext import (
ConvNeXt,
Convnext,
Convnext_base,
Convnext_large,
Convnext_small,
Convnext_tiny,
Convnext_xlarge,
ConvNeXtBase,
ConvNeXtLarge,
ConvNeXtSmall,
ConvNeXtTiny,
ConvNeXtXLarge,
convnext_base,
convnext_large,
convnext_small,
convnext_tiny,
convnext_xlarge,
)
from .daf3d import DAF3D
from .densenet import (
DenseNet,
Expand Down
Loading
Loading