Modality¶
- class stix.core.modality.Modality(*args, **kwargs)¶
-
Per-modality configuration.
All constructor fields default to
Noneso a registry can be built incrementally — e.g. an empty skeleton fromModalityRegistry.from_batch, populated later viaModalityRegistry.set. However,Noneis a construction-time state, not a claim that a field is optional at runtime: the type-level| Noneonly 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 byfrom_batchinterpolant,embedder: required byGenerativeModel.__init__num_categories: required for index-valued (categorical) discretemodalities when the raw
shapedoes not carry the category count (namely, when the modality is not one-hot encoded)
embedded_source_prior: optional; samples the embedded source fortwo-sided interpolants (training fill + sampling). Invalid with a one-sided interpolant.
The sampling-time
generator_typeis inferred from the interpolant (seegenerator_type).An
nnx.Module, so the modality is itself a JAX pytree and the parameters of a storedembedderare 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
Generatorsubclass declared on the interpolant (e.g.VelocityAndScorefor a stochastic continuous path,TransitionRatesfor a discrete DFM interpolant).- Raises:
ValueError – If the interpolant has not been set yet.
- class stix.core.modality.ModalityRegistry(*args, **kwargs)¶
Bases:
ModuleContainer describing a pytree of
Modalityobjects.An
nnx.Module, so a model holding a registry (e.g.GenerativeModel) exposes the registry-heldModalityparams, such as a learnableembedder, to NNX tracing andnnx.split. The innerregistrypytree is wrapped innnx.dataso NNX descends into it rather than treating it as an opaque static leaf; otherwise those params are invisible tonnx.splitand silently frozen at initialisation.- Parameters:
args (Any)
kwargs (Any)
- Return type:
Any
- property treedef¶
The registry’s pytree structure with each
Modalityas a leaf.
- classmethod from_batch(batch)¶
Build a registry skeleton by inferring structure from a training
Batch.The attributes
shape,is_discreteand (when inferable)num_categoriesare read off the raw target array.shape: the per-example target shape (leading batch axis dropped).is_discrete: taken directly frombatch.is_discrete. Whateverbuilds the
Batchmust 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 (seeis_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 leftNone; set it explicitly viaregistry.set("num_categories", ...).
- Parameters:
batch (Batch) – A training batch whose
raw_batchandis_discretedefine the registry structure.- Returns:
A
ModalityRegistrywithshape/is_discrete/num_categoriesfilled where inferable.- Raises:
ValueError – If
batch.is_discreteisNone.- Return type:
- 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_factoryisTrue, each leaf ofvalueis treated as a factoryCallable[[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 factoryCallable[[Modality], Any].filter_fn (Callable[[Modality], bool] | None) – Optional predicate; only modalities returning
Trueare modified (Noneaffects all).use_deepcopy (bool) – Whether to deep-copy each broadcast value so modalities do not alias one shared object. Ignored under
is_factory. Default isTrue. A plain function (e.g. ajr.normalprior) is copied atomically and so survives as the same object; a stateful callable is snapshotted per modality, so passFalseto share one live object across them.is_factory (bool) – Treat every leaf of
valueas 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. aRawSourceTargetPairsplitting intosource/target, where aNonesource changes the leaf count), but its structure above the modalities must match the registry exactly.jax.tree.mapenforces this implicitly when the registry is its first argument (viaflatten_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
fieldsunset (None).- Parameters:
*fields (str) – Field names that must be non-
Noneon every modality.- Raises:
ValueError – If any modality has
Nonefor any requested field.- Return type:
None
- assert_num_categories_set_on_discrete_modalities()¶
Raise if any discrete modality has
num_categoriesnot set.- Return type:
None
- map(f, *trees)¶
jax.tree.mapover the registry, treating eachModalityas a leaf.This is the canonical way to map over per-modality data. It bundles three operations together:
Assert that all pytrees are compatible with the registry.
Run
jax.tree.mapwith the registry as the first pytree.Set
is_leafso that eachModalityis treated as a leaf.
- 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_treeis copied to every modality beneath it.prefix_treemay 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; passis_leafto 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:
- Returns:
A pytree with the registry’s modality structure, holding one
prefix_treeleaf per modality.- Raises:
ValueError – If
prefix_treeis not a prefix of the registry’s modality structure.- Return type:
- split_and_project_key(key)¶
Split
keyinto one independent key per modality, projected onto the registry structure.The PRNG-dual of
broadcast():broadcastshares one value across modalities; this hands each its own key. Feed the result tomap().
- 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);Noneon a one-sided interpolant staysNone;Noneon a two-sided interpolant withembedded_source_prioris 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=Nonemeans every leaf is missing.num_samplesis 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 (
Noneleaves are missing).Noneitself 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_sourceandnum_samplesare bothNone; if provided sources disagree on batch size or disagree withnum_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: