Note

This page is a reference documentation. It only explains the class signature, and not how to use it. Please refer to the user guide for the big picture.

nidl.backbones.volume.VisionTransformer3DMoE

class nidl.backbones.volume.VisionTransformer3DMoE(img_size=(96, 108, 96), patch_size=(12, 12, 12), in_chans=1, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4.0, drop_path_rate=0.0, use_moe=False, moe_params=None, init_std=0.02)[source]

Bases: Module

3D ViT backbone with an optional sparse Mixture-of-Experts (MoE) mixed in at configurable layers.

Parameters:
img_size(int, int, int), default=(96, 108, 96)

Size (in voxels) of the input volume.

patch_size(int, int, int), default=(12, 12, 12)

Size (in voxels) of one cubic patch (“tubelet”).

in_chansint, default=1

Number of input channels.

embed_dimint, default=768

Token embedding dimension.

depthint, default=12

Number of Transformer blocks.

num_headsint, default=12

Number of attention heads.

mlp_ratiofloat, default=4.0

Ratio between the MLP hidden dimension and embed_dim (dense blocks only; MoE blocks use moe_params.moe_inter_dim instead).

drop_path_ratefloat, default=0.0

Maximum stochastic-depth drop rate, linearly increased across blocks.

use_moebool, default=False

Whether to replace the dense MLP with a sparse MoE (see MoE) in the blocks listed in moe_params.moe_layer_indices.

moe_paramsMoEParams or None, default=None

Sparse MoE hyperparameters. Required if use_moe=True.

init_stdfloat, default=0.02

Standard deviation used for truncated-normal weight init.

__init__(img_size=(96, 108, 96), patch_size=(12, 12, 12), in_chans=1, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4.0, drop_path_rate=0.0, use_moe=False, moe_params=None, init_std=0.02)[source]

Initialize internal Module state, shared by both nn.Module and ScriptModule.

forward(x, masks=None)[source]

Encode a volume into patch tokens.

Parameters:
xtorch.Tensor

(B, C, H, W, D) volume.

maskslist of torch.Tensor or None, default=None

Optional list of (B, K) LongTensors: if given, only those patch indices are kept, and the output batch is multiplied by len(masks) (one block per mask, concatenated along batch).

Returns:
tokenstorch.Tensor

(B[*len(masks)], K, E) patch tokens.

moe_scoreslist

One entry per block of router scores, or an empty list if use_moe=False.

property grid_shape

(nH, nW, nD) patch-grid dimensions.