Typing¶
The typing module defines the shared type vocabulary used throughout stix.
Batch¶
- class stix.typing.data.Batch(raw_batch, context_data=None, context_mask=None, attn_mask=None, loss_mask=None, is_discrete=None)¶
Bases:
PyTreeNodeA batch of multimodal data, consumed by the loss pipeline and solver.
A frozen
flax.structdataclass registered as a JAX pytree, so it flows throughjit/vmap/jax.tree.map.- Parameters:
- raw_batch¶
Per-modality source/target pairs in raw (pre-embedding) space.
sourceisNonefor one-sided (noise-to-data) setups and a data array for two-sided (data-to-data) transport.
- context_data¶
Optional conditioning signal fed to the network’s context encoder (e.g. class labels);
Nonewhen unconditional.- Type:
jaxtyping.PyTree[Var] | None
- context_mask¶
Optional per-modality mask gating
context_data(e.g. zeroing a modality’s contribution);Nonemeans fully conditional.- Type:
jaxtyping.PyTree[Mask | None] | None
- attn_mask¶
Optional per-modality attention mask, e.g. to avoid attending to pad positions in variable-length sequences.
- Type:
jaxtyping.PyTree[Mask | None] | None
- loss_mask¶
Optional per-modality mask zeroing the loss at given positions, e.g. pad locations in variable-length sequences.
- Type:
jaxtyping.PyTree[Mask | None] | None
- is_discrete¶
Optional per-modality discreteness flags, mirroring the modality keys of
raw_batch(Truefor discrete modalities,Falsefor continuous ones). Marked as static pytree metadata (pytree_node=False) so it is not treated as a traced leaf underjit/vmapand can be used for Python-level control flow.Nonewhen not provided.- Type:
jaxtyping.PyTree[bool] | None
- replace(**updates)¶
Returns a new object replacing the specified fields with new values.
Variable types¶
- class stix.typing.variable_type.VarType(*values)¶
Bases:
EnumEnum of prediction target types for the stochastic interpolant.
The types of variables we accept are:
RAW_SOURCERAW_TARGETEMBEDDED_SOURCEEMBEDDED_TARGETNOISEVELOCITY_FIELDVELOCITYSCORE
For more information see the paper here.
Note
This is a vocabulary enum: it is the canonical name for each common prediction target, referenced across the tests, docs, and tutorials. This being said, the user can easily target something with their network beyond these quantities.
Variable type aliases¶
- stix.typing.variables.RawCtsVar = RawCtsVar¶
Type alias.
Type aliases are created through the type statement:
type Alias = int
In this example, Alias and int will be treated equivalently by static type checkers.
At runtime, Alias is an instance of TypeAliasType. The __name__ attribute holds the name of the type alias. The value of the type alias is stored in the __value__ attribute. It is evaluated lazily, so the value is computed only if the attribute is accessed.
Type aliases can also be generic:
type ListOrSet[T] = list[T] | set[T]
In this case, the type parameters of the alias are stored in the __type_params__ attribute.
See PEP 695 for more information.
- stix.typing.variables.RawDisVar = RawDisVar¶
Type alias.
Type aliases are created through the type statement:
type Alias = int
In this example, Alias and int will be treated equivalently by static type checkers.
At runtime, Alias is an instance of TypeAliasType. The __name__ attribute holds the name of the type alias. The value of the type alias is stored in the __value__ attribute. It is evaluated lazily, so the value is computed only if the attribute is accessed.
Type aliases can also be generic:
type ListOrSet[T] = list[T] | set[T]
In this case, the type parameters of the alias are stored in the __type_params__ attribute.
See PEP 695 for more information.
- stix.typing.variables.RawVar = RawVar¶
Type alias.
Type aliases are created through the type statement:
type Alias = int
In this example, Alias and int will be treated equivalently by static type checkers.
At runtime, Alias is an instance of TypeAliasType. The __name__ attribute holds the name of the type alias. The value of the type alias is stored in the __value__ attribute. It is evaluated lazily, so the value is computed only if the attribute is accessed.
Type aliases can also be generic:
type ListOrSet[T] = list[T] | set[T]
In this case, the type parameters of the alias are stored in the __type_params__ attribute.
See PEP 695 for more information.
- stix.typing.variables.EmbeddedVar = EmbeddedVar¶
Type alias.
Type aliases are created through the type statement:
type Alias = int
In this example, Alias and int will be treated equivalently by static type checkers.
At runtime, Alias is an instance of TypeAliasType. The __name__ attribute holds the name of the type alias. The value of the type alias is stored in the __value__ attribute. It is evaluated lazily, so the value is computed only if the attribute is accessed.
Type aliases can also be generic:
type ListOrSet[T] = list[T] | set[T]
In this case, the type parameters of the alias are stored in the __type_params__ attribute.
See PEP 695 for more information.
- stix.typing.variables.NoiseVar = NoiseVar¶
Type alias.
Type aliases are created through the type statement:
type Alias = int
In this example, Alias and int will be treated equivalently by static type checkers.
At runtime, Alias is an instance of TypeAliasType. The __name__ attribute holds the name of the type alias. The value of the type alias is stored in the __value__ attribute. It is evaluated lazily, so the value is computed only if the attribute is accessed.
Type aliases can also be generic:
type ListOrSet[T] = list[T] | set[T]
In this case, the type parameters of the alias are stored in the __type_params__ attribute.
See PEP 695 for more information.
Variable pairs¶
- class stix.typing.variables.RawSourceTargetPair(target, source=None)¶
Bases:
NamedTupleRaw-space source/target pair for one modality.
- target¶
Raw target variable \(x_{\mathrm{tgt}}\).
- source¶
Raw source variable \(x_{\mathrm{src}}\), or
Nonewhen the modality has no source distribution (one-sided setups).- Type:
stix.typing.variables.RawVar | None
- class stix.typing.variables.EmbeddedSourceTargetPair(target, source=None)¶
Bases:
NamedTupleEmbedded-space source/target pair for one modality.
Holds \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\), not automatically \((z_0, z_1)\). See
RawSourceTargetPairfor the \(z_{\mathrm{src}}\) / \(z_0\) distinction.- Parameters:
target (EmbeddedVar)
source (EmbeddedVar | None)
- target¶
Embedded target \(z_{\mathrm{tgt}}\).
- source¶
Embedded source \(z_{\mathrm{src}}\), or
Nonewhen absent (one-sided; then \(z_0 \neq z_{\mathrm{src}}\) because there is no source variable).- Type:
- stix.typing.variables.get_batch_size(raw_batch)¶
Infer the batch size from a raw batch, validated across modalities.
- Parameters:
raw_batch (PyTree[stix.typing.variables.RawSourceTargetPair]) – A pytree whose leaves are
RawSourceTargetPair.- Returns:
The common leading (batch) dimension of each modality’s
targetarray.- Raises:
ValueError – If the batch is empty, or if modalities disagree on the leading (batch) dimension of their target arrays.
- Return type:
Core type aliases¶
Core type aliases shared across the library.
- stix.typing.core.PRNGKeyArray¶
A JAX PRNG key. Semantically distinct from a generic
jax.Array.