Layers¶
Reusable low-level building blocks composed by the DiT encoder, decoder, and backbone.
- class stix.nn.layers.SinusoidalFourierFeatures(*args, **kwargs)¶
Bases:
ModuleLearned 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
- class stix.nn.layers.AdaptiveLayerNorm(*args, **kwargs)¶
Bases:
ModuleDiT-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
- class stix.nn.layers.SwiGLUFeedForward(*args, **kwargs)¶
Bases:
ModuleSwiGLU feed-forward network.
Computes: output_proj(silu(gate_proj(x)) * up_proj(x))
- Parameters:
args (Any)
kwargs (Any)
- Return type:
Any
- class stix.nn.layers.MultiHeadSelfAttention(*args, **kwargs)¶
Bases:
ModuleStandard 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
- class stix.nn.layers.TransformerBlock(*args, **kwargs)¶
Bases:
ModuleDiT 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=Nonedegrades to a plain pre-norm transformer block (no AdaLN modulation): shift=0, scale=0, gate=1, i.e.x + attn(LN(x))andx + ffn(LN(x))with no conditioning influence.