Transformer Mathematics and Architecture
WARNING
Experimental & Untested Module The Transformer submodule (EasyHybrid.Transformers) is currently experimental and under active development. APIs, layer specifications, and training workflows may change in the future.
This document outlines the fundamental mathematical and tensor operations occurring within the EasyHybrid.jl Transformer architecture, mapping the math directly to the source code files.
1. Embeddings (embeddings.jl)
FeatureEmbedding (Time Series)
Given an input sequence
PatchEmbedding (Conv2D for Spatio-Temporal Maps)
For spatial grids stride = patch_size, the convolution acts as a linear projection for each patch. For an output feature
Flattening the spatial dimensions
PatchUnEmbedding (ConvTranspose)
To invert the sequence back to spatial grids, PatchUnEmbedding reshapes the sequence back to ConvTranspose operation to upscale the spatial dimensions by
2. Attention Mechanisms (attention.jl)
Let sequence length be
Self-Attention & Grouped Query Attention (GQA)
Self-Attention allows every token to interact with every other token in the sequence. For
Where the output of each head
Note: In GQA, multiple query heads share the same
Cross-Attention
In Cross-Attention, the Queries come from the target sequence (
Rotary Position Embedding (RoPE)
Instead of adding absolute positional embeddings to
Because
3. Architecture Blocks (blocks.jl)
The transformer blocks combine Attention and Feed-Forward Networks using Residual Connections and Pre-Normalization (specifically RMSNorm).
Root Mean Square Normalization (RMSNorm)
Before entering any sub-layer, the feature vectors are normalized to stabilize training. For a feature vector
where
Feed Forward Network (SwiGLU FFN)
The FFN acts on each sequence position independently, introducing non-linearity. Following modern architectures (e.g. LLaMA, PaLM; Shazeer 2020), EasyHybrid implements a SwiGLU gated activation block with 3 linear projections without bias:
Gate projection:
Up projection:
Down projection:
Given input
Where multiple_of = 256).
TransformerBlock (Encoder)
The fundamental building block for encoding sequences. It explicitly combines Pre-RMSNorm, Self-Attention, and SwiGLU FFN with residual (skip) connections:
CrossAttentionBlock (Decoder)
The building block for decoders. It includes a third explicit sub-layer specifically to query the Encoder's memory (
4. Sequence Models (transformer.jl & encoder_decoder.jl)
These modules stack the fundamental blocks to create full architectural loops.
TransformerModel (Encoder-Only)
Given a raw sequence
The data is then processed sequentially through TransformerBlocks:
Finally, a layer normalization and linear projection output the desired target features
EncoderDecoderModel (Seq2Seq)
This model splits processing into two distinct streams. 1. Encoder: Processes the historical sequence TransformerBlocks to produce the latent memory:
2. Decoder: Processes the concurrent/future forcings CrossAttentionBlocks. The cross-attention layers use
The final decoder sequence is projected to the output:
5. Vision Models (vit.jl)
VisionTransformer (Scalar / Classification)
Takes a spatial map
Patch Extraction:
, where is the number of patches.Transformer Stack: The sequence passes through
standardTransformerBlocks yielding .Global Average Pooling: Instead of keeping the full sequence, we average across all
spatial patches to form a single vector representation:
- Projection:
.
This is used exclusively for global scalar regression or classification.
VisionToVisionModel (Map-to-Map Regression)
Used for map-to-map regression or spatial forecasting. It retains full spatial topology without pooling.
Patch Extraction:
.Transformer Stack: The sequence is processed, allowing spatial patches to globally attend to one another, yielding
.Unflattening: The sequence
is reshaped back into a spatial patch grid .PatchUnEmbedding: A
ConvTransposelayer expands the spatial dimensions by the patch size , mapping back to the original grid scale with the target channels:
6. Pedagogical Toy Example: The Full Pipeline
Let's walk through an entire Transformer Block and Encoder-Decoder mechanism for a tiny sequence with real numbers.
Setup: Sequence Length
Phase 1: Embeddings (embeddings.jl)
Assume our embedded input sequence
# Input Sequence X (Features: 4, Time: 2)
X = Float32[
1 0;
0 1;
1 1;
0 0
]
# Lux expects (features, sequence_length, batch)
X_lux = reshape(X, 4, 2, 1)4×2×1 Array{Float32, 3}:
[:, :, 1] =
1.0 0.0
0.0 1.0
1.0 1.0
0.0 0.0Phase 2: RoPE & Self-Attention (attention.jl)
Assume projection matrices
Head 1 (Top 2 rows):
Head 2 (Bottom 2 rows):
Applying RoPE to Head 1's Queries (
Time 1 (
): Vector rotated becomes .Time 2 (
): Vector rotated becomes .
Rotated Queries:
(Keys
Dot Product & Softmax (Head 1):
Scale by
Multiply by Values
Head 2 Output (
Concatenate Heads (
using LinearAlgebra, NNlib
using EasyHybrid, EasyHybrid.Transformers, Lux, Random
# 1. Manual Math Execution
# Split into 2 heads (2 features per head)
Q1 = K1 = V1 = X[1:2, :]
Q2 = K2 = V2 = X[3:4, :]
d = 2
# (a) Standard Self-Attention (without RoPE)
S1_unrot = Q1' * K1
S1_unrot_probs = softmax(S1_unrot ./ sqrt(d); dims=1)
A1_unrot = V1 * S1_unrot_probs'
S2 = Q2' * K2
S2_probs = softmax(S2 ./ sqrt(d); dims=1)
A2 = V2 * S2_probs'
A_manual = vcat(A1_unrot, A2)
println("Manual Standard Attention Output (A): \n", round.(A_manual; digits=2))
# (b) Optional: Applying RoPE to Head 1 (90° for m=1, 180° for m=2)
R_90 = [0 -1; 1 0]
R_180 = [-1 0; 0 -1]
Q1_rotated = hcat(R_90 * Q1[:, 1], R_180 * Q1[:, 2])
K1_rotated = hcat(R_90 * K1[:, 1], R_180 * K1[:, 2])
S1_rope = Q1_rotated' * K1_rotated
S1_rope_probs = softmax(S1_rope ./ sqrt(d); dims=1)
A1_rope = V1 * S1_rope_probs'
A_manual_rope = vcat(A1_rope, A2)
println("\nManual Attention with RoPE on Head 1: \n", round.(A_manual_rope; digits=2))
# 2. Library Validation (EasyHybrid)
rng = Random.default_rng()
attention_layer = MultiHeadSelfAttention(4, 2)
ps, st = Lux.setup(rng, attention_layer)
# Force weights to Identity and zero biases to match theoretical equations
ps.query.weight .= Float32.(I(4))
ps.query.bias .= 0.0f0
ps.key.weight .= Float32.(I(4))
ps.value.weight .= Float32.(I(4))
ps.value.bias .= 0.0f0
ps.out.weight .= Float32.(I(4))
ps.out.bias .= 0.0f0
A_lux, _ = attention_layer(X_lux, ps, st)
println("\nLibrary MultiHeadSelfAttention Output: \n", round.(A_lux[:, :, 1]; digits=2))Manual Standard Attention Output (A):
[0.67 0.33; 0.33 0.67; 1.0 1.0; 0.0 0.0]
Manual Attention with RoPE on Head 1:
[0.8 0.2; 0.2 0.8; 1.0 1.0; 0.0 0.0]
Library MultiHeadSelfAttention Output:
Float32[0.67 0.33; 0.33 0.67; 1.0 1.0; 0.0 0.0]Phase 3: Residual & LayerNorm (blocks.jl)
We add the Attention output
RMSNorm then normalizes each column (feature vector per time step) using root-mean-square:
using Statistics
# 1. Manual Math Execution
H_prime = X .+ A_manual
# Manual RMSNorm per column
rms(x) = sqrt(mean(x.^2) + 1e-5)
H_norm_manual = hcat([H_prime[:, i] ./ rms(H_prime[:, i]) for i in 1:2]...)
println("Manual RMSNorm: \n", round.(H_norm_manual; digits=2))
# 2. Library Validation
norm_layer = EasyHybrid.Transformers.RMSNorm(4; eps = 1.0f-5)
ps_n, st_n = Lux.setup(rng, norm_layer)
H_prime_lux = X_lux .+ A_lux
H_norm_lux, _ = norm_layer(H_prime_lux, ps_n, st_n)
println("\nLibrary RMSNorm Output: \n", round.(H_norm_lux[:, :, 1]; digits=2))Manual RMSNorm:
[1.27 0.25; 0.25 1.27; 1.52 1.52; 0.0 0.0]
Library RMSNorm Output:
Float32[1.27 0.25; 0.25 1.27; 1.52 1.52; 0.0 0.0]Phase 4: FFN (blocks.jl)
Assume
Swish (
# FFN: W1 expands, Swish activates, W2 projects back
W1 = 2.0f0 * I(4)
W1_H = W1 * H_prime
swish(x) = x * sigmoid(x)
H_swish = swish.(W1_H)
println("H_swish: \n", round.(H_swish; digits=2))
# Let's explicitly define the mock M_enc from the text for Phase 5
M_enc = Float32[
2 0;
0 2;
1 1;
-1 -1
]4×2 Matrix{Float32}:
2.0 0.0
0.0 2.0
1.0 1.0
-1.0 -1.0Phase 5: Cross-Attention (encoder_decoder.jl)
Now, the Decoder wants to predict Time 3. It receives a concurrent forcing
This is split into
Head 1:
,Head 2:
,
Concatenating both heads yields:
# 1. Manual Math Execution
# Decoder concurrent forcing (length 1)
X_dec = Float32[1; 0; 1; 0]
Q_dec_1 = X_dec[1:2]
K_enc_1 = M_enc[1:2, :]
V_enc_1 = M_enc[1:2, :]
Q_dec_2 = X_dec[3:4]
K_enc_2 = M_enc[3:4, :]
V_enc_2 = M_enc[3:4, :]
S_cross_1 = Q_dec_1' * K_enc_1
S_cross_probs_1 = softmax(S_cross_1 ./ sqrt(d); dims=2)
A_cross_1 = V_enc_1 * S_cross_probs_1'
S_cross_2 = Q_dec_2' * K_enc_2
S_cross_probs_2 = softmax(S_cross_2 ./ sqrt(d); dims=2)
A_cross_2 = V_enc_2 * S_cross_probs_2'
A_cross_man = vcat(A_cross_1, A_cross_2)
println("Manual Final Cross-Attention Output (A_cross): \n", round.(A_cross_man; digits=2))
# 2. Library Validation
cross_attn = MultiHeadSelfAttention(4, 2)
ps_ca, st_ca = Lux.setup(rng, cross_attn)
ps_ca.query.weight .= Float32.(I(4))
ps_ca.query.bias .= 0.0f0
ps_ca.key.weight .= Float32.(I(4))
ps_ca.value.weight .= Float32.(I(4))
ps_ca.value.bias .= 0.0f0
ps_ca.out.weight .= Float32.(I(4))
ps_ca.out.bias .= 0.0f0
X_dec_lux = reshape(X_dec, 4, 1, 1)
M_enc_lux = reshape(M_enc, 4, 2, 1)
out_ca, _ = cross_attn(X_dec_lux, ps_ca, st_ca; context=M_enc_lux)
println("\nLibrary Cross-Attention Output: \n", round.(out_ca[:, :, 1]; digits=2))Manual Final Cross-Attention Output (A_cross):
[1.61; 0.39; 1.0; -1.0;;]
Library Cross-Attention Output:
Float32[1.61; 0.39; 1.0; -1.0;;]7. Pedagogical Toy Example: Vision Models
Let's walk through the mathematical differences between global classification (VisionTransformer) and map-to-map forecasting (VisionToVisionModel).
Setup: We have a single-channel
We will extract patches of size
Phase 1: PatchEmbedding (Conv2D)
A
(Assume this sequence TransformerBlock stack but remains unchanged for this example).
Scenario A: VisionTransformer (Scalar Classification)
For classification, we want to compress this sequence into a single global scalar representation. We use Global Average Pooling (GAP) to collapse the
A final linear head maps this
Scenario B: VisionToVisionModel (Map-to-Map)
For spatial forecasting, we completely skip GAP. We must reconstruct the grid using PatchUnEmbedding.
Step 1: Unflattening We reshape
Step 2: ConvTranspose We use a
The final map-to-map output grid perfectly retains its spatial integrity:
# Vision Models Demonstration
X_grid = Float32[1 2; 3 4] # 2x2 grid
# Patch Embedding equivalent
W_emb = Float32[1; -1]
# Multiply each scalar pixel by W_emb to get D=2 sequence
H_seq = W_emb * reshape(X_grid, 1, 4)
println("\nPatchEmbedding Sequence (H): \n", H_seq)
# Scenario A: GAP for VisionTransformer
z_gap = mean(H_seq; dims=2)
println("Global Average Pooling (z): \n", z_gap)
# Scenario B: ConvTranspose for VisionToVision
W_out = Float32[1 1]
# Multiply W_out by each sequence token and reshape to 2x2
X_hat_flat = W_out * H_seq
X_hat_grid = reshape(X_hat_flat, 2, 2)
println("PatchUnEmbedding Reconstructed Grid (X_hat): \n", X_hat_grid)
PatchEmbedding Sequence (H):
Float32[1.0 3.0 2.0 4.0; -1.0 -3.0 -2.0 -4.0]
Global Average Pooling (z):
Float32[2.5; -2.5;;]
PatchUnEmbedding Reconstructed Grid (X_hat):
Float32[0.0 0.0; 0.0 0.0]