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
¶
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())
preferred_element_type
class-attribute
instance-attribute
¶
embedding_cls
class-attribute
instance-attribute
¶
embedding_cls: type[EmbeddingModule] = BvalBvecSignalEmbeddingNet
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
embedding_cls
class-attribute
instance-attribute
¶
embedding_cls: type[EmbeddingModule] = BvalBvecSignalEmbeddingNet
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.
embedding_cls
class-attribute
instance-attribute
¶
embedding_cls: type[EmbeddingModule] = BvalBvecSignalEmbeddingNet
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
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
precision_fields
instance-attribute
¶
__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 |
None
|
use_flash_cross_attention
|
bool | None
|
Optional override applied to every
sub-config exposing |
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)
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
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
total_context_dim
instance-attribute
¶
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
¶
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
¶
__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)
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)
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
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_signals
instance-attribute
¶
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_projection
instance-attribute
¶
global_summary_projection = nnx.Linear(self.global_summary_feature_dim, model_dim, rngs=rngs, **precision_kwargs)
global_summary_outnorm
instance-attribute
¶
output_seq_layer
instance-attribute
¶
out_norm_seq
instance-attribute
¶
output_glob_layer
instance-attribute
¶
out_norm_glob
instance-attribute
¶
__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)
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
embed_scalars
instance-attribute
¶
embed_signals
instance-attribute
¶
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_projection
instance-attribute
¶
global_summary_projection = nnx.Linear(self.global_summary_feature_dim, model_dim, rngs=rngs, **linear_kwargs)
output_seq_layer
instance-attribute
¶
output_glob_layer
instance-attribute
¶
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]
tokenizer
¶
Tokenizer
¶
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
|
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
¶
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
¶
Returns a list of model and noise indices corresponding to the provided model and noise types.
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:
- alpha_token (B, 1, token_dim):
- Represents the prior fractions for model components
- Position: First token in the sequence
-
Shape: (batch_dims..., 1, token_dim)
-
idx_tokens (B, T, token_dim):
- Represents the embedded indices for each model and noise component
- 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
- Position: Follows the alpha_token
- Shape: (batch_dims..., T, token_dim) where T is the number of components
- 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:
- val_tokens (B, N, token_dim):
- Represents the encoded parameter values for each component
- Shape: (batch_dims..., N, token_dim) where N is the number of components with parameters
- Each token corresponds to the parameters of a specific model or noise component
-
Components without parameters are excluded from the output
-
tokens_cfg (B, N, token_dim):
- Configuration tokens from embed_cfgs, filtered to match the components with parameters
- 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
- The theta vector is split into components based on the parameter dimensions of each model/noise type
- Each component's parameters are encoded using the corresponding encoding network
- Components without parameters are filtered out
- The encoded parameters are concatenated along the second-to-last axis
- The configuration tokens are filtered to match only the components with parameters
- 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
¶
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.
cfg_emb
instance-attribute
¶
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.
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)
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
total_context_dim
instance-attribute
¶
time_embedding
instance-attribute
¶
kv_in_features
instance-attribute
¶
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)
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
masked_standard_normal_log_prob
¶
Return a spherical Gaussian log density over active dimensions only.