Layers

Reusable low-level building blocks composed by the DiT encoder, decoder, and backbone.

class stix.nn.layers.SinusoidalFourierFeatures(*args, **kwargs)

Bases: Module

Learned Fourier features: cos(2π(x·w + b)).

Encodes a scalar input into a vector of Fourier features with learned frequencies and phases, initialised from N(0, 1).

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

__call__(x)

Encode a scalar into learned Fourier features.

Parameters:

x (Scalar) – Scalar input.

Returns:

A feature vector of shape (num_features,).

Return type:

Array

class stix.nn.layers.AdaptiveLayerNorm(*args, **kwargs)

Bases: Module

DiT-style Adaptive Layer Norm (AdaLN).

Computes: (1 + scale_proj(norm_c(c))) * norm_x(x) + offset_proj(norm_c(c))

Both norm_x and norm_c are LayerNorm without learned scale/bias. When zero_init=True, scale and offset projections are zero-initialised so the block initially passes through norm_x(x) unchanged.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

__call__(x, context)

Apply AdaLN. x: (L, D), context: (L, C)/(C,)/None → (L, D).

context=None degrades to plain LayerNorm (no scale/shift).

Parameters:
Return type:

Array

class stix.nn.layers.SwiGLUFeedForward(*args, **kwargs)

Bases: Module

SwiGLU feed-forward network.

Computes: output_proj(silu(gate_proj(x)) * up_proj(x))

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

__call__(x)

Forward pass. x: (L, D) → (L, D_out).

Parameters:

x (Array)

Return type:

Array

class stix.nn.layers.MultiHeadSelfAttention(*args, **kwargs)

Bases: Module

Standard multi-head self-attention.

No bias on Q/K projections, bias on V/output projections (standard DiT). Uses jax.nn.dot_product_attention for the attention computation.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

__call__(x, attention_mask=None)

Self-attention. x: (L, D), mask: (L, L) | None → (L, D).

Parameters:
Return type:

Array

class stix.nn.layers.TransformerBlock(*args, **kwargs)

Bases: Module

DiT block with AdaLN-Zero modulation (Peebles & Xie, 2023).

A single zero-initialised MLP produces all 6 modulation parameters (shift_attn, scale_attn, gate_attn, shift_ffn, scale_ffn, gate_ffn) from the conditioning vector. This ensures all modulation starts at zero and co-evolves during training.

Forward: modulate(norm(x)) → Attention → gate → residual

→ modulate(norm(x)) → FFN → gate → residual.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

__call__(x, context, attention_mask=None)

Forward pass. x: (L, D), context: (L, C)/(C,)/None → (L, D).

context=None degrades to a plain pre-norm transformer block (no AdaLN modulation): shift=0, scale=0, gate=1, i.e. x + attn(LN(x)) and x + ffn(LN(x)) with no conditioning influence.

Parameters:
Return type:

Array

class stix.nn.layers.RegressionHead(*args, **kwargs)

Bases: Module

Decoder output head: AdaLN + SwiGLU FFN with zero-init final layer.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

__call__(x, context)

Forward pass. x: (L, D), context: (L, C)/(C,)/None → (L, output_dim).

Parameters:
  • x (Var)

  • context (Var | None)

Return type:

Var