Modality

class stix.core.modality.Modality(*args, **kwargs)

Bases: Module, Generic

Per-modality configuration.

All constructor fields default to None so a registry can be built incrementally — e.g. an empty skeleton from ModalityRegistry.from_batch, populated later via ModalityRegistry.set. However, None is a construction-time state, not a claim that a field is optional at runtime: the type-level | None only reflects what is permitted during construction, and each consumer validates (and raises for) the fields it actually needs. E.g.

  • shape, is_discrete: read off the raw target by from_batch

  • interpolant, embedder: required by GenerativeModel.__init__

  • num_categories: required for index-valued (categorical) discrete

    modalities when the raw shape does not carry the category count (namely, when the modality is not one-hot encoded)

  • embedded_source_prior: optional; samples the embedded source for

    two-sided interpolants (training fill + sampling). Invalid with a one-sided interpolant.

The sampling-time generator_type is inferred from the interpolant (see generator_type).

An nnx.Module, so the modality is itself a JAX pytree and the parameters of a stored embedder are visible to NNX tracing.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

property generator_type: type[Generator]

Sampling-time generator class inferred from the interpolant.

Returns:

The Generator subclass declared on the interpolant (e.g. VelocityAndScore for a stochastic continuous path, TransitionRates for a discrete DFM interpolant).

Raises:

ValueError – If the interpolant has not been set yet.

class stix.core.modality.ModalityRegistry(*args, **kwargs)

Bases: Module

Container describing a pytree of Modality objects.

An nnx.Module, so a model holding a registry (e.g. GenerativeModel) exposes the registry-held Modality params, such as a learnable embedder, to NNX tracing and nnx.split. The inner registry pytree is wrapped in nnx.data so NNX descends into it rather than treating it as an opaque static leaf; otherwise those params are invisible to nnx.split and silently frozen at initialisation.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

property treedef

The registry’s pytree structure with each Modality as a leaf.

classmethod from_batch(batch)

Build a registry skeleton by inferring structure from a training Batch.

The attributes shape, is_discrete and (when inferable) num_categories are read off the raw target array.

  • shape: the per-example target shape (leading batch axis dropped).

  • is_discrete: taken directly from batch.is_discrete. Whatever

    builds the Batch must state discreteness explicitly.

  • num_categories: inferred only for one-hot discrete data,

    where it is the last dimension of shape. One-hot vectors are identified by checking they hold binary values and sum to 1 (see is_one_hot_vector()). An integer dtype is taken to be index-valued, never one-hot. For index-valued discrete data (integer targets, shape == (L,)), the count cannot be inferred and is left None; set it explicitly via registry.set("num_categories", ...).

Parameters:

batch (Batch) – A training batch whose raw_batch and is_discrete define the registry structure.

Returns:

A ModalityRegistry with shape / is_discrete / num_categories filled where inferable.

Raises:

ValueError – If batch.is_discrete is None.

Return type:

ModalityRegistry

set(field, value, filter_fn=None, use_deepcopy=True, is_factory=False)

Set a per-modality field, broadcasting each argument over the registry.

One can either pass a single value or a pytree. The value argument is broadcasted to every modality beneath it.

When is_factory is True, each leaf of value is treated as a factory Callable[[Modality], Any] and invoked with the target modality to produce that modality’s value.

Examples

Considering a registry shaped as {"mod_a": ..., "mod_b": {"mod_b1": ..., "mod_b2": ...}}:

# Broadcast one value to every modality.
registry.set("interpolant", interpolant)

# Prefix pytree: mod_a gets its own value; the single value
# under "mod_b" fans out to both mod_b1 and mod_b2.
registry.set(
    "interpolant", {"mod_a": interpolant_a, "mod_b": interpolant_b}
)

# Fully-resolved pytree: one value per modality.
registry.set(
    "interpolant",
    {"mod_a": itp_a, "mod_b": {"mod_b1": itp_b1, "mod_b2": itp_b2}},
)

# Factory: build each modality's value from its Modality.
registry.set(
    "embedder",
    lambda m: IdentityEmbedder(m.shape),
    is_factory=True,
)

# filter_fn: only touch matching modalities.
registry.set(
    "embedder", continuous_embedder, filter_fn=lambda m: not m.is_discrete
)
Parameters:
  • field (str) – The field to set.

  • value (Any | PyTree[Any]) – The value to set the field to; a leaf broadcasts to every modality beneath it. With is_factory, each leaf is instead a factory Callable[[Modality], Any].

  • filter_fn (Callable[[Modality], bool] | None) – Optional predicate; only modalities returning True are modified (None affects all).

  • use_deepcopy (bool) – Whether to deep-copy each broadcast value so modalities do not alias one shared object. Ignored under is_factory. Default is True. A plain function (e.g. a jr.normal prior) is copied atomically and so survives as the same object; a stateful callable is snapshotted per modality, so pass False to share one live object across them.

  • is_factory (bool) – Treat every leaf of value as a factory called with the target modality to build its value.

Return type:

None

assert_compatible(*trees)

Check each tree is structurally compatible with the registry.

The registry’s Modality-leaf treedef must be a prefix of each tree. A tree may resolve finer below each modality (e.g. a RawSourceTargetPair splitting into source/target, where a None source changes the leaf count), but its structure above the modalities must match the registry exactly.

jax.tree.map enforces this implicitly when the registry is its first argument (via flatten_up_to); this check runs it explicitly to raise a readable, registry-specific error instead.

Parameters:

*trees (PyTree) – The pytrees to validate against the registry.

Raises:

ValueError – If any tree is not compatible with the registry structure.

Return type:

None

assert_fields_set(*fields)

Raise if any modality leaf leaves any of fields unset (None).

Parameters:

*fields (str) – Field names that must be non-None on every modality.

Raises:

ValueError – If any modality has None for any requested field.

Return type:

None

assert_num_categories_set_on_discrete_modalities()

Raise if any discrete modality has num_categories not set.

Return type:

None

map(f, *trees)

jax.tree.map over the registry, treating each Modality as a leaf.

This is the canonical way to map over per-modality data. It bundles three operations together:

  1. Assert that all pytrees are compatible with the registry.

  2. Run jax.tree.map with the registry as the first pytree.

  3. Set is_leaf so that each Modality is treated as a leaf.

Parameters:
  • f (Callable) – Callable receiving the modality followed by the corresponding leaf of each tree in trees.

  • *trees (PyTree) – Pytrees to map over, each compatible with the registry.

Returns:

A pytree with the registry’s structure holding the results of f.

Return type:

PyTree

broadcast(prefix_tree, is_leaf=None)

Broadcast a prefix tree up to the registry’s per-modality structure.

The registry’s modality-key structure is the target; each leaf of prefix_tree is copied to every modality beneath it. prefix_tree may be coarser than the modality structure — a single value broadcast to all modalities, or one value per modality group.

Use this for coarse->fine broadcasting (one value to many modalities); use map() for per-modality computation. By default container values are descended as pytrees; pass is_leaf to hold a structured value intact at each modality (so a compound value broadcasts as a single leaf rather than being descended into). A broadcast leaf is shared (the same object) across modalities.

Parameters:
  • prefix_tree (Any | PyTree[Any]) – A value or coarser pytree to broadcast onto the registry.

  • is_leaf (Callable[[Any], bool] | None) – Optional predicate marking values that should not be descended as pytrees during broadcasting.

Returns:

A pytree with the registry’s modality structure, holding one prefix_tree leaf per modality.

Raises:

ValueError – If prefix_tree is not a prefix of the registry’s modality structure.

Return type:

PyTree

split_and_project_key(key)

Split key into one independent key per modality, projected onto the registry structure.

The PRNG-dual of broadcast(): broadcast shares one value across modalities; this hands each its own key. Feed the result to map().

Parameters:

key (Array) – PRNG key to split across modalities.

Returns:

A pytree of independent keys with the registry’s modality structure.

Return type:

PyTree[jax.Array]

sample_initial_state(key, raw_source=None, num_samples=None)

Sample the per-modality initial state \(z_0\).

Resolves \(z_{\mathrm{src}}\) and \(\epsilon\) for every modality, then delegates to sample_initial_state().

Per modality:

  • a raw source array is embedded and used as \(z_{\mathrm{src}}\) (wins over embedded_source_prior);

  • None on a one-sided interpolant stays None;

  • None on a two-sided interpolant with embedded_source_prior is filled from that prior;

  • a raw source on a one-sided interpolant, or a two-sided interpolant with neither a raw source nor a prior, raises.

raw_source=None means every leaf is missing. num_samples is inferred from any provided source; it is required when every leaf is missing. If both are given they must match.

Parameters:
  • key (Array) – PRNG key; split into one independent key per modality.

  • raw_source (PyTree[Var | None] | None) – Optional pytree of raw source arrays aligned with the registry (None leaves are missing). None itself is “all leaves missing”.

  • num_samples (int | None) – Leading batch axis. Required when no raw source is provided; must match the source batch size when both are given.

Returns:

Per-modality initial embedded state \(z_0\).

Raises:
  • ValueError – If raw_source and num_samples are both None; if provided sources disagree on batch size or disagree with num_samples; if a one-sided modality is given a raw source; or if a modality is missing its interpolant or embedder.

  • TypeError – If a two-sided modality has neither a raw source nor embedded_source_prior.

Return type:

PyTree[Var]