DEV Community

Dhruv
Dhruv

Posted on AI-assisted

Building a Transformer from Scratch: Part 1 — The Embedding Layer

In an autoregressive Transformer such as GPT, the work begins before the attention mechanism processes a single tensor. A model cannot operate on raw text, and integer token IDs carry no semantic structure on their own. The embedding layer bridges this gap by converting discrete token IDs into continuous, high-dimensional vectors that encode both the meaning of each token and its position in the sequence.

This article explains how the embedding layer works and presents a clean PyTorch implementation.

The Problem Embeddings Solve

A tokenizer splits raw text into segments and maps each one to an integer index. For example, the word "transformer" might map to 41551 (an illustrative value).

Passing these integers directly into a neural network is problematic, because numerical values imply an ordinal relationship that does not exist. Token 41552 is not "greater than" token 41551 in any meaningful sense.

The solution uses two learned lookup tables:

  • Token embeddings (W_e) map each token ID to a continuous vector of dimension d_model. Through backpropagation, tokens that appear in similar linguistic contexts converge toward similar regions of the vector space.
  • Positional embeddings (W_p) encode word order. Transformers process all tokens in parallel rather than sequentially, so without explicit position information, "dog bites man" and "man bites dog" would be indistinguishable to self-attention. Each position index 0, 1, ..., T-1 is therefore mapped to a learned vector of dimension d_model.

PyTorch Implementation

The module below combines token embeddings with learned positional embeddings, following the architecture used in GPT-2.

import torch
import torch.nn as nn

class TransformerEmbeddings(nn.Module):
    def __init__(self, vocab_size: int, d_model: int, max_seq_len: int):
        super().__init__()
        # Token lookup matrix: shape (vocab_size, d_model)
        self.token_embeddings = nn.Embedding(vocab_size, d_model)

        # Position lookup matrix: shape (max_seq_len, d_model)
        self.position_embeddings = nn.Embedding(max_seq_len, d_model)

    def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
        # Expected input shape: (batch_size, seq_len)
        batch_size, seq_len = input_ids.shape

        # Guard clause: protect against context window overflow
        if seq_len > self.position_embeddings.num_embeddings:
            raise ValueError(
                f"Sequence length ({seq_len}) exceeds maximum context window "
                f"({self.position_embeddings.num_embeddings})"
            )

        # Generate positional indices [0, 1, ..., seq_len - 1]
        # Shape: (1, seq_len), placed on the same device as the input
        positions = torch.arange(0, seq_len, device=input_ids.device).unsqueeze(0)

        # Retrieve representations
        tok_emb = self.token_embeddings(input_ids)      # (batch_size, seq_len, d_model)
        pos_emb = self.position_embeddings(positions)   # (1, seq_len, d_model)

        # Element-wise addition with broadcasting across the batch dimension
        return tok_emb + pos_emb
Enter fullscreen mode Exit fullscreen mode

Key Implementation Details

1. How nn.Embedding Works

nn.Embedding is a trainable weight matrix:

W ∈ R^(vocab_size × d_model)
Enter fullscreen mode Exit fullscreen mode

It is mathematically equivalent to multiplying a one-hot vector by this matrix (x_one_hot @ W), but it is implemented as a direct row lookup. This avoids constructing large, sparse one-hot tensors and removes the associated memory and compute overhead.

2. Device Placement and Broadcasting

  • Device matching. Passing device=input_ids.device ensures that the position indices are created on the same device (CPU, CUDA, or Apple Silicon MPS) as the input. Omitting it defaults to the CPU and raises a device mismatch error when training on a GPU.
  • Broadcasting. The positional embedding tensor has shape (1, seq_len, d_model), while tok_emb has shape (batch_size, seq_len, d_model). When the two are added, PyTorch broadcasts the position vectors across every sequence in the batch.

3. Why Addition Rather Than Concatenation

Concatenating the two vectors would widen each token representation to 2 × d_model, increasing the parameter count and memory usage of every subsequent linear projection.

Element-wise addition preserves the channel dimension at d_model. Conceptually, the positional embedding acts as an offset applied to the token vector, allowing the model to distinguish the same token at different positions.

Next Steps

The combined embeddings leave this layer with shape (batch_size, seq_len, d_model) and feed directly into the first Transformer block. The next article in this series covers the attention mechanism.

Top comments (0)