Encoder Stack Explained
- Input Embedding: Converts tokens into dense vector representations.
- Positional Encoding: Injects sequential order information into embeddings.
- Multi-Head Self-Attention: Each token attends to all other tokens bidirectionally.
- Add & Layer Normalization: Residual connections stabilize learning and preserve gradients.
- Feed Forward Network (FFN): Applies non-linear transformation independently to each token.
- N × Encoder Layers: Repeated stacking deepens contextual understanding.
- Final Contextual Embeddings: Output vectors encode full bidirectional context for downstream tasks.
Minimal BERT Encoder Block (PyTorch)
This simplified block demonstrates bidirectional self-attention with residual connections.
import torch
import torch.nn as nn
class BERTEncoderBlock(nn.Module):
def __init__(self, embed_dim=768, num_heads=12, ff_dim=3072):
super().__init__()
self.attn = nn.MultiheadAttention(embed_dim, num_heads, batch_first=True)
self.ln1 = nn.LayerNorm(embed_dim)
self.ff = nn.Sequential(
nn.Linear(embed_dim, ff_dim),
nn.GELU(),
nn.Linear(ff_dim, embed_dim)
)
self.ln2 = nn.LayerNorm(embed_dim)
def forward(self, x):
attn_output, _ = self.attn(x, x, x)
x = self.ln1(x + attn_output)
ff_output = self.ff(x)
x = self.ln2(x + ff_output)
return x