How to Implement Token Embedding Layers for Language Models in PyTorch
Token embedding layers convert discrete token IDs into continuous vectors using PyTorch's nn.Embedding, which are then combined with positional embeddings to form the input to transformer attention blocks.
Transformer-based language models require a mechanism to transform integer token indices into dense, continuous representations that capture semantic relationships. In the FareedKhan-dev/train-llm-from-scratch repository, this critical transformation occurs within the Transformer class through specialized embedding layers. This guide examines the exact implementation details of token embedding layers for language models, demonstrating how discrete vocabulary indices become differentiable vectors that power modern NLP architectures.
The Architecture of Token and Position Embeddings
Transformer models utilize two distinct embedding mechanisms that operate in parallel. Token embeddings provide semantic meaning to each vocabulary item, while positional embeddings inject sequence order information that self-attention mechanisms inherently lack.
Token Embeddings
In src/models/transformer.py, token embeddings are initialized as a learnable lookup table that maps each vocabulary index to a high-dimensional vector. The implementation uses PyTorch's efficient sparse embedding layer:
# src/models/transformer.py
self.token_embed = nn.Embedding(vocab_size, n_embed)
This layer creates a matrix of shape (vocab_size, n_embed) where each row represents the learned vector for a specific token in the vocabulary. During the forward pass, integer token indices are used to retrieve their corresponding vector rows, resulting in a tensor of shape (batch_size, sequence_length, n_embed).
Positional Embeddings
Since transformer architectures process all tokens simultaneously without inherent sequential bias, they require explicit position information. The repository implements this through a second embedding layer:
# src/models/transformer.py
self.position_embed = nn.Embedding(context_length, n_embed)
Unlike token embeddings that vary based on vocabulary content, positional embeddings depend solely on the token's index within the sequence. The context_length parameter defines the maximum sequence length the model can process.
Combining Embeddings
The final input representation is formed by element-wise addition of token and positional embeddings. In the _pre_attn_pass method, these components are summed to create the transformer block input:
# src/models/transformer.py
tok_embedding = self.token_embed(idx) # shape: (B, T, n_embed)
pos_embedding = self.position_embed(self.pos_idxs[:T]) # shape: (T, n_embed)
return tok_embedding + pos_embedding # combined: (B, T, n_embed)
This combined tensor serves as the input to the stack of transformer blocks defined in src/models/transformer_block.py, each containing multi-head attention and feed-forward layers.
Implementation Details in the Transformer Class
The Transformer class orchestrates the embedding initialization and combination logic. Understanding the initialization parameters and the forward pass flow is essential for implementing custom language models.
Layer Initialization
The embedding layers are instantiated in the Transformer.__init__ method alongside other model components. The constructor accepts vocab_size, n_embed (embedding dimension), and context_length to configure both lookup tables:
# Configuration from src/models/transformer.py
self.token_embed = nn.Embedding(vocab_size, n_embed)
self.position_embed = nn.Embedding(context_length, n_embed)
self.pos_idxs = torch.arange(context_length) # Pre-computed position indices
Pre-computing pos_idxs optimizes the forward pass by avoiding repeated tensor creation for positional lookups.
The _pre_attn_pass Method
The _pre_attn_pass method encapsulates the embedding logic before attention mechanisms process the input. This method handles variable sequence lengths by truncating position indices to the current sequence length T:
def _pre_attn_pass(self, idx):
B, T = idx.size()
tok_embedding = self.token_embed(idx)
pos_embedding = self.position_embed(self.pos_idxs[:T])
return tok_embedding + pos_embedding
The output of this method flows directly into the transformer blocks, where it undergoes layer normalization, multi-head attention, and MLP transformations.
Practical Usage Examples
The repository provides concrete examples demonstrating how to instantiate the model and utilize its embedding layers for training and inference.
Creating Embeddings
To extract embeddings without triggering the full transformer computation, instantiate the model and call the internal embedding method:
import torch
from src.models.transformer import Transformer
# Hyper-parameters
vocab_size = 5000
embed_dim = 256
seq_len = 32
# Random token indices (batch of 2 sequences)
idx = torch.randint(0, vocab_size, (2, seq_len))
model = Transformer(
n_head=8,
n_embed=embed_dim,
context_length=seq_len,
vocab_size=vocab_size,
N_BLOCKS=4,
)
# Get combined token + position embeddings
embeddings, _ = model._pre_attn_pass(idx) # shape: (2, 32, 256)
print(embeddings.shape) # torch.Size([2, 32, 256])
Training Forward Pass
During training, the embeddings feed into the language modeling head to produce vocabulary logits. The standard forward pass handles embedding lookup, transformer processing, and loss calculation:
# Teacher-forcing example with target tokens
logits, loss = model(idx, targets=idx)
print(logits.shape) # (2, 32, 5000) - logits for each vocabulary item
print(loss) # Scalar cross-entropy loss
Text Generation with Learned Embeddings
For inference, the model uses its learned embeddings to process initial tokens and generate subsequent ones autoregressively:
# Start from a single token (e.g., BOS token)
start_idx = torch.tensor([[42]]) # shape: (1, 1)
generated = model.generate(start_idx, max_new_tokens=20)
print(generated.shape) # (1, 21) - original token plus 20 new tokens
Why nn.Embedding is Optimal for Language Models
PyTorch's nn.Embedding class provides three critical advantages for transformer implementations:
- Efficient Sparse Lookup: Only the embedding rows corresponding to input token IDs are accessed, making the operation memory-efficient even with vocabularies of 50,000+ tokens.
- Differentiable Parameters: As a standard
nn.Module, gradients flow back through the embedding matrix during backpropagation, allowing the model to learn contextual representations from training data. - Unified Interface: The same class handles both token and positional embeddings; the only difference is the input domain (vocabulary indices versus position indices).
Summary
- Token embedding layers are implemented in
src/models/transformer.pyusingnn.Embedding(vocab_size, n_embed)to create learnable lookup tables. - Positional embeddings inject sequence order information through a separate
nn.Embedding(context_length, n_embed)layer. - The
_pre_attn_passmethod combines these embeddings via element-wise addition to create the transformer input. - Embedding parameters are trained end-to-end with the rest of the model through standard gradient descent.
- Repository files
scripts/train_transformer.pyandscripts/generate_text.pydemonstrate practical training and inference workflows.
Frequently Asked Questions
What is the difference between token embeddings and word embeddings?
Token embeddings represent subword units or characters processed by the tokenizer, whereas traditional word embeddings map entire words. In the train-llm-from-scratch implementation, the nn.Embedding layer handles arbitrary token IDs, making it agnostic to whether the underlying vocabulary uses words, subwords (BPE), or byte-pair tokens.
Why are positional embeddings added to token embeddings rather than concatenated?
Addition preserves the embedding dimension while injecting position information, making the implementation more parameter-efficient. According to the original transformer architecture implemented in this repository, the sum operation (tok_embedding + pos_embedding) allows the model to learn attention patterns that are sensitive to both semantic content and relative position without doubling the input dimension.
How are embedding weights initialized and updated?
PyTorch initializes nn.Embedding weights from a uniform distribution $\mathcal{U}(-\sqrt{k}, \sqrt{k})$ where $k = 1/\text{embed_dim}$. During training in scripts/train_transformer.py, these parameters are updated automatically by the optimizer alongside attention and MLP weights, allowing the model to refine token representations based on prediction error.
Can token embeddings be shared between the input and output layers?
While the current implementation in src/models/transformer.py uses separate embeddings for input tokens and the final language modeling head, weight tying (using the same nn.Embedding matrix for both) is a common optimization. This reduces parameters and can improve performance, though it requires the embedding dimension to match the vocabulary projection size.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →