Lesson 8.3 · 24 min
Prefix Tuning: Learnable Context Prepended to the Input
Can we teach a frozen language model a new task by training nothing but a handful of invisible “words” placed in front of every input?
In short: Prefix tuning keeps every weight of a pretrained model frozen and learns a short sequence of continuous vectors, the prefix, that is prepended to the keys and values at every attention layer. Real tokens attend to this prefix, which steers the model toward the task. Only around 0.1% of the parameters are trained, and one model can serve many tasks by swapping prefixes.
What is a large language model? (a quick refresher)
A large language model (LLM) is a neural network trained to predict the next token (a word or word piece) given the tokens before it. Most modern LLMs are Transformers: a stack of identical layers, each containing self-attention and a feed-forward network.
Inside self-attention, every token produces three vectors: a query (what am I looking for?), a key (what do I offer?), and a value (the information I pass on). A token compares its query with all keys, turns the scores into weights with softmax, and takes a weighted sum of the values. Prefix tuning works by adding extra keys and values to this mechanism, so keep that picture in mind.
The problem: why full fine-tuning is expensive
Full fine-tuning updates every weight. For each task we pay for gradients and optimizer state on all parameters during training, and we store a complete new copy of the model afterwards. If our company wants ten tasks (summarise tickets, write product descriptions, convert tables to text...), that is ten full models.
Running example for this lesson: our bike-rental shop wants a model that turns a structured booking record such as bike: e-bike | days: 3 | pickup: Station B into a friendly confirmation sentence. This is a classic table-to-text task, the kind prefix tuning was first tested on.
Think of it like briefing a translator We do not retrain an expert translator for every client. We hand them a short brief before each job: “legal tone, British spelling, short sentences”. The expert is unchanged; the brief steers them. Prefix tuning learns the ideal brief, except the brief is written in numbers the model understands directly, not in words.
What is prefix tuning?
Prefix tuning was introduced by Xiang Lisa Li and Percy Liang in 2021. It freezes all model weights and learns a short, task-specific sequence of vectors that is placed before the real input at every layer. The name splits nicely:
- Prefix: a few extra positions in front of the input sequence, like a preface before a book.
- Tuning: those positions hold trainable numbers, adjusted with gradient descent while the model stays fixed.
The model then generates as usual. Because every real token can attend to the prefix positions, the prefix acts like a hidden instruction that shapes every layer’s computation.
How prefix tuning works
From frozen model to task-specific behaviour
- Freeze the model: All pretrained weights are locked. No gradients are stored for them.
- Create the prefix: For each layer, create L trainable key vectors and L trainable value vectors (L is the prefix length, e.g. 10 or 20).
- Prepend to attention: At every layer, the prefix keys and values are placed before the keys and values computed from real tokens.
- Real tokens attend to the prefix: Each real token’s query scores the prefix keys too, so part of its output comes from the prefix values.
- Train on task data: Compute the normal next-token loss on the target text and backpropagate. Gradients flow through the frozen model into the prefix only.
- Save and swap: Store the prefix per task. At serving time, load the prefix for the task the request needs.
The prefix is not real words
A natural question: why not just write a good instruction in text? That is prompt engineering, and the words must come from the vocabulary. Each word maps to a fixed embedding, so we can only choose among existing points in embedding space.
Prefix vectors are continuous: any numbers are allowed, so they can sit between or far away from real word embeddings. This is why the prefix is sometimes called a soft prompt (as opposed to a hard prompt of real tokens). It is far more expressive than any sequence of words, but it is also not human-readable. If we look up the nearest real words to a learned prefix vector, the result is usually meaningless.
Pause and think: Could we print a trained prefix as a sentence and paste it into ChatGPT-style prompts to get the same effect?
No. The prefix is a set of continuous vectors (and, in prefix tuning, per-layer keys and values) that do not correspond to tokens. It only works inside the exact model it was trained with, injected at the activation level.
Where the prefix is added and how it is trained
Where: in prefix tuning the prefix is added at every Transformer layer, as extra key and value activations. This is the key difference from prompt tuning (below), which adds vectors only at the input embedding layer. Deeper placement gives the prefix direct influence on every layer’s attention. For decoder-only models like GPT-2 the prefix goes before the input; for encoder-decoder models like BART the authors added prefixes to both the encoder and decoder.
How: training is ordinary supervised learning. Feed the booking record, compute the cross-entropy loss on the target sentence, backpropagate. Gradients pass through the frozen layers and stop at the prefix, which is the only thing updated.
Below is a self-contained numpy demo of one attention layer. The real tokens’ keys and values are frozen. We train only two prefix positions to steer the layer’s output toward a target vector. We use finite differences for gradients so the code stays short.
prefix_attention.py
import numpy as np
rng = np.random.default_rng(7)
d, n_tok, n_prefix = 8, 5, 2
# Frozen keys/values computed from the real input tokens (one attention layer)
K = rng.normal(size=(n_tok, d))
V = rng.normal(size=(n_tok, d))
q = rng.normal(size=d) # query of the token being generated
target = np.ones(d) # output we want the layer to move toward
# Trainable prefix: 2 virtual key/value vectors (not real words)
P_k = rng.normal(scale=0.1, size=(n_prefix, d))
P_v = rng.normal(scale=0.1, size=(n_prefix, d))
def attend(P_k, P_v):
keys, vals = np.vstack([P_k, K]), np.vstack([P_v, V]) # prefix goes first
s = keys @ q / np.sqrt(d)
w = np.exp(s - s.max()); w /= w.sum()
return w @ vals, w
out, w = attend(P_k, P_v)
print(f"before: weight on prefix {w[:n_prefix].sum():.3f} "
f"distance to target {np.linalg.norm(out - target):.3f}")
lr, eps = 0.5, 1e-5
for step in range(200): # gradient by finite differences
for P in (P_k, P_v):
g = np.zeros_like(P)
for idx in np.ndindex(P.shape):
P[idx] += eps; up = np.linalg.norm(attend(P_k, P_v)[0] - target)
P[idx] -= 2 * eps; dn = np.linalg.norm(attend(P_k, P_v)[0] - target)
P[idx] += eps; g[idx] = (up - dn) / (2 * eps)
P -= lr * g # only the prefix changes
out, w = attend(P_k, P_v)
print(f"after: weight on prefix {w[:n_prefix].sum():.3f} "
f"distance to target {np.linalg.norm(out - target):.3f}")
print("frozen K, V unchanged; trainable numbers:", P_k.size + P_v.size)Output:
before: weight on prefix 0.247 distance to target 2.800 after: weight on prefix 0.560 distance to target 0.010 frozen K, V unchanged; trainable numbers: 32
How small is the prefix really?
Count it: each layer stores L key vectors and L value vectors of size d. So the prefix has layers × 2 × L × d numbers.
- GPT-2 Medium (24 layers, d = 1024, about 345M parameters) with L = 10:
24 × 2 × 10 × 1024 = 491,520, about 0.14% of the model. - A 7B model with 32 layers and d = 4096, L = 20:
32 × 2 × 20 × 4096 = 5,242,880, about 0.07%.
Li and Liang described their prefixes as roughly 0.1% of the model’s parameters, a 1000× reduction in per-task storage compared to saving a fine-tuned copy.
Prefix tuning vs full fine-tuning vs prompt tuning
Prompt tuning (Lester, Al-Rfou and Constant, 2021) is a simpler cousin: it learns soft vectors only at the input embedding layer, not at every layer. It has even fewer parameters, and its authors found it becomes competitive with full fine-tuning mainly for very large models (around 10 billion parameters and up), while lagging on smaller ones. Later work, P-Tuning v2, went back to adding prompts at every layer, essentially prefix tuning, to work well across model sizes.
Pause and think: In the original paper, prefix tuning matched or beat fine-tuning most clearly in which setting: lots of training data or very little?
Very little. Li and Liang reported that prefix tuning was comparable to fine-tuning with full data and tended to outperform it in low-data settings, likely because it changes far fewer parameters and so overfits less.
Advantages, limitations, and where it is used
- Advantage: tiny per-task storage. Save a few hundred thousand to a few million numbers per task instead of a whole model.
- Advantage: modular serving. A single frozen model in memory can mix requests for different tasks in one batch, each with its own prefix.
- Advantage: no forgetting of the base. The model weights never change, so general abilities are untouched when the prefix is removed.
- Limitation: context cost. The prefix occupies L positions in every layer’s attention, slightly increasing compute and reducing usable context.
- Limitation: training sensitivity. Results depend on prefix length, initialisation and learning rate; the reparameterisation MLP helps.
- Limitation: capacity. For tasks needing large changes, LoRA or full fine-tuning usually does better.
Where it is used Prefix tuning was evaluated on table-to-text generation with GPT-2 and on summarisation with BART. Today it appears mostly in research, in PEFT libraries as one option among many, and as the conceptual ancestor of soft-prompt methods. In industry, LoRA has become the more common choice for adapting LLMs.
Common mistake Treating a longer prefix as always better. Very long prefixes add compute and can make training less stable without improving results. Start small (around 10–20) and measure on a validation set.
Worked example, step by step
Let us follow one query through one attention layer by hand, first without a prefix and then with one. We use two real tokens and a single prefix position, with scores small enough to compute on paper. All numbers are illustrative.
One query, with and without a prefix
- Scores without the prefix: The query scores the two real tokens: token 1 gets 1.0, token 2 gets 0.0. Softmax gives
e¹ / (e¹ + e⁰) = 2.72 / 3.72 ≈ 0.73for token 1 and0.27for token 2. - Output without the prefix: The output is the weighted mix of the values:
0.73 · v₁ + 0.27 · v₂. Only real tokens contribute. - Add one prefix position: Training has produced a prefix key that this query scores at 2.0. Now there are three scores: 2.0 (prefix), 1.0 and 0.0.
- Softmax again:
e² = 7.39,e¹ = 2.72,e⁰ = 1. The sum is 11.11. Weights: prefix7.39 / 11.11 ≈ 0.665, token 1≈ 0.245, token 2≈ 0.090. - New output:
0.665 · p_v + 0.245 · v₁ + 0.090 · v₂. Two thirds of this token’s output now comes from the prefix valuep_v, a vector that training chose freely. - What did not change: The model weights, the token keys and the token values are exactly as before. The ratio between the two real tokens is also unchanged:
0.245 / 0.090 ≈ 2.72, the same as0.73 / 0.27.
So a prefix has two separate levers. The prefix key decides how much attention the prefix takes from the real tokens. The prefix value decides what is written into the output with that attention. Training adjusts both, at every layer.
| Position | Score | Weight without prefix | Weight with prefix |
|---|---|---|---|
| Prefix | 2.0 | not present | 0.665 |
| Token 1 | 1.0 | 0.731 | 0.245 |
| Token 2 | 0.0 | 0.269 | 0.090 |
This also shows a failure case. If a prefix key scores very high for every query, attention to the real input collapses and the model starts ignoring what the user wrote. When a prefix-tuned model produces fluent text that does not depend on the input, that is the first thing to suspect.
Practice: try it yourself
We will reproduce the hand calculation in code and then turn a dial: we scale the prefix key from 0 to 2 and watch how much attention the prefix takes and how far the output moves. Nothing is trained here. We are looking at the mechanism itself.
practice_prefix_attention.py
import numpy as np
def softmax(z):
e = np.exp(z - z.max())
return e / e.sum()
q = np.array([1.0, 0.0]) # query of the token being computed
K = np.array([[1.0, 0.0], [0.0, 1.0]]) # frozen keys of 2 real tokens
V = np.array([[1.0, 0.0], [0.0, 1.0]]) # frozen values of 2 real tokens
p_k = np.array([2.0, 0.0]) # learned prefix key
p_v = np.array([0.0, 3.0]) # learned prefix value
w = softmax(K @ q)
print("no prefix weights", w.round(3), " output", (w @ V).round(3))
# Prepend the prefix to keys and values. Scaling its key mimics training.
for scale in [0.0, 0.5, 1.0, 2.0]:
K_all = np.vstack([scale * p_k, K]) # [prefix ; real tokens]
V_all = np.vstack([p_v, V])
w = softmax(K_all @ q) # the query now also sees the prefix
out = w @ V_all # weighted mix of all values
print(f"key x{scale:<3} prefix share {w[0]:.3f} output {out.round(3)}")
# Size of a whole prefix: layers x 2 (key and value) x length x width
layers, length, width = 12, 10, 768
print("prefix size:", layers * 2 * length * width, "numbers")Output:
no prefix weights [0.731 0.269] output [0.731 0.269] key x0.0 prefix share 0.212 output [0.576 0.848] key x0.5 prefix share 0.422 output [0.422 1.422] key x1.0 prefix share 0.665 output [0.245 2.086] key x2.0 prefix share 0.936 output [0.047 2.826] prefix size: 184320 numbers
Look at the row key x0.0. Even a prefix key of all zeros takes 21% of the attention, because a score of 0 still gets a share after softmax.
Now change it:
- Change the query on line 7 to
[0.0, 1.0]. Predict first: will the prefix share atkey x1.0be larger or smaller than 0.665? (Hint: compute the prefix scorep_k · q.) - Set the prefix value on line 11 to
[0.0, 0.0]. Predict what happens to the output as the key grows. Does the prefix still matter? - Change
lengthon line 25 from 10 to 100. Predict the new size, then compare it with a model of about 124 million weights. Is it still a small fraction?
Pause and think: At key x2.0 the prefix takes 93.6% of the attention and the output is [0.047, 2.826]. Why might this be a bad prefix even if the training loss is low?
Almost nothing from the real tokens reaches the output (0.047 in the first dimension, where token 1 used to contribute 0.731). The layer is now driven by the prefix and nearly blind to the input. A prefix like this can fit the training set by producing a typical answer, but it will respond poorly when the input changes. We want the prefix to steer attention, not replace it.
Pause and think: In the code, the prefix is added only to K and V, never to the queries. What would be missing in the output if we also added a prefix query?
Nothing useful would be gained. A query belongs to a position that produces an output. Prefix positions do not need outputs of their own; they only need to be looked at by real tokens. That is why prefix tuning stores just keys and values per layer, which is also where the count layers × 2 × length × width comes from.
Key takeaways
- Prefix tuning freezes the model and learns trainable key/value vectors prepended at every layer.
- Real tokens attend to the prefix, which steers the model like a hidden, learned instruction.
- The prefix is continuous, not words, so it is more expressive than a text prompt but not readable.
- Size is layers × 2 × L × d, typically around 0.1% of the model.
- Prompt tuning acts only at the input layer; LoRA is the more common PEFT choice today.
Key terms
- Prefix: Learned key/value vectors placed before the real tokens at each attention layer.
- Soft prompt: A prompt made of trainable continuous vectors instead of real tokens.
- Prompt tuning: Learning soft vectors only at the input embedding layer of a frozen model.
- Key and value: Vectors each position offers in attention: the key is matched against queries, the value is passed on.
- Reparameterisation: Producing the prefix through a small MLP during training for stability, then keeping only the result.
- Prefix length (L): The number of virtual positions in the prefix.
← 8.2 LoRA: Parameter-Efficient Fine-Tuning via Low-Rank Matrices · 8.4 Knowledge Distillation: Compressing Large Models into Small Ones →