Course Content
LLMs Deep Dive
10 sections · 40 lessons
How are gradients computed with respect to embeddings?
What you need to know
Lookup as a matrix multiply
With a vocabulary of 5 and width 3, looking up token 2 is:
one_hot(2) = [0, 0, 1, 0, 0]one_hot(2) x E (5 x 3) = row 2 of EThe gradient of the loss with respect to E is one_hot^T × (gradient at the output), which is zero everywhere except row 2.
A small demo
1import torch23emb = torch.nn.Embedding(10, 4) # 10 tokens, 4 dims4ids = torch.tensor([3, 7, 3]) # token 3 used twice5emb(ids).sum().backward()67print(emb.weight.grad[:, 0])8# tensor([0., 0., 0., 2., 0., 0., 0., 1., 0., 0.])Row 3 receives the sum of two gradients (2.0), row 7 one (1.0), and the other eight rows zero. In a real model the incoming gradients are different vectors, but the routing is the same.
Consequences
- Sparse updates — a batch of 4,096 tokens touches perhaps a couple of thousand of the 128,000 rows. Frameworks offer sparse gradients to skip the zeros, mainly useful for very large tables such as recommendation models.
- Rare tokens learn slowly — a token that appears once in a billion gets very few updates. Tokens that exist in the tokenizer but almost never in the training data can end up with poorly trained embeddings and cause strange outputs.
- Optimiser subtlety — with Adam or AdamW, a row with zero gradient this step can still move a little, because of momentum from earlier steps and weight decay.
- Weight tying — if the output projection reuses the embedding table, every row gets a gradient from the output softmax on every step, not just when its token is an input.
Fine-tuning choices
- Freeze embeddings — saves memory and protects pretrained meanings. LoRA leaves them frozen by default.
- Add new tokens — resize the table, initialise new rows (the mean of existing embeddings, or of the token's old subword pieces, works better than random), and train those rows.
A real-life example
A bank adds 40 new tokens for product names and codes (SURAKSHA_FD, YUVA_CARD, internal status codes) so they stop being split into fragments. In the first fine-tune, the new rows start random and the model outputs gibberish whenever one appears, because only 2,000 of the 30,000 training chats contain any of them.
The fix: initialise each new row as the average of the embeddings of its old subword pieces, train the new rows (plus LoRA adapters) at a higher learning rate, and add 3,000 more examples mentioning the products. Now each new row gets enough gradient to settle, and the product names are handled as well as ordinary words.
Follow-up questions to expect
- "Why do rare tokens have worse embeddings?" — Their row only gets gradient when they appear, so they receive far fewer updates.
- "What is weight tying and why use it?" — Sharing the input embedding matrix with the output projection. It saves parameters (hundreds of millions in a big vocabulary) and helps small models; many large models keep them separate.
- "Can you backpropagate into the input text?" — Not into token IDs, which are discrete, but you can take gradients with respect to the embedding vectors; that is how gradient-based adversarial prompt search works.