Network

class stix.nn.Network(*args, **kwargs)

Bases: Module, ABC

Abstract base defining the network call contract.

A GenerativeModel treats its network as a black box invoked as network(z_t, t, context_data, context_mask, attention_mask) (see get_network_output()). Subclass this and implement __call__() with a compatible signature to plug a network into a generative model.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

abstractmethod __call__(z_t, t, context_data=None, context_mask=None, attention_mask=None)

Map per-modality noisy state z_t at time t to per-modality output.

Parameters:
  • z_t (PyTree[Var]) – Per-modality noisy state in embedded space.

  • t (Time) – Time scalar.

  • context_data (PyTree[Var] | None) – Context conditioning (a pytree) consumed by the network’s context encoder.

  • context_mask (PyTree[Mask | None] | None) – Context masks gating context_data.

  • attention_mask (PyTree[Mask | None] | None) – Per-modality attention masks.

Returns:

Per-modality network output; a bare Var for one-sided, a (forward, reverse) tuple for two-sided.

Return type:

PyTree[Var | tuple[Var, Var]]

class stix.nn.EncoderBackboneDecoderNetwork(*args, **kwargs)

Bases: Network

Wraps instantiated encoders, backbone, and decoders.

Data flow: 1. context_encoder (if any) builds the context vector from t, context_data, and context_mask. 2. Per-modality encoders project inputs into embedding space. 3. Backbone fuses embeddings with time/context, returns per-modality outputs. 4. Per-modality decoders project outputs back to data space with context.

The encoders, decoders and z_t share one per-modality pytree structure.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

__call__(z_t, t, context_data=None, context_mask=None, attention_mask=None)

Forward pass: build context -> encode -> backbone -> decode.

Parameters:
  • z_t (PyTree[Var]) – Per-modality noisy state in embedded space.

  • t (Time) – Time scalar.

  • context_data (PyTree[Var] | None) – Context conditioning (a pytree) forwarded to context_encoder unchanged. Its structure is the context encoder’s contract.

  • context_mask (PyTree[Mask | None] | None) – Context masks forwarded to context_encoder unchanged.

  • attention_mask (PyTree[Mask | None] | None) – Forwarded to the backbone unchanged; semantics are backbone-specific (typically sequence-padding masks).

Returns:

Per-modality network output (same structure as z_t); a single Var for one-sided decoders, or a tuple[Var, Var] for two-sided decoders (e.g. via MultiHeadDecoder).

Return type:

PyTree[Var | tuple[Var, Var]]

class stix.nn.NetworkDimsConfig(*, embedding_dim=128, context_dim=64, ffn_hidden_dim=256, dtype=dtype('float32'))

Bases: BaseModel

Dimensions shared across the DiT encoder, backbone, decoder and context encoder.

Construct once and pass the same instance to each component so the shared widths stay consistent by construction.

Parameters:
  • embedding_dim (int)

  • context_dim (int)

  • ffn_hidden_dim (int)

  • dtype (dtype)

embedding_dim

Embedding width shared by encoders, backbone and decoders.

Type:

int

context_dim

Conditioning-vector width shared by the time context encoder, backbone and decoders.

Type:

int

ffn_hidden_dim

Feed-forward hidden width shared by backbone and decoders.

Type:

int

dtype

Parameter/compute dtype used throughout.

Type:

numpy.dtype

model_config = {'arbitrary_types_allowed': True}

Configuration for the model, should be a dictionary conforming to [ConfigDict][pydantic.config.ConfigDict].

class stix.nn.MultiHeadDecoder(*args, **kwargs)

Bases: Module

Bundles multiple decoder heads sharing one backbone output.

Each head is called with the same (x, context=context) and the results are returned as a tuple, in the order given. Use this when a modality needs more than one prediction (e.g. velocity + score for two-sided stochastic interpolants).

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

__call__(x, context=None)

Call each head with the same (x, context) and return results as a tuple.

Parameters:
  • x (Var)

  • context (Var | None)

Return type:

tuple[Var, …]