Skip to content
36 changes: 34 additions & 2 deletions monai/networks/nets/swin_unetr.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,7 @@ def __init__(
hyena_omega_0: float = 10.0,
hyena_l_cache: int = 32,
hyena_short_conv_fft_chunks: int = 0,
use_flash_attention: bool = False,
) -> None:
"""
Args:
Expand Down Expand Up @@ -149,6 +150,7 @@ def __init__(
hyena_omega_0: SIREN frequency. Default 10.0 (stable).
hyena_l_cache: SIREN coordinate-grid cache size per spatial dim.
hyena_short_conv_fft_chunks: channel chunk size for the FFT short conv (0 = no chunking).
use_flash_attention: use flash attention (scaled dot product attention) at inference.

Examples::

Expand Down Expand Up @@ -224,6 +226,7 @@ def __init__(
hyena_omega_0=hyena_omega_0,
hyena_l_cache=hyena_l_cache,
hyena_short_conv_fft_chunks=hyena_short_conv_fft_chunks,
use_flash_attention=use_flash_attention,
)

self.encoder1 = UnetrBasicBlock(
Expand Down Expand Up @@ -513,6 +516,7 @@ def __init__(
qkv_bias: bool = False,
attn_drop: float = 0.0,
proj_drop: float = 0.0,
use_flash_attention: bool = False,
) -> None:
"""
Args:
Expand All @@ -522,12 +526,17 @@ def __init__(
qkv_bias: add a learnable bias to query, key, value.
attn_drop: attention dropout rate.
proj_drop: dropout rate of output.
use_flash_attention: if True, use ``torch.nn.functional.scaled_dot_product_attention`` for the
windowed attention. Equivalent to the default path but faster at inference; only used when
autograd is disabled (e.g. under ``torch.no_grad()`` or ``torch.inference_mode()``, not
``eval()`` alone) and the module is not scripted.
"""

super().__init__()
self.dim = dim
self.window_size = window_size
self.num_heads = num_heads
self.use_flash_attention = use_flash_attention
head_dim = dim // num_heads
self.scale = head_dim**-0.5
mesh_args = torch.meshgrid.__kwdefaults__
Expand Down Expand Up @@ -584,12 +593,26 @@ def forward(self, x, mask):
b, n, c = x.shape
qkv = self.qkv(x).reshape(b, n, 3, self.num_heads, c // self.num_heads).permute(2, 0, 3, 1, 4)
q, k, v = qkv[0], qkv[1], qkv[2]
q = q * self.scale
attn = q @ k.transpose(-2, -1)
relative_position_bias = self.relative_position_bias_table[
self.relative_position_index.clone()[:n, :n].reshape(-1) # type: ignore[operator]
].reshape(n, n, -1)
relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous()
if self.use_flash_attention and not torch.jit.is_scripting() and not torch.is_grad_enabled():
Comment thread
coderabbitai[bot] marked this conversation as resolved.
# additive bias combines the relative position bias and, for shifted windows, the attention mask
if mask is not None:
nw = mask.shape[0]
bias = relative_position_bias.view(1, 1, self.num_heads, n, n) + mask.reshape(1, nw, 1, n, n)
bias = bias.expand(b // nw, nw, self.num_heads, n, n).reshape(b, self.num_heads, n, n)
else:
bias = relative_position_bias.unsqueeze(0)
x = torch.nn.functional.scaled_dot_product_attention(
q, k, v, attn_mask=bias.to(q.dtype), dropout_p=0.0, scale=self.scale
)
x = x.transpose(1, 2).reshape(b, n, c)
return self.proj_drop(self.proj(x))

q = q * self.scale
attn = q @ k.transpose(-2, -1)
attn = attn + relative_position_bias.unsqueeze(0)
if mask is not None:
nw = mask.shape[0]
Expand Down Expand Up @@ -628,6 +651,7 @@ def __init__(
act_layer: str = "GELU",
norm_layer: type[LayerNorm] = nn.LayerNorm,
use_checkpoint: bool = False,
use_flash_attention: bool = False,
) -> None:
"""
Args:
Expand All @@ -643,6 +667,7 @@ def __init__(
act_layer: activation layer.
norm_layer: normalization layer.
use_checkpoint: use gradient checkpointing for reduced memory usage.
use_flash_attention: use flash attention (scaled dot product attention) at inference.
"""

super().__init__()
Expand All @@ -660,6 +685,7 @@ def __init__(
qkv_bias=qkv_bias,
attn_drop=attn_drop,
proj_drop=drop,
use_flash_attention=use_flash_attention,
)

self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
Expand Down Expand Up @@ -924,6 +950,7 @@ def __init__(
hyena_omega_0: float = 10.0,
hyena_l_cache: int = 32,
hyena_short_conv_fft_chunks: int = 0,
use_flash_attention: bool = False,
) -> None:
"""
Args:
Expand All @@ -946,6 +973,7 @@ def __init__(
hyena_use_chunked_fft, hyena_use_fft_short_conv, hyena_omega_0, hyena_l_cache,
hyena_short_conv_fft_chunks: forwarded to :class:`HyenaTransformerBlock`. See its
docstring for semantics.
use_flash_attention: use flash attention (scaled dot product attention) at inference.
"""

super().__init__()
Expand Down Expand Up @@ -996,6 +1024,7 @@ def __init__(
drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,
norm_layer=norm_layer,
use_checkpoint=use_checkpoint,
use_flash_attention=use_flash_attention,
)
for i in range(depth)
]
Expand Down Expand Up @@ -1079,6 +1108,7 @@ def __init__(
hyena_omega_0: float = 10.0,
hyena_l_cache: int = 32,
hyena_short_conv_fft_chunks: int = 0,
use_flash_attention: bool = False,
) -> None:
"""
Args:
Expand Down Expand Up @@ -1112,6 +1142,7 @@ def __init__(
hyena_use_chunked_fft, hyena_use_fft_short_conv, hyena_omega_0, hyena_l_cache,
hyena_short_conv_fft_chunks: HyenaND configuration. See
:class:`monai.networks.blocks.HyenaTransformerBlock` for semantics.
use_flash_attention: use flash attention (scaled dot product attention) at inference.
"""

super().__init__()
Expand Down Expand Up @@ -1193,6 +1224,7 @@ def __init__(
hyena_omega_0=hyena_omega_0,
hyena_l_cache=hyena_l_cache,
hyena_short_conv_fft_chunks=hyena_short_conv_fft_chunks,
use_flash_attention=use_flash_attention,
)
if i_layer == 0:
self.layers1.append(layer)
Expand Down
13 changes: 13 additions & 0 deletions tests/networks/nets/test_swin_unetr.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,19 @@ def test_invalid_input_shape(self):
with self.assertRaises(ValueError):
net_2d(torch.randn(1, 1, 48, 33)) # 33 is not divisible by 32

@skipUnless(has_einops, "Requires einops")
def test_flash_attention(self):
input_param = {"in_channels": 1, "out_channels": 2, "feature_size": 12, "spatial_dims": 3}
net_ref = SwinUNETR(use_flash_attention=False, **input_param).double()
net_flash = SwinUNETR(use_flash_attention=True, **input_param).double()
net_flash.load_state_dict(net_ref.state_dict())
x = torch.randn(1, 1, 64, 64, 64, dtype=torch.float64)
with eval_mode(net_ref, net_flash):
ref = net_ref.swinViT(x, net_ref.normalize)
out = net_flash.swinViT(x, net_flash.normalize)
for a, b in zip(ref, out, strict=True):
assert_allclose(a, b, atol=1e-6, rtol=1e-6, type_test=False)

def test_patch_merging(self):
dim = 10
t = PatchMerging(dim)(torch.zeros((1, 21, 20, 20, dim)))
Expand Down
Loading