# Implementing Scaled Dot-Product Attention with Key, Query, and Value Projections

> Implement key, query, and value projections for transformer attention. Learn to compute scaled dot-product attention scores and aggregate outputs efficiently.

- Repository: [Fareed Khan/train-llm-from-scratch](https://github.com/FareedKhan-dev/train-llm-from-scratch)
- Tags: how-to-guide
- Published: 2026-05-31

---

**Implement key, query, and value projections using bias-free `nn.Linear` layers, compute scaled dot-product attention scores with causal masking, and aggregate outputs through weighted sums to build multi-head transformer attention mechanisms.**

The `train-llm-from-scratch` repository provides a pure PyTorch implementation of transformer attention mechanisms without relying on high-level abstractions. Learning how to implement attention with key query value projections from scratch reveals the mathematical foundations that power modern large language models.

## Projecting Inputs into Key, Query, and Value Representations

The attention mechanism begins by transforming input embeddings into three distinct representations. In [[`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py)](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py#L29-L31), the `Head` class initializes three separate `nn.Linear` projection layers:

```python
self.key = nn.Linear(n_embed, head_size, bias=False)
self.query = nn.Linear(n_embed, head_size, bias=False)
self.value = nn.Linear(n_embed, head_size, bias=False)

```

**Bias is intentionally disabled** in these projections because the subsequent scaled dot-product calculation operates on relative similarities; additive constants would cancel out during the matrix multiplication and softmax operations. Each layer maps the input embedding dimension (`n_embed`) to the head-specific dimension (`head_size`).

During the forward pass, the input tensor `x` of shape `(B, T, n_embed)`—where **B** represents batch size and **T** represents sequence length—undergoes simultaneous projection:

```python
k = self.key(x)    # (B, T, head_size)

q = self.query(x)  # (B, T, head_size)

v = self.value(x)  # (B, T, head_size)

```

## Computing Scaled Dot-Product Attention

After projection, the implementation calculates attention scores using scaled dot-product similarity. According to the source code in [[`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py)](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py#L49-L51), the operation follows the standard transformer formula:

```python
attn_scores = q @ k.transpose(-2, -1) * (1.0 / math.sqrt(k.size(-1)))

```

**Scaling by the square root of the head size** prevents the dot-product values from growing too large in high-dimensional spaces, which would push the softmax function into regions with extremely small gradients.

To maintain autoregressive properties required for language modeling, the repository applies a **causal mask** using `torch.tril` ([lines 52-54](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py#L52-L54)). This triangular mask ensures tokens only attend to themselves and previous positions, never to future tokens.

## Aggregating Values and Building Multi-Head Output

The attention weights undergo softmax normalization to create a probability distribution ([line 55](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py#L55)), followed by a weighted aggregation with the value vectors ([line 56](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py#L56)):

```python
attn_weights = F.softmax(attn_scores, dim=-1)
out = attn_weights @ v  # (B, T, head_size)

```

This computation occurs independently across multiple heads. The `MultiHeadAttention` class ([line 71](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py#L71)) instantiates `n_head` separate `Head` modules, processing the same input through parallel projections. The outputs are concatenated along the feature dimension and passed through a final linear projection ([line 78](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py#L78)):

```python
class MultiHeadAttention(nn.Module):
    def __init__(self, n_head, n_embed, head_size, dropout):
        super().__init__()
        self.heads = nn.ModuleList([Head(n_embed, head_size, dropout) for _ in range(n_head)])
        self.proj = nn.Linear(n_embed, n_embed)  # Line 78

    
    def forward(self, x):
        out = torch.cat([h(x) for h in self.heads], dim=-1)
        return self.proj(out)

```

## Integrating Attention into Transformer Blocks

The attention mechanism serves as the core component within larger transformer architectures. In [[`src/models/transformer_block.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer_block.py)](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer_block.py), the `MultiHeadAttention` module is combined with layer normalization, residual connections, and feed-forward networks:

```python
from src.models.transformer_block import TransformerBlock
import torch

# Configuration

batch_size, seq_len, embed_dim = 2, 10, 64
n_head = 4
head_size = embed_dim // n_head

# Initialize transformer block

block = TransformerBlock(
    n_head=n_head,
    n_embed=embed_dim,
    head_size=head_size,
    dropout=0.1
)

# Forward pass

x = torch.randn(batch_size, seq_len, embed_dim)
output = block(x)  # Shape: (2, 10, 64)

```

## Summary

- **Key, query, and value projections** are implemented as separate `nn.Linear` layers without bias in the `Head` class to preserve relative similarity calculations.
- **Scaled dot-product attention** divides scores by √head_size to maintain stable gradients during training.
- **Causal masking** via `torch.tril` enforces autoregressive behavior by preventing future token visibility.
- **Multi-head concatenation** combines parallel attention computations through a final linear projection at line 78.
- The complete architecture stacks these attention blocks within `TransformerBlock` modules to build full transformer models.

## Frequently Asked Questions

### Why are biases disabled in the key, query, and value projections?

**Biases are unnecessary** in attention projections because the scaled dot-product operation (`Q @ K^T`) computes relative similarities between vectors. Any additive constant introduced by bias would appear in both the query and key calculations, effectively canceling out during the dot-product operation while adding redundant parameters to the model.

### What purpose does the causal mask serve in the attention implementation?

The **causal mask ensures autoregressive generation** by preventing tokens from attending to future positions in the sequence. Using `torch.tril` to create a lower-triangular matrix of ones, the implementation masks out upper-triangular positions (future tokens) before the softmax operation, which is essential for training language models to predict the next token based solely on previous context.

### How does multi-head attention improve model performance compared to single-head attention?

**Multi-head attention** allows the model to jointly attend to information from different representation subspaces at different positions. By concatenating outputs from `n_head` parallel `Head` instances—each learning distinct attention patterns—the model captures richer contextual relationships than any single attention mechanism could achieve alone, followed by a learned projection matrix that mixes these representations.

### What is the difference between the `Head` and `MultiHeadAttention` classes in this implementation?

The **`Head` class** implements a single attention mechanism with individual key, query, and value projections, computing scaled dot-product attention for one representation subspace. The **`MultiHeadAttention` class** orchestrates multiple `Head` instances (typically 4-16 heads), concatenates their outputs, and applies a final linear projection, effectively combining parallel attention mechanisms into a unified representation.