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,BaseEstimatorImplementation 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, andforward(x, masks=None). SeeVisionTransformer3DMoEfor a reference implementation (with or without a sparse MoE backbone – passuse_moe=Trueto that constructor and setuse_moe=Truehere 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
- __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]¶
- 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.
- training_step(batch, batch_idx)[source]¶
Perform one training step and compute the training loss.
- Parameters:
- batchtorch.Tensor or (torch.Tensor, torch.Tensor)
Xor(X, Y)whereXhas 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).