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: PyTreeNode

A batch of multimodal data, consumed by the loss pipeline and solver.

A frozen flax.struct dataclass registered as a JAX pytree, so it flows through jit / vmap / jax.tree.map.

Parameters:
raw_batch

Per-modality source/target pairs in raw (pre-embedding) space. source is None for one-sided (noise-to-data) setups and a data array for two-sided (data-to-data) transport.

Type:

jaxtyping.PyTree[stix.typing.variables.RawSourceTargetPair]

context_data

Optional conditioning signal fed to the network’s context encoder (e.g. class labels); None when unconditional.

Type:

jaxtyping.PyTree[Var] | None

context_mask

Optional per-modality mask gating context_data (e.g. zeroing a modality’s contribution); None means 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 (True for discrete modalities, False for continuous ones). Marked as static pytree metadata (pytree_node=False) so it is not treated as a traced leaf under jit / vmap and can be used for Python-level control flow. None when 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: Enum

Enum of prediction target types for the stochastic interpolant.

The types of variables we accept are:

  • RAW_SOURCE

  • RAW_TARGET

  • EMBEDDED_SOURCE

  • EMBEDDED_TARGET

  • NOISE

  • VELOCITY_FIELD

  • VELOCITY

  • SCORE

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: NamedTuple

Raw-space source/target pair for one modality.

Parameters:
target

Raw target variable \(x_{\mathrm{tgt}}\).

Type:

stix.typing.variables.RawVar

source

Raw source variable \(x_{\mathrm{src}}\), or None when the modality has no source distribution (one-sided setups).

Type:

stix.typing.variables.RawVar | None

class stix.typing.variables.EmbeddedSourceTargetPair(target, source=None)

Bases: NamedTuple

Embedded-space source/target pair for one modality.

Holds \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\), not automatically \((z_0, z_1)\). See RawSourceTargetPair for the \(z_{\mathrm{src}}\) / \(z_0\) distinction.

Parameters:
target

Embedded target \(z_{\mathrm{tgt}}\).

Type:

stix.typing.variables.EmbeddedVar

source

Embedded source \(z_{\mathrm{src}}\), or None when absent (one-sided; then \(z_0 \neq z_{\mathrm{src}}\) because there is no source variable).

Type:

stix.typing.variables.EmbeddedVar | None

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 target array.

Raises:

ValueError – If the batch is empty, or if modalities disagree on the leading (batch) dimension of their target arrays.

Return type:

int

Core type aliases

Core type aliases shared across the library.

type stix.typing.core.Mask = Annotated[Array, 'boolean mask']

A boolean mask array.

stix.typing.core.PRNGKeyArray

A JAX PRNG key. Semantically distinct from a generic jax.Array.

type stix.typing.core.Scalar = Annotated[Array, 'scalar']

A scalar array.

type stix.typing.core.Shape = Sequence[int]

An array shape as a sequence of dimension sizes.

type stix.typing.core.Time = Annotated[Array, 'scalar time in [0, 1]']

Scalar interpolation time in [0, 1].

type stix.typing.core.Var = Array

A single modality’s variable array (data or network state).