Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
119 changes: 119 additions & 0 deletions src/mrpro/nn/nets/MLP.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,119 @@
"""Multi-layer perceptron."""

from collections.abc import Sequence
from itertools import pairwise
from typing import Literal

import torch
from torch.nn import GELU, LeakyReLU, Linear, ReLU, SiLU

from mrpro.nn.FiLM import FiLM
from mrpro.nn.LayerNorm import LayerNorm
from mrpro.nn.Sequential import Sequential


class MLP(Sequential):
"""Multi-layer perceptron.

A series of linear layers, normalization and activation.
Allows FiLM conditioning.
Order is Linear -> Norm (optional) -> FiLM (optional) -> Activation.

If you need more flexibility, use `~mrpro.nn.Sequential` directly.
"""

features_last: bool

def __init__(
self,
n_channels_in: int,
n_channels_out: int,
norm: Literal['layer', 'none'] = 'none',
activation: Literal['gelu', 'relu', 'silu', 'leaky_relu'] = 'gelu',
n_features: Sequence[int] = (256, 256),
cond_dim: int = 0,
features_last: bool = True,
):
"""Initialize a MLP.

Parameters
----------
n_channels_in
The number of input channels.
n_channels_out
The number of output channels.
norm
The type of normalization to use. If `layer`, use layer normalization.
If `none`, use no normalization.
activation
The type of activation to use. If `gelu`, use GELU.
If `relu`, use ReLU. If `silu`, use SiLU. If `leaky_relu`, use LeakyReLU.
n_features
The number of features in the hidden layers. The length of this sequence determines the number of hidden
layers. The total number of linear layers is `len(n_features) + 1`.
cond_dim
The dimension of the condition tensor. If 0, no FiLM conditioning is applied.
Otherwise, between linear layers, after normalization, FiLM conditioning is applied.
features_last
Whether the features are in the last dimension, as common in transformer models,
or in the second dimension, as common in image models.
"""
super().__init__()
use_film = cond_dim > 0
self.features_last = features_last

if len(n_features) == 0:
self.append(Linear(n_channels_in, n_channels_out))
return

self.append(Linear(n_channels_in, n_features[0]))

for c_in, c_out in pairwise((*n_features, n_channels_out)):
if norm.lower() == 'layer':
self.append(LayerNorm(c_in, features_last=True))
elif norm.lower() != 'none':
raise ValueError(f'Invalid normalization type: {norm}')

if use_film:
self.append(FiLM(c_in, cond_dim, features_last=True))

if activation.lower() == 'gelu':
self.append(GELU(approximate='tanh'))
elif activation.lower() == 'relu':
self.append(ReLU())
elif activation.lower() == 'silu':
self.append(SiLU())
elif activation.lower() == 'leaky_relu':
self.append(LeakyReLU())
else:
raise ValueError(f'Invalid activation type: {activation}')

self.append(Linear(c_in, c_out))

def __call__(self, x: torch.Tensor, *, cond: torch.Tensor | None = None) -> torch.Tensor: # type: ignore[override]
"""Apply the MLP to the input tensor.

Parameters
----------
x
The input tensor.
cond
The condition tensor. If None, no FiLM conditioning is applied.

Returns
-------
The output tensor.
"""
return super().__call__(x, cond=cond)

def forward(self, *x: torch.Tensor, cond: torch.Tensor | None = None) -> torch.Tensor:
"""Apply the MLP to the input tensor."""
if len(x) != 1:
raise ValueError(f'Mlp expects exactly one input tensor, got {len(x)}')
tensor = x[0]
if not self.features_last:
tensor = tensor.moveaxis(1, -1)
out = super().forward(tensor, cond=cond)
if not self.features_last:
out = out.moveaxis(-1, 1)
return out
8 changes: 5 additions & 3 deletions src/mrpro/nn/nets/__init__.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
from mrpro.nn.nets.BasicCNN import BasicCNN
from mrpro.nn.nets.UNet import AttentionGatedUNet, UNet
from mrpro.nn.nets.MLP import MLP

__all__ = [
'AttentionGatedUNet',
'BasicCNN',
'UNet',
"AttentionGatedUNet",
"BasicCNN",
"MLP",
"UNet",
]
89 changes: 89 additions & 0 deletions tests/nn/test_mlp.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,89 @@
"""Tests for Mlp module."""

from typing import cast

import pytest
import torch
from mrpro.nn.nets import MLP
from mrpro.utils import RandomGenerator


@pytest.mark.parametrize('torch_compile', [True, False], ids=['compiled', 'uncompiled'])
@pytest.mark.parametrize(
'device',
[
pytest.param('cpu', id='cpu'),
pytest.param('cuda', marks=pytest.mark.cuda, id='cuda'),
],
)
def test_mlp_forward(torch_compile: bool, device: str) -> None:
"""Test the forward pass of the Mlp."""
mlp = MLP(
n_channels_in=8,
n_channels_out=4,
norm='layer',
activation='gelu',
n_features=(16,),
cond_dim=12,
features_last=False,
).to(device)
x = torch.zeros(1, 8, 9, 7, device=device)
cond = torch.zeros(1, 12, device=device)
if torch_compile:
mlp = cast(MLP, torch.compile(mlp))
y = mlp(x, cond=cond)
assert y.shape == (1, 4, 9, 7)


def test_mlp_backward() -> None:
"""Test the backward pass of the Mlp."""
mlp = MLP(
n_channels_in=6,
n_channels_out=3,
norm='none',
activation='silu',
n_features=(12, 12),
cond_dim=10,
features_last=True,
)
rng = RandomGenerator(seed=42)
x = rng.float32_tensor((1, 20, 6)).requires_grad_(True)
cond = rng.float32_tensor((1, 10)).requires_grad_(True)
y = mlp(x, cond=cond)
y.sum().backward()
assert x.grad is not None, 'x.grad is None'
assert not x.grad.isnan().any(), 'x.grad is NaN'
assert cond.grad is not None, 'cond.grad is None'
assert not cond.grad.isnan().any(), 'cond.grad is NaN'
for name, parameter in mlp.named_parameters():
assert parameter.grad is not None, f'{name}.grad is None'
assert not parameter.grad.isnan().any(), f'{name}.grad is NaN'


def test_mlp_features_last() -> None:
"""Test Mlp with features_last=True vs features_last=False."""
rng = RandomGenerator(seed=42)
x = rng.float32_tensor((1, 3, 4, 5)).requires_grad_(True)

mlp_last = MLP(
n_channels_in=3,
n_channels_out=4,
norm='layer',
activation='relu',
n_features=(6,),
cond_dim=0,
features_last=True,
)
mlp = MLP(
n_channels_in=3,
n_channels_out=4,
norm='layer',
activation='relu',
n_features=(6,),
cond_dim=0,
features_last=False,
)
mlp.load_state_dict(mlp_last.state_dict())
y_last = mlp_last(x.moveaxis(1, -1))
y = mlp(x)
torch.testing.assert_close(y, y_last.moveaxis(-1, 1))
Loading