Skip to content

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   (: features, : time, : batch), a linear layer maps it to the hidden dimension :

PatchEmbedding (Conv2D for Spatio-Temporal Maps) ​

For spatial grids  , a 2D convolution extracts local patches of size  . The number of patches is   . By setting stride = patch_size, the convolution acts as a linear projection for each patch. For an output feature at spatial patch index on the coarse grid:

Flattening the spatial dimensions   yields the sequence:

PatchUnEmbedding (ConvTranspose) ​

To invert the sequence back to spatial grids, PatchUnEmbedding reshapes the sequence back to    and applies a ConvTranspose operation to upscale the spatial dimensions by . For a pixel in the high-resolution output grid, it maps to coarse grid cell        with internal patch offsets       :


2. Attention Mechanisms (attention.jl) ​

Let sequence length be . Input  .

Self-Attention & Grouped Query Attention (GQA) ​

Self-Attention allows every token to interact with every other token in the sequence. For heads (each of dimension  ), we explicitly define the operation over  :

Where the output of each head is calculated using linearly projected Queries, Keys, and Values:

Note: In GQA, multiple query heads share the same and projections to reduce memory footprint. is an optional causal mask matrix.

Cross-Attention ​

In Cross-Attention, the Queries come from the target sequence ( ), while the Keys and Values are explicitly drawn from the Encoder's memory ( ):

Rotary Position Embedding (RoPE) ​

Instead of adding absolute positional embeddings to , RoPE rotates the query and key vectors in the complex plane based on sequence position . For a pair of features in  , it rotates by angle where  :

Because  , the dot product intrinsically encodes the relative distance between tokens.


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   is a learnable scaling parameter.

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   , is the Hadamard (element-wise) product, and the hidden dimension is scaled as    (default 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  , it is embedded and positional information is added to form the initial layer input :

The data is then processed sequentially through identical 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 (length ) through TransformerBlocks to produce the latent memory:

2. Decoder: Processes the concurrent/future forcings (length ) through CrossAttentionBlocks. The cross-attention layers use for Keys and Values:

The final decoder sequence is projected to the output:


5. Vision Models (vit.jl) ​

VisionTransformer (Scalar / Classification) ​

Takes a spatial map   and maps it to a single global vector  .

  1. Patch Extraction:   , where is the number of patches.

  2. Transformer Stack: The sequence passes through standard TransformerBlocks yielding .

  3. Global Average Pooling: Instead of keeping the full sequence, we average across all spatial patches to form a single vector representation:

  1. 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.

  1. Patch Extraction:   .

  2. Transformer Stack: The sequence is processed, allowing spatial patches to globally attend to one another, yielding .

  3. Unflattening: The sequence   is reshaped back into a spatial patch grid  .

  4. PatchUnEmbedding: A ConvTranspose layer 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  , Model Dimension  , Heads  ,  .

Phase 1: Embeddings (embeddings.jl) ​

Assume our embedded input sequence   (Time 1 and Time 2) is:

julia
# 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.0

Phase 2: RoPE & Self-Attention (attention.jl) ​

Assume projection matrices are Identity matrices, so initially    . We split into   heads.

Head 1 (Top 2 rows):

Head 2 (Bottom 2 rows):

Applying RoPE to Head 1's Queries (): Let's rotate the vectors based on their time positions   and  . Assume for simplicity the base angle results in a () rotation for   and () for  .

  • Time 1 ( ): Vector rotated becomes .

  • Time 2 ( ): Vector rotated becomes .

Rotated Queries:

(Keys would also be rotated identically. For this example, let's assume we proceed with the standard unrotated for the dot product to keep the arithmetic obvious).

Dot Product & Softmax (Head 1):

Scale by   and apply Softmax column-wise:

Multiply by Values :

Head 2 Output ():

Concatenate Heads ():

julia
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 back to the original input (Residual Connection):

RMSNorm then normalizes each column (feature vector per time step) using root-mean-square:

julia
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 simply multiplies everything by 2.

Swish ( ) activates these values, and projects them back to form the final Encoder output . Let's say the final Encoder memory is:

julia
# 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.0

Phase 5: Cross-Attention (encoder_decoder.jl) ​

Now, the Decoder wants to predict Time 3. It receives a concurrent forcing   (length 1):

This is split into   heads (each  ):

  • Head 1:  ,      

  • Head 2:  ,      

Concatenating both heads yields:

julia
# 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   spatial map (   ).

We will extract patches of size  . Total patches  . Model dimension  .

Phase 1: PatchEmbedding (Conv2D) ​

A   Conv2D acts as a linear map from    . Assume our filter matrix   is . In Julia's column-major ordering, the spatial grid is flattened column-by-column:  ,  ,  ,  . Multiplying each scalar pixel by produces our sequence of   tokens, each with dimension  :

(Assume this sequence now passes through the 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   patches:

A final linear head maps this   vector into class logits!

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   back into the   spatial dimensions (  ):

Step 2: ConvTranspose We use a   transposed convolution to map the   features back to   target channels. Assume   is . For each pixel in the final grid, we compute the dot product:

The final map-to-map output grid perfectly retains its spatial integrity:

julia
# 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]