Experience Replay Architectures
Instead of learning directly from experiences as they happen, an agent stores its memories in a buffer and reviews them in random shuffled batches so it does not get trapped in local feedback loops.
Why Does This Exist?
In standard reinforcement learning, an agent collects consecutive transitions sequentially in real time. Training a deep neural network directly on this live stream violates the foundational assumption of stochastic gradient descent: that training samples are independent and identically distributed (i.i.d.).
Consecutive states are strongly autocorrelated: an agent driving a car on a highway observes consecutive video frames that are identical. When gradients from nearly identical states update the network weights consecutively, the model overfits to the immediate local corridor, causing parameter oscillation or catastrophic policy divergence. Furthermore, rare critical events (like hitting an obstacle or collecting a sparse bonus) are experienced once and immediately discarded, causing the network to forget them before sufficient gradient signal propagates backward.
Experience Replay stores transitions into a sliding memory buffer and draws randomized mini-batches for optimization. This simple mechanism breaks temporal correlation, restores i.i.d. sampling assumptions, and allows the agent to reuse rare experiences multiple times across its training lifetime.
Think of It Like This
Studying for a medical licensing exam with flashcards
Imagine a medical student attending emergency room rounds. On Monday, five consecutive patients present with broken ankles; on Tuesday, four consecutive patients arrive with allergic rashes.
If the student studied only the most recent patient in real time, on Monday evening they would prescribe ankle splints for every symptom, and by Tuesday evening they would prescribe antihistamines for broken bones. The sequence of clinical encounters is biased by time and coincidence.
Instead, the student writes each patient's symptoms and diagnosis onto a flashcard and drops it into a study box. Every evening, they shuffle the box thoroughly, pull out twenty random cards from across the past six months, and quiz themselves. Rare cases like venomous snakebites are marked with bright flags and reviewed more frequently than trivial common colds.
How It Actually Works
Temporal Decorrelation and Sum-Tree Prioritization
At each environment step , the agent stores transition tuple into a replay memory of capacity (often to transitions). When capacity is exceeded, the oldest transition is evicted in first-in, first-out (FIFO) order.
In standard Uniform Replay, transitions are sampled uniformly at random . This decorrelates consecutive samples, ensuring within each optimization batch.
However, uniform sampling wastes compute: transitions where the agent already predicts the outcome accurately provide near-zero gradient signal, while rare, highly surprising transitions are sampled no more often than mundane steps.
Prioritized Experience Replay (PER) samples transitions in proportion to their temporal difference (TD) error magnitude :
The probability of sampling transition is defined as:
where priority (with small preventing zero-probability starvation), and controls prioritization strength ( is uniform sampling; is pure greedy prioritization).
Because prioritizing samples alters the visitation frequency of the state distribution, it introduces estimation bias into the Q-learning expected value. PER corrects this bias using Importance Sampling (IS) weights:
where anneals from a small value up to at the end of training. Weights are normalized by so that updates only scale downward.
To maintain sampling and update speed instead of an linear scan, PER implements a binary Sum-Tree:
- Every leaf node stores transition priority .
- Every internal parent node stores the sum of its two child nodes: .
- The root node holds the total sum .
- Sampling partitions into equal intervals and performs binary search down the tree in time.
Worked Example
Trace a mini Sum-Tree with 4 leaf transitions holding priorities: , with , buffer size , and bias correction .
-
Sum-Tree Structure:
- Leaves: , , , .
- Intermediate nodes: , .
- Root node: .
-
Sampling Probabilities:
-
Sample with random query value :
- At Root (), check left child (). Since , descend to right child and subtract left sum: .
- At (), check left leaf (). Since , descend to left child .
- Leaf chosen: index 2 (priority 4.0).
-
Compute Importance Sampling Weight for index 2:
For the minimum probability transition (index 0, ):
Normalized weight:
The high-priority transition update is scaled down by to correct the sampling bias.
Code
import numpy as npfrom typing import Tuple, List
class SumTree: """Binary Sum-Tree storing priorities in O(log N) lookup and update.""" def __init__(self, capacity: int) -> None: self.capacity = capacity # Tree array: parents occupy 1 to capacity-1, leaves occupy capacity to 2*capacity-1 self.tree = np.zeros(2 * capacity, dtype=np.float32) self.data_pointer = 0
def add(self, priority: float) -> int: tree_idx = self.data_pointer + self.capacity self.update(tree_idx, priority) data_idx = self.data_pointer self.data_pointer = (self.data_pointer + 1) % self.capacity return tree_idx
def update(self, tree_idx: int, priority: float) -> None: change = priority - self.tree[tree_idx] self.tree[tree_idx] = priority # Propagate changes up to root parent = tree_idx // 2 while parent >= 1: self.tree[parent] += change parent //= 2
def total_priority(self) -> float: return float(self.tree[1])
def get_leaf(self, value: float) -> Tuple[int, float]: idx = 1 while idx < self.capacity: left = 2 * idx right = left + 1 if value <= self.tree[left]: idx = left else: value -= self.tree[left] idx = right return idx, float(self.tree[idx])
# Demonstrate O(log N) priority samplingtree = SumTree(capacity=4)tree.add(2.0) # idx 4tree.add(6.0) # idx 5tree.add(4.0) # idx 6tree.add(8.0) # idx 7
print(f"Total Priority in Root: {tree.total_priority():.1f}")# Sample with value 11.5leaf_idx, prio = tree.get_leaf(11.5)data_idx = leaf_idx - tree.capacityprint(f"Sampled Leaf Index: {leaf_idx}, Data Index: {data_idx}, Priority: {prio:.1f}")# -> Total Priority in Root: 20.0# -> Sampled Leaf Index: 6, Data Index: 2, Priority: 4.0Watch Out For
Buffer staleness under rapidly evolving target policies
As the Q-network trains over millions of steps, old transitions stored near the bottom of a large buffer () were collected by an exploratory, highly suboptimal policy . If sampled alongside fresh transitions without discount adjustments, off-policy divergence can destabilize deep Q-learning.
To mitigate buffer staleness, dynamically scale the buffer size relative to the learning rate, or implement multi-step targets -step returns with or Retrace correction. When using Prioritized Replay, recalculate each transition's TD error immediately after each backward pass and refresh its tree priority.
The Quick Version
- Eliminates autocorrelation across sequential environment steps, restoring i.i.d. conditions necessary for stable gradient descent.
- Prioritized Experience Replay focuses learning capacity on surprising transitions with large TD errors.
- Sum-Tree binary heap enables sampling and priority updates, while Importance Sampling weights prevent parameter estimation bias.