DiT components¶
The DiT (diffusion transformer) backbone and its matching encoder, decoder, and time/noise context encoder.
- class stix.nn.DiTEncoder(*args, **kwargs)¶
Bases:
ModulePer-modality encoder: linear projection + learned positional encoding.
Optionally applies input standardisation when
beta_fnandgamma_fnare provided:z_t = z_t / clip(sqrt(beta_t**2 + gamma_t**2), min=1e-4). Standardisation must happen before the linear projection (the layer has bias, so order matters).- Parameters:
args (Any)
kwargs (Any)
- Return type:
Any
- class stix.nn.DiTBackbone(*args, **kwargs)¶
Bases:
ModuleTransformer backbone with AdaLN-Zero conditioning.
Concatenates per-modality encoded inputs along the sequence axis, runs through
TransformerBlocklayers conditioned on a context vector, then splits back to per-modality outputs.The modalities are fused, so this backbone works on their flattened leaves rather than leaf by leaf, and hands back the structure it was given.
- Parameters:
args (Any)
kwargs (Any)
- Return type:
Any
- __call__(encoded_modalities, t, context, attention_mask=None)¶
Forward pass through transformer backbone.
- Parameters:
encoded_modalities (PyTree[Var]) – Per-modality encoded inputs (L_k, embedding_dim).
t (Time) – Time value in [0, 1] (unused; conditioning is via context).
context (Var | None) – Conditioning vector from the context encoder;
Nonedegrades eachTransformerBlockto plain pre-norm (no AdaLN modulation).attention_mask (PyTree[Mask | None] | None) – Optional per-modality boolean masks. A whole
Noneattends over everything; aNonein place of one modality’s mask leaves that modality fully attended to.
- Returns:
Per-modality embeddings (same structure that was passed in).
- Raises:
ValueError – If the encoded token counts disagree with
modality_num_tokens.- Return type:
- class stix.nn.DiTDecoder(*args, **kwargs)¶
Bases:
ModulePer-modality decoder: RegressionHead with context-dependent AdaLN.
Wraps RegressionHead behind the
decoder(x, context=...)call signature thatEncoderBackboneDecoderNetworkexpects from every per-modality decoder. Optionally squeezes the sequence dimension for single-token modalities.- Parameters:
args (Any)
kwargs (Any)
- Return type:
Any
- __call__(x, context)¶
Decode with context-dependent AdaLN.
- Parameters:
- Returns:
Decoded output, shape
(output_dim,)ifsqueeze_sequenceelse(L, output_dim).- Return type:
- class stix.nn.TimeNoiseContextEncoder(*args, **kwargs)¶
Bases:
ModuleEncodes time t and noise level into a context vector via Fourier features.
Pipeline: [t, 0.25·log(clamp(gamma_fn(t)))] → separate Fourier features → concatenate → Linear → SiLU → Linear → context vector.
Encoding both t directly and log(gamma_t) ensures the model can distinguish symmetric time points (e.g. t=0.2 vs t=0.8 with gamma_fn=sqrt(2t(1-t))), which is critical for two-sided stochastic interpolants.
NOTE: Assumes gamma_fn(t) > 0 for all t in (0, 1). Values below 1e-6 are clamped before log, producing a floor at 0.25 * log(1e-6) ≈ −3.45.
- Parameters:
args (Any)
kwargs (Any)
- Return type:
Any