Skip to content

Neural Nets API

Core network building blocks for embedding, model selection, and inference. Start with dmri_reconstruction_model for the assembled model; use the remaining modules when customizing an architecture.

For a small in-process example, see standalone training. Use dmri train for normal checkpointed training runs.

Reconstruction and selection

dmri_reconstruction_model

EmbeddingModule module-attribute

TokenizerType module-attribute

TokenizerType = type[DMRITokenizer]

AcquisitionSchemeLike module-attribute

AcquisitionSchemeLike = acquisition_scheme | ssfp_acquisition_scheme

DMRIInferenceModelConfig dataclass

DMRIInferenceModelConfig(simulator: type[MultiCompartment], model_dim: int = 64, use_attention_mask: bool = True, inference_loss_type: str = 'v', dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, tokenizer_cls: TokenizerType = DMRITokenizer, embedding_cls: type[EmbeddingModule] = BvalBvecSignalEmbeddingNet, embedding_cfg: Any = DMRIEmbeddingConfig(), model_selection_cfg: Any = DMRIModelSelectionConfig(), theta_inference_cfg: DMRIThetaInferenceConfig = DMRIThetaInferenceConfig())

simulator instance-attribute

simulator: type[MultiCompartment]

model_dim class-attribute instance-attribute

model_dim: int = 64

use_attention_mask class-attribute instance-attribute

use_attention_mask: bool = True

inference_loss_type class-attribute instance-attribute

inference_loss_type: str = 'v'

dtype class-attribute instance-attribute

dtype: DTypeLike | None = None

param_dtype class-attribute instance-attribute

param_dtype: DTypeLike | None = None

precision class-attribute instance-attribute

precision: PrecisionLike | None = None

preferred_element_type class-attribute instance-attribute

preferred_element_type: DTypeLike | None = None

tokenizer_cls class-attribute instance-attribute

tokenizer_cls: TokenizerType = DMRITokenizer

embedding_cls class-attribute instance-attribute

embedding_cfg class-attribute instance-attribute

embedding_cfg: Any = field(default_factory=DMRIEmbeddingConfig)

model_selection_cfg class-attribute instance-attribute

model_selection_cfg: Any = field(default_factory=DMRIModelSelectionConfig)

theta_inference_cfg class-attribute instance-attribute

theta_inference_cfg: DMRIThetaInferenceConfig = field(default_factory=DMRIThetaInferenceConfig)

DMRIInferenceModelConfigMaskPriorAmortized dataclass

DMRIInferenceModelConfigMaskPriorAmortized(simulator: type[MultiCompartment], model_dim: int = 64, use_attention_mask: bool = True, inference_loss_type: str = 'v', dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, tokenizer_cls: TokenizerType = DMRITokenizer, embedding_cls: type[EmbeddingModule] = BvalBvecSignalEmbeddingNet, embedding_cfg: Any = DMRIEmbeddingConfig(), model_selection_cfg: Any = DMRIModelSelectionAmortizedPriorConfig(), theta_inference_cfg: DMRIThetaInferenceConfig = DMRIThetaInferenceConfig())

Bases: DMRIInferenceModelConfig

model_selection_cfg class-attribute instance-attribute

model_selection_cfg: Any = field(default_factory=DMRIModelSelectionAmortizedPriorConfig)

DMRIInferenceModelConfigMaskPriorAmortizedPP dataclass

DMRIInferenceModelConfigMaskPriorAmortizedPP(simulator: type[MultiCompartment], model_dim: int = 64, use_attention_mask: bool = True, inference_loss_type: str = 'v', dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, tokenizer_cls: TokenizerType = DMRITokenizerPP, embedding_cls: type[EmbeddingModule] = BvalBvecSignalEmbeddingNet, embedding_cfg: Any = DMRIEmbeddingConfig(), model_selection_cfg: Any = DMRIModelSelectionAmortizedPriorConfig(), theta_inference_cfg: DMRIThetaInferenceConfig = DMRIThetaInferenceConfig())

Bases: DMRIInferenceModelConfigMaskPriorAmortized

tokenizer_cls class-attribute instance-attribute

tokenizer_cls: TokenizerType = DMRITokenizerPP

embedding_cls class-attribute instance-attribute

embedding_cfg class-attribute instance-attribute

embedding_cfg: Any = field(default_factory=DMRIEmbeddingConfig)

DMRIInferenceModelConfigMaskPriorAmortizedPPP dataclass

DMRIInferenceModelConfigMaskPriorAmortizedPPP(simulator: type[MultiCompartment], model_dim: int = 64, use_attention_mask: bool = True, inference_loss_type: str = 'v', dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, tokenizer_cls: TokenizerType = DMRITokenizerPPP, embedding_cls: type[EmbeddingModule] = BvalBvecSignalEmbeddingNet, embedding_cfg: Any = DMRIEmbeddingConfig(), model_selection_cfg: Any = DMRIModelSelectionAmortizedPriorConfig(), theta_inference_cfg: DMRIThetaInferenceConfig = DMRIThetaInferenceConfig())

Bases: DMRIInferenceModelConfigMaskPriorAmortized

Legacy architecture used by early multishell pretrained checkpoints.

tokenizer_cls class-attribute instance-attribute

embedding_cls class-attribute instance-attribute

embedding_cfg class-attribute instance-attribute

embedding_cfg: Any = field(default_factory=DMRIEmbeddingConfig)

SSFPInferenceModelConfig dataclass

SSFPInferenceModelConfig(simulator: type[MultiCompartment], model_dim: int = 64, use_attention_mask: bool = True, inference_loss_type: str = 'v', dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, tokenizer_cls: TokenizerType = DMRITokenizerPP, embedding_cls: type[EmbeddingModule] = SSFPEmbeddingNet, embedding_cfg: Any = SSFPEmbeddingNetConfig(), model_selection_cfg: Any = DMRIModelSelectionConfig(), theta_inference_cfg: DMRIThetaInferenceConfig = DMRIThetaInferenceConfig())

Bases: DMRIInferenceModelConfig

tokenizer_cls class-attribute instance-attribute

tokenizer_cls: TokenizerType = DMRITokenizerPP

embedding_cls class-attribute instance-attribute

embedding_cls: type[EmbeddingModule] = SSFPEmbeddingNet

embedding_cfg class-attribute instance-attribute

embedding_cfg: Any = field(default_factory=SSFPEmbeddingNetConfig)

DMRIInferenceModel

DMRIInferenceModel(cfg: DMRIInferenceModelConfig, rngs: Rngs)

Bases: Module

cfg instance-attribute

precision_fields instance-attribute

precision_fields: tuple[str, ...] = ('dtype', 'param_dtype', 'precision', 'preferred_element_type')

__call__

__call__(model_mask: Array, theta: Array, x: Array, acq: AcquisitionSchemeLike, mask_prior: Array | None = None, alpha_prior: Array | None = None, model_idx: list[int] | None = None, noise_idx: list[int] | None = None, t: ArrayLike | None = None) -> tuple[Array, Array]

set_precision

set_precision(*, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, use_flash_attention: bool | None = None, use_flash_cross_attention: bool | None = None) -> None

Reinitialize the model with updated precision/attention settings.

Parameters:

Name Type Description Default
dtype DTypeLike | None

Optional override for computation dtype.

None
param_dtype DTypeLike | None

Optional override for parameter dtype.

None
precision PrecisionLike | None

Optional override for XLA matmul precision.

None
preferred_element_type DTypeLike | None

Optional override for matmul element type.

None
use_flash_attention bool | None

Optional override applied to every sub-config that exposes a use_flash_attention flag.

None
use_flash_cross_attention bool | None

Optional override applied to every sub-config exposing use_flash_cross_attention.

None
Note

Calling this method discards the current parameter values because all modules are rebuilt from scratch.

embed_inputs

embed_inputs(model_mask: Array, x: Array, acq: AcquisitionSchemeLike, mask_prior: Array | None = None, alpha_prior: Array | None = None, model_idx: list[int] | None = None, noise_idx: list[int] | None = None) -> tuple[Array, Array | None, Array, Array | None]

theta_mask

theta_mask(model_mask: Array, model_idx: list[int] | None = None, noise_idx: list[int] | None = None) -> Array

marginalization_mask

marginalization_mask(model_mask: Array, model_idx: list[int] | None = None, noise_idx: list[int] | None = None) -> Array

loss_fn

loss_fn(rng: RngKey, model_mask: Array, theta: Array, x: Array, acq: AcquisitionSchemeLike, mask_prior: Array | None = None, alpha_prior: Array | None = None, target_score: Array | None = None, model_idx: list[int] | None = None, noise_idx: list[int] | None = None, permute_order: bool = False, use_loss_mask: bool = False, weight_by_complexity: bool = False, label_smoothing: float = 0.0, cut_off_tsm: float = 0.1) -> Array

sample_mask

sample_mask(rng: RngKey, acq: AcquisitionSchemeLike, x: Array, mask_prior: Array | None = None, y_ctx: Array | None = None, y: Array | None = None) -> Array

log_prob_mask

log_prob_mask(model_mask: Array, acq: AcquisitionSchemeLike, x: Array, mask_prior: Array | None = None, y_ctx: Array | None = None, y: Array | None = None) -> Array

sample_theta

sample_theta(rng: RngKey, acq: AcquisitionSchemeLike, x: Array, model_mask: Array, sample_method: str = 'ode', num_steps: int = 64, last_euler_step: bool = True, t_min: float | None = None, t_max: float | None = None, y_ctx: Array | None = None, y: Array | None = None) -> Array

log_prob_theta

log_prob_theta(theta: Array, acq: AcquisitionSchemeLike, x: Array, model_mask: Array, num_steps: int = 64, t_min: float | None = None, t_max: float | None = None) -> Array

score_theta

score_theta(theta: Array, acq: AcquisitionSchemeLike, x: ArrayLike, model_mask: Array, t: ArrayLike | None = None) -> Array

autoregressive

DMRIModelSelectionConfig dataclass

DMRIModelSelectionConfig(num_layers: int = 4, num_heads: int = 4, widening_factor: int = 3, attn_size: int = 16, dropout_rate: float = 0.0, prior_params_embed_dim: int = 0, mask_prior_dim: int | None = None, kv_in_features: int | None = None, use_flash_attention: bool = False, use_flash_cross_attention: bool = False, normalize_qk_attn: bool = False, normalize_qk_cross_attn: bool = False, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, outlayer: str = 'linear', outnorm: bool = True)

num_layers class-attribute instance-attribute

num_layers: int = 4

num_heads class-attribute instance-attribute

num_heads: int = 4

widening_factor class-attribute instance-attribute

widening_factor: int = 3

attn_size class-attribute instance-attribute

attn_size: int = 16

dropout_rate class-attribute instance-attribute

dropout_rate: float = 0.0

prior_params_embed_dim class-attribute instance-attribute

prior_params_embed_dim: int = 0

mask_prior_dim class-attribute instance-attribute

mask_prior_dim: int | None = None

kv_in_features class-attribute instance-attribute

kv_in_features: int | None = None

use_flash_attention class-attribute instance-attribute

use_flash_attention: bool = False

use_flash_cross_attention class-attribute instance-attribute

use_flash_cross_attention: bool = False

normalize_qk_attn class-attribute instance-attribute

normalize_qk_attn: bool = False

normalize_qk_cross_attn class-attribute instance-attribute

normalize_qk_cross_attn: bool = False

dtype class-attribute instance-attribute

dtype: DTypeLike | None = None

param_dtype class-attribute instance-attribute

param_dtype: DTypeLike | None = None

precision class-attribute instance-attribute

precision: PrecisionLike | None = None

preferred_element_type class-attribute instance-attribute

preferred_element_type: DTypeLike | None = None

outlayer class-attribute instance-attribute

outlayer: str = 'linear'

outnorm class-attribute instance-attribute

outnorm: bool = True

DMRIModelSelectionAmortizedPriorConfig dataclass

DMRIModelSelectionAmortizedPriorConfig(num_layers: int = 4, num_heads: int = 4, widening_factor: int = 3, attn_size: int = 16, dropout_rate: float = 0.0, prior_params_embed_dim: int = 64, mask_prior_dim: int | None = 1, kv_in_features: int | None = None, use_flash_attention: bool = False, use_flash_cross_attention: bool = False, normalize_qk_attn: bool = False, normalize_qk_cross_attn: bool = False, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, outlayer: str = 'linear', outnorm: bool = True)

Bases: DMRIModelSelectionConfig

prior_params_embed_dim class-attribute instance-attribute

prior_params_embed_dim: int = 64

mask_prior_dim class-attribute instance-attribute

mask_prior_dim: int | None = 1

BinaryAutoregressiveDecoder

BinaryAutoregressiveDecoder(rngs: Rngs, model_dim: int = 64, num_heads: int = 4, num_layers: int = 4, widening_factor: int = 4, attn_size: int = 16, dropout_rate: float = 0.0, prior_params_embed_dim: int = 0, mask_prior_dim: int | None = None, additional_context_dim: int = 0, kv_in_features: int | None = None, enable_cross_attention: bool = True, use_flash_attention: bool = False, use_flash_cross_attention: bool = False, normalize_qk_attn: bool = False, normalize_qk_cross_attn: bool = False, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, outlayer: str = 'linear', outnorm: bool = True)

Bases: Module

model_dim class-attribute instance-attribute

model_dim: int = model_dim

num_heads class-attribute instance-attribute

num_heads: int = num_heads

num_layers class-attribute instance-attribute

num_layers: int = num_layers

widening_factor class-attribute instance-attribute

widening_factor: int = widening_factor

attn_size class-attribute instance-attribute

attn_size: int = attn_size

prior_params_embed_dim instance-attribute

prior_params_embed_dim = prior_params_embed_dim

additional_context_dim instance-attribute

additional_context_dim = additional_context_dim

total_context_dim instance-attribute

total_context_dim = self.prior_params_embed_dim + self.additional_context_dim

context_dim instance-attribute

context_dim = context_dim

use_flash_attention instance-attribute

use_flash_attention = use_flash_attention

mask_prior_dim instance-attribute

mask_prior_dim = mask_prior_dim

mask_prior_embed instance-attribute

mask_prior_embed = GaussianFourierEmbedding(mask_prior_dim, self.prior_params_embed_dim, rngs=rngs, **precision_kwargs)

kv_in_features instance-attribute

kv_in_features = kv_in_features if kv_in_features is not None else model_dim

transformer instance-attribute

transformer = Transformer(model_dim, self.num_heads, self.num_layers, self.attn_size, kv_in_features=self.kv_in_features, context_dim=context_dim, widening_factor=self.widening_factor, rngs=rngs, dropout_rate=dropout_rate, attention_fn=attn_fn, cross_attention_fn=cross_attn_fn, enable_cross_attention=enable_cross_attention, normalize_qk_attn=normalize_qk_attn, normalize_qk_cross_attn=normalize_qk_cross_attn, **precision_kwargs)

out_norm instance-attribute

out_norm = nnx.LayerNorm(model_dim, rngs=rngs) if outnorm else (lambda value: value)

output instance-attribute

output = nnx.Linear(model_dim, 1, rngs=rngs, **precision_kwargs)

__call__

__call__(model_mask: Array, tokenizer: Tokenizer, mask_prior: Array | None = None, additional_context: Array | None = None, y: Array | None = None, mask: Array | None = None, decode: bool = False, deterministic: bool = False, **kwargs: Any) -> Array

loss_fn

loss_fn(model_mask: ArrayLike, tokenizer: Tokenizer, y: Array, rng: RngKey | None = None, permute_order: bool = False, label_smoothing: float = 0.0, mask_prior: Array | None = None, additional_context: Array | None = None, tokens_cfg: Array | None = None, **kwargs: Any) -> Array

sample

sample(key, tokenizer, y, dim, mask_prior: Array | None = None, additional_context: Array | None = None)

log_prob

log_prob(model_mask: Array, tokenizer: Tokenizer, y: Array, mask_prior: Array | None = None, additional_context: Array | None = None, **kwargs: Any) -> Array

naive_autoregressive_decoding

naive_autoregressive_decoding(model: BinaryAutoregressiveDecoder, key: RngKey, tokenizer: Tokenizer, y: Array, dim: int, mask_prior: Array | None = None, additional_context: Array | None = None) -> Array

Embeddings and tokenization

embedding_net

DMRIEmbeddingConfig dataclass

DMRIEmbeddingConfig(num_layers: int = 3, num_heads: int = 4, widening_factor: int = 2, use_flash_attention: bool = False, attn_size: int = 16, dropout_rate: float = 0.0, bvals_embed_dim: int = 3, signals_embed_dim: int = 3, bvec_repeats: int = 1, y_seq_dim: int | None = None, y_glob_dim: int | None = None, min_bval: float = 0.0, max_bval: float = 4000.0, min_signal: float = 0.0, max_signal: float = 1.0, soft_squash_signals: bool = True, log_transform_signals: bool = False, embed_signals: str = 'repeat', embed_bvals: str = 'fourier', dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, use_global_summary_token: bool = False, global_summary_bins: int = 8, out_norm: bool = False, reduce_factor: int = 1)

num_layers class-attribute instance-attribute

num_layers: int = 3

num_heads class-attribute instance-attribute

num_heads: int = 4

widening_factor class-attribute instance-attribute

widening_factor: int = 2

use_flash_attention class-attribute instance-attribute

use_flash_attention: bool = False

attn_size class-attribute instance-attribute

attn_size: int = 16

dropout_rate class-attribute instance-attribute

dropout_rate: float = 0.0

bvals_embed_dim class-attribute instance-attribute

bvals_embed_dim: int = 3

signals_embed_dim class-attribute instance-attribute

signals_embed_dim: int = 3

bvec_repeats class-attribute instance-attribute

bvec_repeats: int = 1

y_seq_dim class-attribute instance-attribute

y_seq_dim: int | None = None

y_glob_dim class-attribute instance-attribute

y_glob_dim: int | None = None

min_bval class-attribute instance-attribute

min_bval: float = 0.0

max_bval class-attribute instance-attribute

max_bval: float = 4000.0

min_signal class-attribute instance-attribute

min_signal: float = 0.0

max_signal class-attribute instance-attribute

max_signal: float = 1.0

soft_squash_signals class-attribute instance-attribute

soft_squash_signals: bool = True

log_transform_signals class-attribute instance-attribute

log_transform_signals: bool = False

embed_signals class-attribute instance-attribute

embed_signals: str = 'repeat'

embed_bvals class-attribute instance-attribute

embed_bvals: str = 'fourier'

dtype class-attribute instance-attribute

dtype: DTypeLike | None = None

param_dtype class-attribute instance-attribute

param_dtype: DTypeLike | None = None

precision class-attribute instance-attribute

precision: PrecisionLike | None = None

preferred_element_type class-attribute instance-attribute

preferred_element_type: DTypeLike | None = None

use_global_summary_token class-attribute instance-attribute

use_global_summary_token: bool = False

global_summary_bins class-attribute instance-attribute

global_summary_bins: int = 8

out_norm class-attribute instance-attribute

out_norm: bool = False

reduce_factor class-attribute instance-attribute

reduce_factor: int = 1

BvalBvecSignalEmbeddingNet

BvalBvecSignalEmbeddingNet(rngs: Rngs, model_dim: int = 64, num_heads: int = 4, num_layers: int = 3, widening_factor: int = 2, attn_size: int = 16, dropout_rate: float = 0.0, bvals_embed_dim: int = 3, signals_embed_dim: int = 3, bvec_repeats: int = 1, y_seq_dim: int | None = None, y_glob_dim: int | None = None, min_bval: float = 0.0, max_bval: float = 4000.0, min_signal: float = 0.0, max_signal: float = 1.0, soft_squash_signals: bool = True, log_transform_signals: bool = False, use_flash_attention: bool = False, embed_signals: str = 'repeat', embed_bvals: str = 'fourier', use_global_summary_token: bool = False, global_summary_bins: int = 8, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, out_norm: bool = False, reduce_factor: int = 1)

Bases: Module

num_heads class-attribute instance-attribute

num_heads: int = num_heads

num_layers class-attribute instance-attribute

num_layers: int = num_layers

widening_factor class-attribute instance-attribute

widening_factor: int = widening_factor

attn_size class-attribute instance-attribute

attn_size: int = attn_size

log_transform_signals instance-attribute

log_transform_signals = log_transform_signals

bvec_repeats instance-attribute

bvec_repeats = bvec_repeats

min_bval instance-attribute

min_bval = min_bval

max_bval instance-attribute

max_bval = max_bval

min_signal instance-attribute

min_signal = min_signal

max_signal instance-attribute

max_signal = max_signal

soft_squash_signals instance-attribute

soft_squash_signals = soft_squash_signals

use_global_summary_token instance-attribute

use_global_summary_token = use_global_summary_token

global_summary_bins instance-attribute

global_summary_bins = global_summary_bins

out_norm instance-attribute

out_norm = out_norm

reduce_factor instance-attribute

reduce_factor = reduce_factor

model_dim class-attribute instance-attribute

model_dim: int = model_dim // reduce_factor

initial_layer instance-attribute

initial_layer = nnx.Linear(bvals_embed_dim + signals_embed_dim + 3 * bvec_repeats, model_dim, rngs=rngs, **precision_kwargs)

embed_bvals instance-attribute

embed_bvals = GaussianFourierEmbedding(1, bvals_embed_dim, rngs=rngs, **precision_kwargs)

embed_signals instance-attribute

embed_signals = lambda x: jnp.repeat(x, signals_embed_dim, axis=-1)

transformer instance-attribute

transformer = Transformer(model_dim, self.num_heads, self.num_layers, self.attn_size, widening_factor=self.widening_factor, rngs=rngs, dropout_rate=dropout_rate, attention_fn=attention_fn, **precision_kwargs)

global_summary_feature_dim instance-attribute

global_summary_feature_dim = 3 * self.global_summary_bins + 12

global_summary_projection instance-attribute

global_summary_projection = nnx.Linear(self.global_summary_feature_dim, model_dim, rngs=rngs, **precision_kwargs)

global_summary_outnorm instance-attribute

global_summary_outnorm = nnx.LayerNorm(model_dim, rngs=rngs)

y_seq_dim instance-attribute

y_seq_dim = y_seq_dim

output_seq_layer instance-attribute

output_seq_layer = nnx.Linear(model_dim, y_seq_dim, rngs=rngs, **precision_kwargs)

out_norm_seq instance-attribute

out_norm_seq = nnx.LayerNorm(y_seq_dim, rngs=rngs) if out_norm else (lambda x: x)

y_glob_dim instance-attribute

y_glob_dim = y_glob_dim

output_glob_layer instance-attribute

output_glob_layer = nnx.Linear(model_dim, y_glob_dim, rngs=rngs, **precision_kwargs)

out_norm_glob instance-attribute

out_norm_glob = nnx.LayerNorm(y_glob_dim, rngs=rngs) if out_norm else (lambda x: x)

transform_bvals

transform_bvals(bvals: ArrayLike) -> Array

transform_signals

transform_signals(signals: ArrayLike) -> Array

global_summary_token

global_summary_token(bvals, bvecs, signals) -> Array

__call__

__call__(acq: AcquisitionScheme, x: ArrayLike, deterministic: bool | None = None, decode: bool = False) -> tuple[Array | None, Array]

SSFPEmbeddingNetConfig dataclass

SSFPEmbeddingNetConfig(num_layers: int = 3, num_heads: int = 4, widening_factor: int = 2, use_flash_attention: bool = False, attn_size: int = 16, dropout_rate: float = 0.0, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, use_global_summary_token: bool = False, global_summary_bins: int = 8, y_seq_dim: int | None = None, y_glob_dim: int | None = None)

num_layers class-attribute instance-attribute

num_layers: int = 3

num_heads class-attribute instance-attribute

num_heads: int = 4

widening_factor class-attribute instance-attribute

widening_factor: int = 2

attn_size class-attribute instance-attribute

attn_size: int = 16

dropout_rate class-attribute instance-attribute

dropout_rate: float = 0.0

use_flash_attention class-attribute instance-attribute

use_flash_attention: bool = False

dtype class-attribute instance-attribute

dtype: DTypeLike | None = None

param_dtype class-attribute instance-attribute

param_dtype: DTypeLike | None = None

precision class-attribute instance-attribute

precision: PrecisionLike | None = None

preferred_element_type class-attribute instance-attribute

preferred_element_type: DTypeLike | None = None

use_global_summary_token class-attribute instance-attribute

use_global_summary_token: bool = False

global_summary_bins class-attribute instance-attribute

global_summary_bins: int = 8

y_seq_dim class-attribute instance-attribute

y_seq_dim: int | None = None

y_glob_dim class-attribute instance-attribute

y_glob_dim: int | None = None

SSFPEmbeddingNet

SSFPEmbeddingNet(rngs: Rngs, model_dim: int = 64, num_heads: int = 4, num_layers: int = 3, widening_factor: int = 2, attn_size: int = 16, dropout_rate: float = 0.0, use_flash_attention: bool = False, use_global_summary_token: bool = False, global_summary_bins: int = 8, y_seq_dim: int | None = None, y_glob_dim: int | None = None, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None)

Bases: Module

model_dim class-attribute instance-attribute

model_dim: int = model_dim

num_heads class-attribute instance-attribute

num_heads: int = num_heads

num_layers class-attribute instance-attribute

num_layers: int = num_layers

widening_factor class-attribute instance-attribute

widening_factor: int = widening_factor

attn_size class-attribute instance-attribute

attn_size: int = attn_size

use_global_summary_token instance-attribute

use_global_summary_token = use_global_summary_token

global_summary_bins instance-attribute

global_summary_bins = global_summary_bins

embed_scalars instance-attribute

embed_scalars = GaussianFourierEmbedding(7, scalar_embed_dim, rngs=rngs)

embed_signals instance-attribute

embed_signals = GaussianFourierEmbedding(1, signal_embed_dim, rngs=rngs)

embed_bvecs instance-attribute

embed_bvecs: Callable[[ArrayLike], Array] = _repeat_bvecs

transformer instance-attribute

transformer = Transformer(model_dim, self.num_heads, self.num_layers, self.attn_size, widening_factor=self.widening_factor, rngs=rngs, dropout_rate=dropout_rate, attention_fn=attention_fn, **transformer_kwargs)

global_summary_feature_dim instance-attribute

global_summary_feature_dim = 3 * self.global_summary_bins + 18

global_summary_projection instance-attribute

global_summary_projection = nnx.Linear(self.global_summary_feature_dim, model_dim, rngs=rngs, **linear_kwargs)

y_seq_dim instance-attribute

y_seq_dim = y_seq_dim

output_seq_layer instance-attribute

output_seq_layer = nnx.Linear(model_dim, y_seq_dim, rngs=rngs, **linear_kwargs)

y_glob_dim instance-attribute

y_glob_dim = y_glob_dim

output_glob_layer instance-attribute

output_glob_layer = nnx.Linear(model_dim, y_glob_dim, rngs=rngs, **linear_kwargs)

global_summary_token

global_summary_token(acq: ssfp_acquisition_scheme, signals: ArrayLike) -> Array

__call__

__call__(acq: ssfp_acquisition_scheme, signals: ArrayLike, deterministic: bool | None = None, decode: bool = False) -> tuple[Array | None, Array]

soft_squash

soft_squash(signal, tau=0.5)

tokenizer

Tokenizer

Bases: Module

__call__

__call__(*args: Any, **kwds: Any) -> Array

encode abstractmethod

encode(*args: Any, **kwargs: Any) -> Array

decode abstractmethod

decode(*args: Any, **kwargs: Any) -> Array

DMRITokenizer

DMRITokenizer(simulator: type[MultiCompartment], *, token_dim: int = 64, theta_encode_nets: List[Module | None] | None = None, theta_decode_nets: List[Module | None] | None = None, init_component_embeddings: Callable | None = None, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None, rngs: Rngs)

Bases: Tokenizer

A tokenizer for dMRI multi-compartment model configurations. It handles embedding and decoding of model types, noise types, and associated parameters.

Initializes a DMRITokenizer instance.

Parameters:

Name Type Description Default
simulator type[MultiCompartment]

The multi-compartment simulator class which defines model_types, noise_types, etc.

required
rngs Any

Random number generator(s) used for parameter initialization.

required
token_dim int

Dimension of tokens for embedding. Defaults to 64.

64
theta_encode_nets Optional[List[Module]]

Encoding modules for parameters. Defaults to None.

None
theta_decode_nets Optional[List[Module]]

Decoding modules for parameters. Defaults to None.

None

simulator instance-attribute

simulator = nnx.static(simulator)

num_models instance-attribute

num_models = len(simulator.model_types)

num_noises instance-attribute

num_noises = len(simulator.noise_types)

params_dims instance-attribute

params_dims = tuple(simulator.split_idx())

model_indices instance-attribute

model_indices: tuple[int, ...] = tuple(range(self.num_models))

noise_indices instance-attribute

noise_indices: tuple[int, ...] = tuple(range(self.num_noises))

embed_idx instance-attribute

embed_idx = nnx.Embed(rngs=rngs, num_embeddings=len(simulator.model_types) + len(simulator.noise_types), features=token_dim, embedding_init=self._init_class_embeddings if init_component_embeddings is None else init_component_embeddings)

embed_fraction instance-attribute

embed_fraction = nnx.Linear(len(simulator.model_types), token_dim, rngs=rngs, **self._linear_kwargs)

shared_parameter_embed instance-attribute

shared_parameter_embed = nnx.Linear(simulator.shared_parameter_type.theta_dim, token_dim, rngs=rngs, **self._linear_kwargs)

shared_parameter_decode instance-attribute

shared_parameter_decode = nnx.Linear(token_dim, simulator.shared_parameter_type.theta_dim, rngs=rngs, **self._linear_kwargs)

shared_parameter_idx instance-attribute

shared_parameter_idx = nnx.Embed(rngs=rngs, num_embeddings=1, features=token_dim)

theta_encode_nets instance-attribute

theta_encode_nets = theta_encode_nets

theta_decode_nets instance-attribute

theta_decode_nets = theta_decode_nets

encode

encode(theta: ArrayLike | None = None, model_mask: ArrayLike | None = None, tokens_cfg: Array | None = None, alpha_prior: ArrayLike | None = None, model_idx: Sequence[int] | None = None, noise_idx: Sequence[int] | None = None) -> Array

decode

decode(tokens: ArrayLike, model_idx: Sequence[int] | None = None, noise_idx: Sequence[int] | None = None, **kwargs: Any) -> Array

get_indices_with_params

get_indices_with_params(model_idx: Sequence[int], noise_idx: Sequence[int]) -> list[int]

Returns a list of model and noise indices corresponding to the provided model and noise types.

theta_fraction_mask staticmethod

theta_fraction_mask(model_mask: Array) -> Array

theta_token_mask

theta_token_mask(model_mask: Array, model_idx: Sequence[int] | None = None, noise_idx: Sequence[int] | None = None) -> Array

embed_cfgs

embed_cfgs(model_mask: Array, alpha_prior: ArrayLike | None = None, model_idx: Sequence[int] | None = None, noise_idx: Sequence[int] | None = None) -> Array

Embeds the configuration of model and noise types into tokens.

Token Structure

The output tokens have the following structure:

  1. alpha_token (B, 1, token_dim):
  2. Represents the prior fractions for model components
  3. Position: First token in the sequence
  4. Shape: (batch_dims..., 1, token_dim)

  5. idx_tokens (B, T, token_dim):

  6. Represents the embedded indices for each model and noise component
  7. Value: If the component mask is True the token will be the embedded index, if the component mask is False the token will be zero
  8. Position: Follows the alpha_token
  9. Shape: (batch_dims..., T, token_dim) where T is the number of components
  10. Components that are not active (masked out) will have zero token values

The final output is a concatenation of these tokens along the second-to-last axis, resulting in a tensor of shape (batch_dims..., 1+T, token_dim).

Parameters:

Name Type Description Default
model_mask ArrayLike

A binary mask indicating active model components.

required
alpha_prior Optional[ArrayLike]

Prior fractions for model components.

None
model_idx Optional[Sequence[int]]

Indices of model components.

None
noise_idx Optional[Sequence[int]]

Indices of noise components.

None

Returns:

Name Type Description
ArrayLike Array

The embedded tokens for each model/noise component.

embed_theta

embed_theta(theta: Array, tokens_cfg: Array, model_idx: Sequence[int] | None = None, noise_idx: Sequence[int] | None = None, model_mask: Array | None = None) -> Array

Embeds the continuous parameter vector theta into token representation.

Parameters:

Name Type Description Default
theta ArrayLike

The parameters of each model/noise component.

required
tokens_cfg ArrayLike

Configuration tokens returned by embed_cfgs.

required
model_idx Optional[Sequence[int]]

Subset of model indices to include.

None
noise_idx Optional[Sequence[int]]

Subset of noise indices to include.

None
model_mask Optional[Array]

Optional mask over components.

None

Returns:

Name Type Description
ArrayLike Array

The token representation augmented with encoded parameters.

Token Structure

The output tokens have the following structure:

  1. val_tokens (B, N, token_dim):
  2. Represents the encoded parameter values for each component
  3. Shape: (batch_dims..., N, token_dim) where N is the number of components with parameters
  4. Each token corresponds to the parameters of a specific model or noise component
  5. Components without parameters are excluded from the output

  6. tokens_cfg (B, N, token_dim):

  7. Configuration tokens from embed_cfgs, filtered to match the components with parameters
  8. Shape: (batch_dims..., N, token_dim)

The final output is the sum of val_tokens and tokens_cfg, resulting in a tensor of shape (batch_dims..., N, token_dim). This combines the parameter information with the component identity information in a single token representation.

Processing Steps
  1. The theta vector is split into components based on the parameter dimensions of each model/noise type
  2. Each component's parameters are encoded using the corresponding encoding network
  3. Components without parameters are filtered out
  4. The encoded parameters are concatenated along the second-to-last axis
  5. The configuration tokens are filtered to match only the components with parameters
  6. The final tokens are the sum of the encoded parameters and the filtered configuration tokens

decode_theta

decode_theta(tokens: Array, model_idx: Sequence[int] | None = None, noise_idx: Sequence[int] | None = None, model_mask: Array | None = None, **kwargs: Any) -> Array

Decode tokens back into the continuous parameter vector theta.

DMRITokenizerPP

DMRITokenizerPP(simulator, rngs, token_dim=64, theta_encode_nets=None, theta_decode_nets=None, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None)

Bases: DMRITokenizer

fraction_embed instance-attribute

fraction_embed = nnx.Embed(len(self.simulator.model_types) - 1, token_dim, rngs=rngs, embedding_init=nnx.initializers.orthogonal())

theta_fraction_mask staticmethod

theta_fraction_mask(model_mask: Array) -> Array

Creates a mask for the model fractions.

embed_theta

embed_theta(theta: Array, tokens_cfg: Array, model_idx: Sequence[int] | None = None, noise_idx: Sequence[int] | None = None, model_mask: Array | None = None) -> Array

decode_theta

decode_theta(tokens: Array, model_idx: Sequence[int] | None = None, noise_idx: Sequence[int] | None = None, model_mask: Array | None = None, **kwargs: Any) -> Array

DMRITokenizerPPP

DMRITokenizerPPP(simulator, rngs, token_dim=64, theta_encode_nets=None, theta_decode_nets=None, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None)

Bases: DMRITokenizerPP

Legacy mask-aware tokenizer retained for published checkpoints.

start_end_token instance-attribute

start_end_token = nnx.Embed(2, token_dim // 2, rngs=rngs)

model_mask_emb instance-attribute

model_mask_emb = nnx.Embed(2, token_dim, rngs=rngs)

cfg_emb instance-attribute

cfg_emb = nnx.Linear(token_dim, token_dim - token_dim // 2, rngs=rngs, **self._linear_kwargs)

embed_cfgs

embed_cfgs(model_mask: Array, alpha_prior: ArrayLike | None = None, model_idx: Sequence[int] | None = None, noise_idx: Sequence[int] | None = None) -> Array

Embed the component configuration without masking the index tokens.

Unlike the base tokenizer, this variant must not zero the index token of inactive components: embed_model_mask packs cfg token j into decoder position j, which is exactly the position that predicts bit j. Masking here would leak the target into its own input. The mask enters only through embed_model_mask, shifted by one position.

embed_model_mask

embed_model_mask(tokens_cfg: Array, model_mask: Array) -> Array

Transformer backbones

simformer

DMRIThetaInferenceConfig dataclass

DMRIThetaInferenceConfig(num_layers: int = 6, num_heads: int = 4, widening_factor: int = 3, attn_size: int = 16, time_embed_dim: int = 64, dropout_rate: float = 0.0, use_flash_attention: bool = False, use_flash_cross_attention: bool = False, normalize_qk_attn: bool = False, normalize_qk_cross_attn: bool = False, gate_attention: bool = False, gate_mlp: bool = False, kv_in_features: int | None = None, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None)

num_layers class-attribute instance-attribute

num_layers: int = 6

num_heads class-attribute instance-attribute

num_heads: int = 4

widening_factor class-attribute instance-attribute

widening_factor: int = 3

attn_size class-attribute instance-attribute

attn_size: int = 16

time_embed_dim class-attribute instance-attribute

time_embed_dim: int = 64

dropout_rate class-attribute instance-attribute

dropout_rate: float = 0.0

use_flash_attention class-attribute instance-attribute

use_flash_attention: bool = False

use_flash_cross_attention class-attribute instance-attribute

use_flash_cross_attention: bool = False

normalize_qk_attn class-attribute instance-attribute

normalize_qk_attn: bool = False

normalize_qk_cross_attn class-attribute instance-attribute

normalize_qk_cross_attn: bool = False

gate_attention class-attribute instance-attribute

gate_attention: bool = False

gate_mlp class-attribute instance-attribute

gate_mlp: bool = False

kv_in_features class-attribute instance-attribute

kv_in_features: int | None = None

dtype class-attribute instance-attribute

dtype: DTypeLike | None = None

param_dtype class-attribute instance-attribute

param_dtype: DTypeLike | None = None

precision class-attribute instance-attribute

precision: PrecisionLike | None = None

preferred_element_type class-attribute instance-attribute

preferred_element_type: DTypeLike | None = None

DiffusionTransformer

DiffusionTransformer(rngs: Rngs, model_dim: int = 64, time_embed_dim: int = 64, additional_context_dim: int = 0, num_heads: int = 4, num_layers: int = 6, attn_size: int = 16, widening_factor: int = 3, dropout_rate: float = 0.0, enable_cross_attention: bool = True, use_flash_attention: bool = False, use_flash_cross_attention: bool = False, normalize_qk_attn: bool = False, normalize_qk_cross_attn: bool = False, gate_attention: bool = False, gate_mlp: bool = False, kv_in_features: int | None = None, dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None)

Bases: Module

time_embed_dim instance-attribute

time_embed_dim = time_embed_dim

additional_context_dim instance-attribute

additional_context_dim = additional_context_dim

total_context_dim instance-attribute

total_context_dim = self.time_embed_dim + self.additional_context_dim

time_embedding instance-attribute

time_embedding = GaussianFourierEmbedding(1, self.time_embed_dim, rngs=rngs, **precision_kwargs)

kv_in_features instance-attribute

kv_in_features = kv_in_features if kv_in_features is not None else model_dim

transformer instance-attribute

transformer = Transformer(model_dim, num_heads=num_heads, num_layers=num_layers, attn_size=attn_size, kv_in_features=self.kv_in_features, widening_factor=widening_factor, enable_cross_attention=enable_cross_attention, dropout_rate=dropout_rate, rngs=rngs, context_dim=transformer_context_dim, attention_fn=attn_fn, cross_attention_fn=cross_attn_fn, normalize_qk_attn=normalize_qk_attn, normalize_qk_cross_attn=normalize_qk_cross_attn, attn_fuse_cls=attn_fuse_cls, mlp_fuse_cls=mlp_fuse_cls, **precision_kwargs)

out_norm instance-attribute

out_norm = nnx.LayerNorm(model_dim, rngs=rngs)

__call__

__call__(t: ArrayLike, x: ArrayLike, tokenizer: Tokenizer, y: Array | None = None, context: Array | None = None, attention_mask: Array | None = None, **kwargs: Any) -> Array

EDMSimformer

EDMSimformer(rngs: Rngs, model_dim: int = 64, time_embed_dim: int = 64, additional_context_dim: int = 0, num_heads: int = 4, num_layers: int = 4, attn_size: int = 16, widening_factor: int = 3, dropout_rate: float = 0.0, enable_cross_attention: bool = True, use_flash_attention: bool = False, use_flash_cross_attention: bool = False, normalize_qk_attn: bool = False, normalize_qk_cross_attn: bool = False, gate_attention: bool = False, gate_mlp: bool = False, kv_in_features: int | None = None, loss_type: str = 'x0', dtype: DTypeLike | None = None, param_dtype: DTypeLike | None = None, precision: PrecisionLike | None = None, preferred_element_type: DTypeLike | None = None)

Bases: EDM

loss

loss(rng: RngKey, theta: Array, tokenizer: Tokenizer, y: Array, model_mask: ArrayLike, tokens_cfg: Array, target_score: Array | None = None, attention_mask: ArrayLike | None = None, loss_mask: ArrayLike | None = None, weight_by_complexity: bool = False, cut_off_tsm: float = 0.1, context: ArrayLike | None = None) -> Array

sample

sample(rng: RngKey, tokenizer: Tokenizer, y: Array, dim: int, tokens_cfg: Array | None = None, model_mask: ArrayLike | None = None, context: ArrayLike | None = None, attention_mask: ArrayLike | None = None, sample_method: str = 'ode', num_steps: int = 64, last_euler_step: bool = False, t_min: float | None = None, t_max: float | None = None) -> Array

log_prob

log_prob(x: Array, tokenizer: Tokenizer, y: Array, tokens_cfg: Array | None = None, context: ArrayLike | None = None, attention_mask: ArrayLike | None = None, model_mask: ArrayLike | None = None, t_min: float | None = None, t_max: float | None = None, num_steps: int = 64) -> Array

masked_standard_normal_log_prob

masked_standard_normal_log_prob(x: Array, sigma: Array, mask: Array) -> Array

Return a spherical Gaussian log density over active dimensions only.