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:
Module3D 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 bylen(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.