🎮 Reinforcement Learning · Lecture 9 of 21

Function Approximation in RL: From Tables to Neural Networks

Tabular methods cannot scale to large or continuous state spaces. We replace tables with parameterised functions, derive semi-gradient TD, discuss linear features and tile coding, and understand the deadly triad that makes deep RL unstable.

Backgammon has about $10^{20}$ states; Go about $10^{170}$; a robot's joint angles are continuous. A table with one value per state is impossible to store — and useless anyway, because most states are never visited twice. The solution is to generalise: represent the value function with a parameterised function $\hat{v}(s; \mathbf{w})$ or $\hat{q}(s, a; \mathbf{w})$ whose parameters are shared across states, so experience in one state informs predictions in similar states.

The prediction objective#

We want $\hat{v}(s; \mathbf{w}) \approx V^\pi(s)$. With far fewer parameters than states, we cannot be exact everywhere; we minimise a weighted error:

$$ \overline{\text{VE}}(\mathbf{w}) = \sum_s\mu(s)\big[V^\pi(s) - \hat{v}(s; \mathbf{w})\big]^2 $$

where $\mu(s)$ is the fraction of time spent in $s$ (the on-policy distribution). States visited more often get more accuracy.

Gradient Monte Carlo#

With the true return $G_t$ as an unbiased target, stochastic gradient descent gives

$$ \mathbf{w} \leftarrow \mathbf{w} + \alpha\big[G_t - \hat{v}(S_t; \mathbf{w})\big]\nabla_{\mathbf{w}}\hat{v}(S_t; \mathbf{w}) $$

which converges to a local optimum of $\overline{\text{VE}}$.

Semi-gradient TD#

With the TD target $R_{t+1} + \gamma\hat{v}(S_{t+1}; \mathbf{w})$:

$$ \mathbf{w} \leftarrow \mathbf{w} + \alpha\big[R_{t+1} + \gamma\hat{v}(S_{t+1}; \mathbf{w}) - \hat{v}(S_t; \mathbf{w})\big]\nabla_{\mathbf{w}}\hat{v}(S_t; \mathbf{w}) $$

It is called semi-gradient because we differentiate only the prediction, not the target — even though the target also depends on $\mathbf{w}$. This is not true gradient descent on any objective, yet it usually learns faster than Monte Carlo. For linear approximation with on-policy data, semi-gradient TD(0) converges to a fixed point (the "TD fixed point") whose error is within a factor $\frac{1}{1-\gamma}$ of the best achievable.

In code, the semi-gradient is implemented by detaching the target from the computation graph:

python
import torch
import torch.nn as nn

value = nn.Sequential(nn.Linear(4, 64), nn.ReLU(), nn.Linear(64, 1))
opt = torch.optim.Adam(value.parameters(), lr=1e-3)

def td_update(s, r, s_next, done, gamma=0.99):
    with torch.no_grad():                                  # target treated as a constant
        target = r + gamma * (1 - done) * value(s_next).squeeze(-1)
    loss = ((target - value(s).squeeze(-1)) ** 2).mean()
    opt.zero_grad(); loss.backward(); opt.step()
    return loss.item()

s, s2 = torch.randn(32, 4), torch.randn(32, 4)
print(td_update(s, torch.ones(32), s2, torch.zeros(32)))

Linear methods and features#

Linear approximation $\hat{v}(s; \mathbf{w}) = \mathbf{w}^\top\mathbf{x}(s)$ is well understood and stable. Everything depends on features $\mathbf{x}(s)$:

  • Polynomial and Fourier bases for low-dimensional continuous states.
  • Radial basis functions — Gaussian bumps around prototype states.
  • Tile coding (coarse coding): cover the state space with several overlapping, offset grids ("tilings"); each state activates one tile per tiling, giving a sparse binary feature vector. Fast, effective, and a classic choice for problems like Mountain Car.

Deep RL replaces hand-designed features with learned representations: a neural network maps raw inputs (pixels, sensor readings) to values.

Control with function approximation#

Semi-gradient SARSA and Q-learning update $\hat{q}(s, a; \mathbf{w})$ analogously. For discrete actions, a network typically outputs one value per action from the state, so a single forward pass gives all $Q$ values for $\arg\max$.

The deadly triad#

Sutton and Barto identify three ingredients that together can cause instability and divergence:

  1. Function approximation (generalising across states);
  2. Bootstrapping (TD targets built from estimates);
  3. Off-policy training (learning from data not generated by the target policy).

Any two are manageable; all three can make values diverge to infinity even in simple problems (Baird's counterexample demonstrates divergence with linear features). Q-learning with neural networks has all three. The DQN innovations in the next lecture — experience replay and target networks — are engineering responses to exactly this danger. Gradient-TD methods offer provable convergence for linear off-policy learning at some cost in complexity.

Generalisation vs discrimination#

Good representations generalise between states with similar values and discriminate between states with different values. Too coarse, and the agent cannot tell important situations apart; too fine, and it learns slowly. Representation learning is therefore central to deep RL — and auxiliary tasks (predicting future observations, rewards or features) often help agents learn better representations.

JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST · Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

🎮 Reinforcement Learning

Multi-Armed Bandits: The Exploration–Exploitation Dilemma

Bandits are RL without states — the purest form of the exploration problem. We compare ε-greedy, optimistic initialisation, UCB and Thompson sampling, define regret, and look at contextual bandits for real decisions.

Intermediate⏱ 5 min#228
🎮 Reinforcement Learning

Deep Q-Networks (DQN): Human-Level Atari from Pixels

In 2013–2015 DeepMind's DQN learned to play dozens of Atari games from raw pixels with one algorithm. We dissect the Q-network, experience replay, target networks and preprocessing, and implement DQN for CartPole.

Advanced⏱ 5 min#230
🎮 Reinforcement Learning

Q-Learning and SARSA: Model-Free Control

TD methods for control learn action values and improve the policy on the fly. We derive on-policy SARSA and off-policy Q-learning, implement both on the cliff-walking problem, and explain why they learn different paths.

Intermediate⏱ 5 min#227