# N4. LSTM and GRU

> How do gates let a network remember for longer?

LLM by Hand · Foundations · side trip: Classic networks · runs in your browser · last part on your computer · interactive page: https://llm.liko.page/learn/lstm/

**Boss challenge.** Write a complete GRU step in NumPy and pass every test in your browser. Then, on your computer, write the LSTM cell inside boss.py yourself and train it on text we generate until at least 90% of the lines it writes follow every rule.

In level N3, a plain RNN could not remember the first symbol for 40 steps. Every step passed its hidden state through tanh and multiplied it by $w_h$,
so what was added early got smaller and smaller. An **LSTM** (the name stands for “long short-term memory”) adds a second path,
the **cell state** c, where the default is to keep what is there.

The boss at the end has a part that runs on your computer. If you have not set that up yet, see
[Run the code on your computer](/setup/).

## 1. The cell state: a moving belt

Think of c as an object on a moving belt (a conveyor belt): it travels from step to step, 
and at each step a little can be removed from it or added to it.

Three **gates** decide what happens to c at each step. Each gate is a number between 0 and 1:

- The **forget gate** f: how much of the old c to keep.
- The **input gate** i: how much of the new candidate g to add.
- The **output gate** o: how much of c to show as the output h.

$$
c_t = f \cdot c_{t-1} + i \cdot g
$$

$$
h_t = o \cdot \tanh(c_t)
$$

*[Interactive lab: Lstm cell — open the page to use it]*

**Question.** cₚᵣₑᵥ = 1.0, f = 0.9, i = 0.5, g = 0.8. What is c = f·cₚᵣₑᵥ + i·g?

*Answer it on the page to check your work.*

**Question.** c = 1.3 and the output gate o = 0.6. Use tanh(1.3) ≈ 0.86. What is h = o·tanh(c)? (2 decimals)

*Answer it on the page to check your work.*

Where do the gates come from? Each one is a small layer with a sigmoid, so its output lands between 0 and 1.
It reads the new input and the previous h, just like the RNN did.

**Question.** The forget gate is f = sigmoid(`w_f`·x + `u_f`·hₚᵣₑᵥ + `b_f`). With `w_f` = 2, x = 1, `u_f` = 0, hₚᵣₑᵥ = 0 and `b_f` = 1, what is f? Use e⁻³ ≈ 0.05. (2 decimals)

*Answer it on the page to check your work.*

## 2. Why the cell keeps its value

Look at the path c takes from one step to the next: multiply by f, add something. There is no tanh and no weight matrix on that path.
If nothing new is added, the cell keeps a fraction f of what it had, every step. Move the f slider in the lab and watch the LSTM curve.

**Question.** Nothing new is added (i = 0) and f = 0.9 at every step. c starts at 1. Use 0.9⁵ ≈ 0.59. What is c after 10 steps? (2 decimals)

*Answer it on the page to check your work.*

**Predict.** Nothing new is added. To keep half of c after 50 steps, the forget gate f must be about…

A. 0.5
B. 0.9
C. 0.99

*Answer it on the page to check your work.*

When the loss at a late step sends its signal back to step 1 along the cell, each step multiplies it by f.
With f close to 1, the signal arrives. In the plain RNN, each step multiplied it by about 0.5.

One detail decides whether training starts to work at all. Before training, a gate’s output is about sigmoid(its bias),
because the weights start small. With the forget bias at 0, f starts at sigmoid(0) = 0.5. With the forget bias at 2,
f starts at sigmoid(2) ≈ 0.88. The lab below trained two LSTMs that differ only in this starting bias.

**Predict.** Two LSTMs differ only in the starting bias of their forget gate: 0 or 2. The task is “remember the first symbol” with k = 80, 1,500 training steps and a 32-number state. A plain RNN only guesses (about 50%). Which LSTM reaches about 100%?

A. Both LSTMs: gates are enough
B. Only the LSTM whose forget bias starts at 2
C. Neither: 80 steps is too long

*Answer it on the page to check your work.*

*[Interactive lab: Memory — open the page to use it]*

**If you are stuck: Why does the starting value of the forget bias matter so much?**

Before training, a gate’s output is about sigmoid(bias). With bias 0, f starts at 0.5, and over 79 steps the cell path keeps
0.5⁷⁹ ≈ 1.65e-24 of the signal. The network cannot even notice that step 1 matters, so it never learns to raise f.

With bias 2, f starts at 0.8808, and 0.8808⁷⁹ ≈ 4.42e-5. That is small, but it is enough for training to see the connection.
Training then pushes f closer to 1 by itself. Starting the forget bias at 1 or 2 is standard practice.

## 3. The GRU: the same idea with two gates

The LSTM carries two vectors, h and c, and uses three gates. A **GRU** (gated recurrent unit) keeps the same “keep by default”
idea with fewer parts: one state h and two gates.

Both gates are sigmoids, between 0 and 1, each computed from $x_t$ and $h_{t-1}$ with its own weights, like the LSTM’s gates:

$$
z = \sigma(x_t W_z + h_{t-1} U_z + b_z) \qquad r = \sigma(x_t W_r + h_{t-1} U_r + b_r)
$$

- The **update gate** z decides how much of the new candidate $\tilde h$ to take. The rest, 1 − z, is the old h, kept as it was:

$$
h_t = (1 - z)\, h_{t-1} + z\, \tilde h
$$

- The **reset gate** r decides how much of the old h to use when making the candidate. The candidate has its own weights $W_h$, $U_h$, $b_h$:

$$
\tilde h = \tanh\bigl(x_t W_h + (r \cdot h_{t-1})\, U_h + b_h\bigr)
$$

Here $r \cdot h_{t-1}$ multiplies number by number, like `r * h` in NumPy. With r near 0, the candidate ignores the old h.

**Question.** A GRU step with hₚᵣₑᵥ = 0.6, update gate z = 0.25 and candidate h̃ = −0.2. What is the new h = (1 − z)·hₚᵣₑᵥ + z·h̃?

*Answer it on the page to check your work.*

*[Interactive lab: Gru — open the page to use it]*

When z is near 0, the old h passes through almost unchanged, like c in the LSTM, so the gradient can also pass back through it.
The GRU has no separate output gate and no separate cell, so it has fewer weights.

**Question.** Inputs have D = 10 numbers and the state has H = 20. Each block of a recurrent layer has a (D, H) input matrix, an (H, H) state matrix and H biases: 10×20 + 20×20 + 20 = 620 numbers. An LSTM has 4 blocks. How many numbers does a GRU have?

*Answer it on the page to check your work.*

## 4. Reading both ways

Every RNN so far read left to right, so the state at a word knows only the words before it. Sometimes the words after
matter more. In “he sat by the bank of the river”, it is “river”, four words later, that says which kind of bank this is (a bank for money, or the side of a river).

A **bidirectional** RNN runs two RNNs over the same sentence: one left to right, one right to left. The output at each
position is the two states side by side, so every word sees the whole sentence, from both ends.

**Question.** The sentence is “he sat by the bank of the river” (8 words, counting from 0, “bank” is word 4). The backward RNN starts at “river”. When it reaches “bank”, how many words has it read, counting “bank”?

*Answer it on the page to check your work.*

*[Interactive lab: Bi rnn — open the page to use it]*

**Question.** Each direction has H = 32. For an 8-word sentence, the output is one row per word, each row the forward and backward states side by side. What shape is it?

*Answer it on the page to check your work.*

**Predict.** Two jobs. (1) Label every word of a finished sentence as noun, verb and so on. (2) Write a sentence one word at a time. Which job can a bidirectional RNN do?

A. Both jobs
B. Only labeling the words
C. Only writing
D. Neither

*Answer it on the page to check your work.*

**If you are stuck: If reading both ways is better, why don’t text generators do it?**

A generator writes one word at a time. When it picks word 5, words 6, 7 and 8 do not exist yet, so there is nothing for a
right-to-left pass to read. Bidirectional models are for jobs where the whole input is there before you answer:
labeling each word, classifying a sentence, or the encoder in level N5, which reads a whole input before the
decoder starts writing.

## 5. Lines that follow the rules

Our own generator writes lines like these:

```
the gray bee wants the gray cup.
the blue fox sees the blue box.
```

The rules: `the <color> <animal> <verb> the <same color> <thing>.` There are 6 colors, 6 animals, 6 verbs and 6 things.
A character-level model reads one letter at a time and learns to predict the next one. To follow the last rule, it has to remember the first
color for 17 characters, across the animal and the verb.

**Question.** There are 6 colors. If a model gets the grammar right but picks the second color at random, what fraction of its lines follow every rule? (percent, 1 decimal)

*Answer it on the page to check your work.*

**Predict.** Guess before you look: the same training (1,800 steps, a 128-number state), but with a plain RNN. What happens to the color rule?

A. It learns it, just later
B. It learns the grammar but stays near 1 in 6 for the colors
C. It learns neither

*Answer it on the page to check your work.*

*[Interactive lab: Lstm samples — open the page to use it]*

The LSTM learns the grammar first: by step 450, 98% of its lines follow the grammar. The colors still differ, at about the 1-in-6 chance level.
Then, between step 900 and step 1,200, it suddenly starts carrying the color across. The plain RNN learns the grammar too, but never the color.

**If you are stuck: Why is there a long flat stretch before the jump?**

Getting the grammar right helps at every character, so it is learned first. The color rule helps at only one character per line,
and only once the network finds a way to carry the color forward. Until some gate starts holding it, the gradient for that rule is weak.
Once a gate starts holding the color, the gain grows quickly, and the fraction of good lines jumps from 17% to 94% in 300 steps.

Gates have limits too. We made the rule harder: the second color comes 40 characters later
(for example `the big <thing> and the <same color> <thing>`). Our LSTM did not learn that version in 6,000 steps. In level 14, attention solves this kind of long-distance rule directly.

## 6. Write it yourself

Here is one LSTM step with vectors. As everywhere in the course, vectors are rows: the input x has D numbers, h and c have H.
All four gates come from one matrix multiply, `x @ W + h @ U + b`, which gives 4H numbers.
The first H belong to the forget gate, the next H to the input gate, then H for the candidate g, then H for the output gate.
So W is (D, 4H) and U is (H, 4H). `z[:H]` takes the first H numbers, `z[H:2 * H]` the next H, and so on.

Write the line that updates the cell.

**Code question.** Write the cell update: keep f of the old c, and add i times the candidate g.

Fill in the blank (`____`):

```python
def sigmoid(z):
    return 1 / (1 + np.exp(-z))

def lstm_step(x, h, c, W, U, b):
    z = x @ W + h @ U + b              # all four gates at once: 4H numbers
    H = h.shape[0]
    f = sigmoid(z[:H])                 # forget
    i = sigmoid(z[H:2 * H])            # input
    g = np.tanh(z[2 * H:3 * H])        # candidate
    o = sigmoid(z[3 * H:])             # output
    c = ____
    h = o * np.tanh(c)
    return h, c

h, c = lstm_step(np.array([0.]), np.zeros(1), np.array([1.]),
                 np.zeros((1, 4)), np.zeros((1, 4)), np.array([2., 0, 1, 0]))
print("c =", c, " h =", h)
```

*Answer it on the page to check your work.*

**Deeper: PyTorch orders the gates differently**

This page and `boss.py` order the four blocks f, i, g, o. PyTorch’s built-in `nn.LSTM` orders them i, f, g, o.
If you ever copy weights between the two, swap the first two blocks of W, U and b. The math is the same; only the order of the
blocks inside the big matrix differs.

### The boss, part 1: a GRU step

Now a cell you have not seen as code. Write one GRU step from the formulas in section 3: the update gate z, the reset gate r,
the candidate $\tilde h$, and the new h. The weights are in a dictionary, so `P["Wz"]` is $W_z$.

**Code question.** The boss, part 1. Write one GRU step: the update gate z, the reset gate r, the candidate h̃ and the new h, all from section 3. Replace ____ with as many lines as you need.

Fill in the blank (`____`):

```python
def sigmoid(z):
    return 1 / (1 + np.exp(-z))

def gru_step(x, h, P):
    # rows: x (D,), h (H,).
    # P holds Wz, Wr, Wh (D, H), Uz, Ur, Uh (H, H) and bz, br, bh (H,)
    ____
    return h_new

P = {k: np.zeros((1, 1)) for k in ["Wz", "Uz", "Wr", "Ur", "Wh", "Uh"]}
P.update({k: np.zeros(1) for k in ["bz", "br", "bh"]})
print(gru_step(np.array([3.]), np.array([0.8]), P))
```

*Answer it on the page to check your work.*

### The boss, part 2: your LSTM cell on your computer

Download [boss.py](/files/lstm/boss.py) into a folder of its own and open it there ([setup](/setup/)). It has everything for training except
one part: `LSTMCell.forward`. Write it in PyTorch, using the step from section 6 (`torch.sigmoid` and `torch.tanh` work like the NumPy ones,
and `x` now holds a whole batch, one row per example, so slice the columns: `z[:, :H]`).

Then run `python boss.py`. It first checks your cell against the formulas on this page, then trains a character model that uses it
for a few minutes and samples 200 lines. At 180 or more good lines it prints a line that starts with `N4 PASS`. Paste that line here.
The code at the end of the line only shows that you pasted it unchanged; this part relies on your honesty.

The level counts as cleared once part 1 (the GRU step) and this line both pass.

**Code question.** The boss, part 2. Paste the PASS line that boss.py printed for your own LSTM cell, between the quotes.

Fill in the blank (`____`):

```python
line = "____"      # it looks like: N4 PASS 0.900 1a2b3c4d
print(line)
```

*Answer it on the page to check your work.*

**Try it: Then try**

Turn your cell into a plain RNN and train again. In `forward`, right after the line that computes `z`, add
`return torch.tanh(z[:, :self.H]), c`. Then put a `#` in front of `check_cell()` in `main`, because this cell fails that check on
purpose. In our run the LSTM cell wrote 199 good lines out of 200, and the plain RNN wrote 6. Read the three lines it prints and
see which rules they break.

## You can now

- Compute one LSTM step by hand: the new cell c from f, i and g, then h from o and tanh(c).
- Explain why the cell keeps its value: the path from c to c is only “multiply by f, add”.
- Write a GRU step in NumPy from its formulas, and an LSTM cell in PyTorch.
