Skip to main content
Create your own

Implementing RNNs and Backpropagation Through Time

Hello! Welcome to the next module in our journey through AI. We've just concluded our deep dive into generative vision models, culminating in an analysis of model safety and censorship. Now, we pivot from static images to dynamic sequences, a fundamental shift that opens up a new world of applications like language translation, time-series forecasting, and even music generation.

This lesson kicks off Module 10: Sequence Modeling with RNNs and Attention. Our goal is to address the learning outcome: Implement a basic Recurrent Neural Network (RNN) cell and apply backpropagation through time (BPTT).

We will start from first principles to understand why traditional neural networks are ill-suited for sequential data. You will then learn the architecture of a basic RNN cell, focusing on its core component: the hidden state that provides it with "memory." Finally, we will tackle the main challenge of training these models by deriving and implementing the Backpropagation Through Time (BPTT) algorithm from scratch. Given your background in Python and AI theory, we'll be implementing this directly in NumPy to ensure a deep, foundational understanding.

1. Why Standard Neural Networks Fail at Sequences

Imagine trying to predict the next word in a sentence. A standard Multi-Layer Perceptron (MLP) would require a fixed-size input. We could try to feed it the last, say, 10 words. But what if the sentence is shorter or longer? How does the model understand that the order of words matters and that the relationship between word 1 and 2 is the same as between word 2 and 3?

Traditional dense networks lack two crucial properties for handling sequences:

  1. They cannot handle variable-length inputs.
  2. They do not share parameters across different positions in the sequence, forcing them to re-learn patterns at every position.

Recurrent Neural Networks were designed specifically to solve these problems.

RNN From Scratch In Python

To see these limitations illustrated and to get a first look at the RNN concept, please watch the beginning of the following video from Dataquest.

Watch from the beginning until 06:22. The video explains why dense networks are inefficient for sequences and introduces the core idea of an RNN: a recurrent hidden step that processes sequence elements one by one, building up a 'memory'.

2. The RNN Cell: A Forward Pass

The "magic" of an RNN lies in its recurrent structure. The network processes a sequence element by element. At each time step , it takes two inputs:

  1. The input for the current time step, .
  2. The hidden state from the previous time step, .

It then computes the new hidden state and an output . This process is repeated for the entire sequence.

Unrolled Recurrent Neural Network (RNN) Architecture
This diagram shows an RNN 'unrolled' in time. Notice how the same weight matrices (U, W, V) are used at every time step. The hidden state `h` acts as a conveyor of information through the sequence.

The core equations for a basic RNN cell are:

  1. Hidden State Calculation:
  2. Output Calculation:

Where:

  • is the input vector at time .
  • is the hidden state vector at time . is typically initialized to a vector of zeros.
  • is the output vector at time .
  • are the weight matrices for the input, recurrent, and output connections, respectively.
  • are bias vectors.
  • and are activation functions. For the hidden state, is typically the hyperbolic tangent (), while depends on the task (e.g., softmax for classification, or no activation for regression).

Why tanh for the Hidden State?

You might recall ReLU being a popular choice for activation functions. However, in RNNs, is often preferred. This is because the repeated application of the recurrent formula () can cause the values in the hidden state to either shrink to zero or grow uncontrollably. The function, which squashes its output to the range , helps to keep the hidden state values bounded, mitigating the exploding gradient problem.

Let's now see how to implement this forward pass in code, including the switch to tanh.

RNN From Scratch In Python

The next segment of the Dataquest video walks through the math and then a step-by-step NumPy implementation of the forward pass. It also explains the motivation for using the tanh activation function.

Watch from 06:22 to 22:49. This section covers: The detailed mathematical operations of the forward pass. A manual, step-by-step coding example. Why tanh is preferred over ReLU for RNNs, with a visual plot. A refactored implementation using a for loop, which is how you'd typically code it.

Test your understanding!

Consider an RNN for a character-level language model. The vocabulary size is 50 characters, and you've chosen a hidden state size of 128. What are the dimensions of the weight matrices , , and ? Assume the input is a one-hot encoded vector representing a character.

Show answer
  • Input : A one-hot vector of shape (50, 1).
  • Hidden state : A vector of shape (128, 1).
  • Weight matrix U (input-to-hidden): To transform (50, 1) into a vector that can be added to the hidden state calculation, must have the shape (128, 50). ().
  • Weight matrix W (hidden-to-hidden): To transform the previous hidden state (128, 1) into a vector of the same size, must have the shape (128, 128). ().
  • Weight matrix V (hidden-to-output): To transform the hidden state (128, 1) into an output vector of size 50 (for a probability distribution over characters), must have the shape (50, 128). ().

3. Training an RNN: Backpropagation Through Time (BPTT)

Training an RNN means updating the shared parameters . This is done using backpropagation, but with a twist. Because the output at time depends on computations from all previous time steps (), the gradient of the loss at time must be propagated "through time" all the way back to the beginning of the sequence. This unrolled version of backpropagation is called Backpropagation Through Time (BPTT).

Backpropagation Through Time in RNNs
This diagram illustrates BPTT. The forward pass (black arrows) computes states and outputs. The backward pass (red arrows) computes gradients. Notice how the gradient at `s2` depends on the error `E2` and also on the gradient flowing back from `s3`.

The key insight is that the total gradient for a shared weight (like ) is the sum of its gradients calculated at every time step.

The Math of BPTT

Let's formalize the gradient flow. The following resource provides a clean derivation of the BPTT equations. Don't worry about the PyTorch code for now; focus on the mathematical formulas.

Backpropagation Through Time (BPTT) tutorial

This document provides a concise mathematical derivation of BPTT. Focus on understanding how the gradients are defined.

Read the text from the beginning until the 'Implementation in PyTorch' section. Pay close attention to the final set of equations under 'So, the complete gradients are'. Notice two key things: The gradients for U, V, and W are sums over all time steps t. The gradient for the hidden state, h_bar_t, has two components: one from the output at the current step (W^T * s_bar_t) and one from the hidden state at the next step (V^T * r_bar_{t+1}). This is the essence of BPTT.

The most crucial equation from that derivation is the one for the gradient with respect to the hidden state :

Here, using the bar notation :

  • is the gradient of the total loss with respect to the hidden state .
  • is the gradient flowing back from the output at the current time step .
  • is the gradient flowing back from the hidden state calculation at the next time step .

This shows that to calculate gradients at step , we need the gradients from step . This is why the backward pass must iterate backward from the last time step to the first.

4. Implementing BPTT from Scratch

With the theory in place, we can now translate it into code. The implementation requires carefully looping backward through the sequence and accumulating the gradients for and .

RNN From Scratch In Python

Let's return to the Dataquest video, which now implements the BPTT logic we just derived. The video will connect the concepts of gradient flow to a concrete NumPy implementation.

Watch from 22:49 to 39:26. This is the most complex part of the lesson. First, the video gives a conceptual overview of the backward pass. Then, it dives into the step-by-step implementation. Pay close attention to how it loops backward and how it calculates and combines the gradient from the output (o_grad) and the gradient from the next hidden state (next_hidden). This directly implements the math we just saw.

The video then continues to show how to assemble the forward and backward passes into a complete training loop to solve a weather prediction problem. While watching the full training loop is optional, it's good context to see how these components fit together in a real application.

Conclusion

In this lesson, we've built a Recurrent Neural Network from the ground up. We started by understanding the limitations of standard networks for sequential data and saw how the RNN's recurrent hidden state provides a form of memory. We implemented the forward pass and then tackled the more complex concept of Backpropagation Through Time (BPTT), deriving the mathematics and translating them into a NumPy implementation.

Key Takeaways:

  • RNNs process sequences by iterating a cell that updates a hidden state, allowing information to persist across time steps.
  • Parameters () are shared across all time steps, making the model efficient and capable of handling variable-length sequences.
  • BPTT is the algorithm for training RNNs. It involves unrolling the network in time and propagating gradients backward from the end of the sequence to the beginning.
  • The gradient at any given time step is influenced by the error at the current output and the gradient propagated from all future time steps.

Preview of the next lesson:
While powerful, the basic RNN we built today suffers from a major weakness: the vanishing gradient problem. As gradients are propagated back through many time steps, they can shrink exponentially, making it difficult for the network to learn long-range dependencies. In our next lesson, we will explore two more advanced RNN architectures designed to solve this: Long Short-Term Memory (LSTM) and Gated Recurrent Units (GRU).

Can't find a good explanation? Sign up and we'll make it for you

Sign up