DiT components

The DiT (diffusion transformer) backbone and its matching encoder, decoder, and time/noise context encoder.

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

Bases: Module

Per-modality encoder: linear projection + learned positional encoding.

Optionally applies input standardisation when beta_fn and gamma_fn are 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

__call__(x, t=None)

Encode input with optional input standardisation.

Parameters:
  • x (Var) – Input array, shape (L, input_dim) or (input_dim,).

  • t (Time | None) – Time value, required when beta_fn and gamma_fn are both set for standardisation.

Return type:

Var

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

Bases: Module

Transformer backbone with AdaLN-Zero conditioning.

Concatenates per-modality encoded inputs along the sequence axis, runs through TransformerBlock layers 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; None degrades each TransformerBlock to plain pre-norm (no AdaLN modulation).

  • attention_mask (PyTree[Mask | None] | None) – Optional per-modality boolean masks. A whole None attends over everything; a None in 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:

PyTree[Var]

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

Bases: Module

Per-modality decoder: RegressionHead with context-dependent AdaLN.

Wraps RegressionHead behind the decoder(x, context=...) call signature that EncoderBackboneDecoderNetwork expects 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:
  • x (Var) – Backbone output for this modality, shape (L, embedding_dim).

  • context (Var | None) – Conditioning vector, or None to fall back to plain AdaLN-less normalisation inside the regression head.

Returns:

Decoded output, shape (output_dim,) if squeeze_sequence else (L, output_dim).

Return type:

Var

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

Bases: Module

Encodes 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

__call__(t)

Encode t into context vector of shape (context_dim,).

Parameters:

t (Time)

Return type:

Array