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.estimators.ssl.NeuroJEPA

class nidl.estimators.ssl.NeuroJEPA(encoder, mask_scale_configs=(MaskScaleConfig(spatial_scale=(0.0, 0.2), depth_scale=(0.0, 1.0), aspect_ratio=(0.75, 1.5), num_blocks=32, total_mask_ratio=0.75), MaskScaleConfig(spatial_scale=(0.2, 0.5), depth_scale=(0.0, 1.0), aspect_ratio=(0.75, 1.5), num_blocks=16, total_mask_ratio=0.75), MaskScaleConfig(spatial_scale=(0.5, 0.7), depth_scale=(0.0, 1.0), aspect_ratio=(0.75, 1.5), num_blocks=4, total_mask_ratio=0.75)), foreground_aware=True, foreground_threshold=0.0, min_foreground_fraction=0.1, loss_exp=1.0, bg_weight=0.1, use_moe=False, moe_bias_update_rate=0.0001, moe_bias_clip=0.3, predictor_embed_dim=384, predictor_depth=6, predictor_num_heads=12, optimizer='adamW', learning_rate=0.0006, weight_decay=0.04, exclude_bias_and_norm_wd=True, ema_start=0.99925, ema_end=1.0, optimizer_kwargs=None, lr_scheduler='warmup_cosine', lr_scheduler_kwargs=None, **kwargs)[source]

Bases: TransformerMixin, BaseEstimator

Implementation of Neuro-JEPA [1].

This solver predicts the representations of missing parts of the input (here: 3D brain MRI patches) from a visible context, using a context encoder, an EMA target encoder, and a predictor.

Compared to I-JEPA (3d), it uses:

  • multi-scale block masking (several differently-shaped-but-equal-ratio maskings per volume instead of one),

  • a sparse Mixture-of-Experts backbone (via the encoder you pass in),

  • a foreground-aware loss that down-weights background (non-brain) voxel-patches.

Parameters:
encodernn.Module

3D ViT-like encoder. Must expose embed_dim, patch_size, grid_shape, blocks, and forward(x, masks=None). See VisionTransformer3DMoE for a reference implementation (with or without a sparse MoE backbone – pass use_moe=True to that constructor and set use_moe=True here too so the MoE bias update runs during training).

mask_scale_configssequence of MaskScaleConfig, default=3-scale config from [1]

One entry per masking “scale”: each draws num_blocks blocks with sizes controlled by spatial_scale/depth_scale/aspect_ratio, unions them, and adjusts to hit exactly total_mask_ratio of all patches. By default, a small-block (32 blocks, 0-20% spatial scale), a medium-block (16 blocks, 20-50% spatial scale), and a large-block (4 blocks, 50-70% spatial scale) masking are all drawn for every volume, at every step, each targeting a 75% mask ratio.

foreground_awarebool, default=True

Whether to compute a per-patch foreground map (voxel-intensity based) and use it to (a) bias mask erosion/dilation toward removing background patches first, and (b) down-weight background patches in the loss (see bg_weight).

foreground_thresholdfloat, default=0.0

Fallback per-sample voxel-intensity threshold used by compute_foreground_patches when the sample’s data-driven (2nd/98th percentile based) threshold cannot be estimated.

min_foreground_fractionfloat, default=0.1

Minimum fraction of foreground voxels a patch must contain to be itself counted as a foreground patch (see compute_foreground_patches).

loss_expfloat, default=1.0

Exponent of the per-token L1-style loss (|pred - target|^p / p).

bg_weightfloat, default=0.1

Loss weight for a pure-background patch; a pure-foreground patch always has weight 1.0. Set to 1.0 to disable foreground weighting (recovers a uniform loss, i.e. IJEPA’s SmoothL1Loss in spirit).

use_moebool, default=False

Whether encoder.blocks contains sparse MoE layers that need the post-step, aux-loss-free bias update (see vision_transformer_3d.moe_bias_update). Set to match however you built encoder.

moe_bias_update_ratefloat, default=1e-4

Step size of the MoE router bias update (ignored if use_moe=False).

moe_bias_clipfloat, default=0.3

Maximum absolute value of the MoE router bias after each update (ignored if use_moe=False).

predictor_embed_dimint, default=384

Dimension of the predictor hidden layers. It can be different from the encoder output dimension.

predictor_depthint, default=6

Number of Transformer blocks in the predictor.

predictor_num_headsint, default=12

Number of attention heads in the predictor.

optimizer{‘sgd’, ‘adam’, ‘adamW’} or Optimizer, default=’adamW’

Optimizer for training the model. If a string is given, it can be:

  • ‘sgd’: Stochastic Gradient Descent (with optional momentum).

  • ‘adam’: First-order gradient-based optimizer.

  • ‘adamW’ (default): Adam with decoupled weight decay regularization (see “Decoupled Weight Decay Regularization”, Loshchilov and Hutter, ICLR 2019).

learning_ratefloat, default=6e-4

Initial learning rate.

weight_decayfloat, default=0.04

Weight decay in the optimizer.

exclude_bias_and_norm_wdbool, default=True

Whether the bias terms and normalization layers get weight decay during optimization or not.

ema_startfloat, default=0.99925

Base value for the weighting coefficient in the target encoder momentum update with exponential moving average. A cosine annealing scheme is used.

ema_endfloat, default=1.0

Final value for the weighting coefficient in the target encoder momentum update.

optimizer_kwargsdict or None, default=None

Extra named arguments for the optimizer.

lr_scheduler{“none”, “warmup_cosine”}, LRSchedulerPLType or None, default=”warmup_cosine”

Learning rate scheduler to use.

lr_scheduler_kwargsdict or None, default=None

Extra named arguments for the scheduler. By default, it is set to {“warmup_epochs”: 10, “warmup_start_lr”: 1e-6, “min_lr”: 0.0, “interval”: “step”}.

**kwargsdict, optional

Extra named arguments for the BaseEstimator class (given to the PL Trainer), such as max_epochs, max_steps, callbacks, etc.

Attributes:
context_encoderNeuroJEPAEncoderWrapper

Wraps the trainable copy of encoder. Used at inference (transform).

target_encoderNeuroJEPAEncoderWrapper

Wraps the EMA copy of encoder; not used at inference.

predictorVisionTransformerPredictor3D

Predicts target-encoder latents from context-encoder latents.

maskerMultiScaleMaskCollator

Generates the per-step multi-scale (context, target) mask pairs.

References

[1] (1,2)

Huang et al., “Learning Sparse Latent Predictive Foundation Model for Multimodal Neuroimaging”, arXiv:2606.14957.

__init__(encoder, mask_scale_configs=(MaskScaleConfig(spatial_scale=(0.0, 0.2), depth_scale=(0.0, 1.0), aspect_ratio=(0.75, 1.5), num_blocks=32, total_mask_ratio=0.75), MaskScaleConfig(spatial_scale=(0.2, 0.5), depth_scale=(0.0, 1.0), aspect_ratio=(0.75, 1.5), num_blocks=16, total_mask_ratio=0.75), MaskScaleConfig(spatial_scale=(0.5, 0.7), depth_scale=(0.0, 1.0), aspect_ratio=(0.75, 1.5), num_blocks=4, total_mask_ratio=0.75)), foreground_aware=True, foreground_threshold=0.0, min_foreground_fraction=0.1, loss_exp=1.0, bg_weight=0.1, use_moe=False, moe_bias_update_rate=0.0001, moe_bias_clip=0.3, predictor_embed_dim=384, predictor_depth=6, predictor_num_heads=12, optimizer='adamW', learning_rate=0.0006, weight_decay=0.04, exclude_bias_and_norm_wd=True, ema_start=0.99925, ema_end=1.0, optimizer_kwargs=None, lr_scheduler='warmup_cosine', lr_scheduler_kwargs=None, **kwargs)[source]
configure_optimizers()[source]

Initialize the optimizer and learning rate scheduler.

forward_target(x, masks_pred)[source]

Full (unmasked) forward through the EMA target encoder, layer-normed, then sliced into per-scale target latents at the masked positions. Mirrors IJEPA.forward_target.

on_fit_start()[source]

Decorrelate mask geometry across DDP ranks.

self.trainer.global_rank only exists once a Trainer/strategy is attached (i.e. not yet at __init__ time, when the estimator is merely constructed) – on_fit_start is the first hook guaranteed to run after that setup and before the first training step, so it’s the right place to set it. Without this, every rank’s self.masker starts its own step counter at 0 and increments in lockstep with every other rank, with nothing else rank-dependent in the seed – so all ranks would draw the identical mask geometry for the sample at a given local batch index every step, even though the actual volume there differs per rank. Single-GPU / single-process runs are unaffected (global_rank == 0, offset is a no-op).

on_train_batch_end(outputs, batch, batch_idx)[source]

Performs the teacher momentum update (and, if use_moe=True, the MoE aux-loss-free bias update), matching the official train_one_epoch’s post-step ordering: MoE bias update, then EMA.

Parameters:
outputsdict

Outputs of the training step (ignored).

batchtorch.Tensor or pair of torch.Tensor

Ignored.

batch_idxint

Ignored.

test_step(batch, batch_idx)[source]

Skip the test step.

training_step(batch, batch_idx)[source]

Perform one training step and compute the training loss.

Parameters:
batchtorch.Tensor or (torch.Tensor, torch.Tensor)

X or (X, Y) where X has shape (B, C, H, W, D). Y (eventual labels) is ignored.

batch_idxint

Ignored.

Returns:
outputsdict

"loss", "z_pred" (list[num_scales] of predictor outputs), "z_target" (list[num_scales] of target-encoder outputs).

transform_step(batch, batch_idx, dataloader_idx=0)[source]

Encode the input data into the latent space (no masking, no predictor – the predictor is only used during training).

Parameters:
batchtorch.Tensor

(B, C, H, W, D) volume batch, given as-is to the context encoder.

batch_idx, dataloader_idxint

Ignored.

Returns:
featurestorch.Tensor

Context-encoder features averaged across the token dimension, shape (B, embed_dim).

validation_step(batch, batch_idx)[source]

Performs one validation step and computes the validation loss. Same return structure as training_step.