Network¶
- class stix.nn.Network(*args, **kwargs)¶
-
Abstract base defining the network call contract.
A
GenerativeModeltreats its network as a black box invoked asnetwork(z_t, t, context_data, context_mask, attention_mask)(seeget_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_tat timetto 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
Varfor one-sided, a(forward, reverse)tuple for two-sided.- Return type:
- class stix.nn.EncoderBackboneDecoderNetwork(*args, **kwargs)¶
Bases:
NetworkWraps instantiated encoders, backbone, and decoders.
Data flow: 1.
context_encoder(if any) builds the context vector fromt,context_data, andcontext_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_tshare 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_encoderunchanged. Its structure is the context encoder’s contract.context_mask (PyTree[Mask | None] | None) – Context masks forwarded to
context_encoderunchanged.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 singleVarfor one-sided decoders, or atuple[Var, Var]for two-sided decoders (e.g. viaMultiHeadDecoder).- Return type:
- class stix.nn.NetworkDimsConfig(*, embedding_dim=128, context_dim=64, ffn_hidden_dim=256, dtype=dtype('float32'))¶
Bases:
BaseModelDimensions 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.
- context_dim¶
Conditioning-vector width shared by the time context encoder, backbone and decoders.
- Type:
Feed-forward hidden width shared by backbone and decoders.
- Type:
- dtype¶
Parameter/compute dtype used throughout.
- Type:
- 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:
ModuleBundles 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