Prioritized Experience Replay (PER)
Prioritized Experience Replay replays transitions with high TD errors more frequently, accelerating deep Q-network learning by focusing on surprising events.
Why Does This Exist?
In standard uniform experience replay, every stored transition has an identical probability of selection:
In complex environments, however, transitions vary wildly in their information content. The overwhelming majority of transitions are mundane and predictable—such as driving along an empty highway or waiting in an open hallway—yielding near-zero Temporal Difference (TD) error (). In contrast, rare, critical transitions—such as colliding with an obstacle, reaching a sparse goal, or uncovering a trap—carry massive TD errors ().
Uniform replay treats these transitions as equals. A network wastes over 80% of its mini-batch compute repeatedly re-learning transitions it has already mastered, while critical bottleneck transitions may only be sampled once every tens of thousands of steps.
Introduced by Tom Schaul et al. (DeepMind, 2016), Prioritized Experience Replay (PER) replaces uniform selection by prioritizing transitions in proportion to the magnitude of their expected learning progress, proxied by their TD error . By focusing gradient updates on transitions the current network finds most surprising, PER substantially accelerates training efficiency across challenging deep RL benchmarks.
Because non-uniform sampling fundamentally skews the state visitation distribution, PER introduces Importance Sampling (IS) weights to correct sampling bias, ensuring gradient descent converges to the true expected Bellman fixed point.
Think of It Like This
Flashcard Study with Leitner Boxes
Imagine a medical student studying for board exams using a box of 1,000 flashcards.
If the student practiced using uniform replay, they would shuffle the entire deck randomly every evening and draw 50 cards at random. Most nights, they would waste valuable study hours reviewing elementary concepts they mastered on day one ("What is the primary function of the heart?"), while rare, complex pharmacology interactions they failed yesterday appear only once every three weeks. Learning would be slow and inefficient.
With prioritized replay (the Leitner box system), the student sorts cards into difficulty bins based on surprise and error:
- Box 1 (High Error / Priority): Cards the student got wrong yesterday. These are reviewed every single evening ( is high).
- Box 2 (Moderate Error): Cards answered correctly after hesitation. Reviewed every three days.
- Box 3 (Low Error): Cards answered effortlessly. Reviewed once a month.
Once the student finally masters a difficult card from Box 1, its priority drops, moving it to Box 2 or 3.
The Importance Sampling Correction: If the student spends 90% of their time reviewing rare, acute diseases, they risk developing diagnostic tunnel vision. Importance sampling weights mathematically discount the repeated questions during practice, ensuring the student's overall diagnostic model remains calibrated across the entire medical syllabus.
Where the analogy stops: The Leitner system groups cards into coarse discrete boxes. PER operates on a continuous priority distribution powered by an underlying binary Sum-Tree, updating continuous priority scores dynamically after every mini-batch gradient step.
How It Actually Works
Mathematical Formulation and Sum-Tree Architecture
1. Transition Priority Metric
For transition , its priority is determined by the absolute magnitude of its TD error:
Where:
- is the TD error.
- is a small positive regularization constant (e.g., ) ensuring that transitions with zero TD error still retain a non-zero probability of being replayed.
2. Sampling Probability Distribution
The probability of sampling transition from a buffer of size is defined by:
Where:
- dictates the degree of prioritization.
- When , (pure uniform sampling).
- When , sampling is strictly proportional to priority.
- In practice, balances greediness against diversity.
3. Importance Sampling (IS) Bias Correction
Prioritized sampling shifts the distribution of mini-batches away from the true policy state distribution. To compensate for this bias, each sample in mini-batch is weighted by an importance sampling weight :
Where controls the degree of bias correction. In practice, is annealed linearly from an initial value up to at the end of training. Setting completely neutralizes non-uniform sampling bias.
To ensure stability and keep gradients within standard step sizes, weights are normalized by the maximum weight in the current mini-batch:
The weighted loss minimized by gradient descent is:
4. The Binary Sum-Tree Data Structure
A naive implementation that calculates and samples across items requires an linear pass per sample, making PER computationally intractable.
To achieve sampling and priority updates, PER employs a Binary Sum-Tree:
- The tree is stored as a flat array of size .
- The leaf nodes store the individual priorities .
- Each parent node stores the sum of its two children: .
- The root node (index 0) contains the total sum of all priorities: .
Stratified Sampling in : To draw a mini-batch of size , the range is divided into equal segments. A value is drawn uniformly from each segment:
Starting at the root, the algorithm walks down the tree: if , descend left; otherwise, subtract from and descend right. A leaf is reached in exactly comparisons.
Priority Update in : After computing new TD errors on the sampled mini-batch, each leaf's priority is updated. The change propagates upward through parent pointers to the root in operations.
Worked Numerical Example
Consider a buffer of transitions with , , and .
1. Transition Setup & Sum-Tree Construction
Four transitions arrive with the following TD error magnitudes:
- :
- : (High TD error)
- : (Low TD error)
- :
Total Priority:
Sum-Tree Array Structure ( nodes):
- Leaf nodes (indices 3, 4, 5, 6): , , , .
- Parent of leaves 3 & 4 (index 1): .
- Parent of leaves 5 & 6 (index 2): .
- Root node (index 0): .
- Array: .
2. Sampling Probabilities
- (Sampled 4x more often than )
3. Importance Sampling Weights ()
Raw weight formula: :
- :
- :
- : ()
- :
Normalize by :
Notice how , which is sampled most frequently (), is discounted the most (). This down-weighting prevents frequent updates on from dominating and destabilizing the overall gradient.
4. Post-SGD Priority Update
Suppose the network trains on and its TD error drops from to :
- Delta: .
- Leaf 4 becomes .
- Parent node 1 becomes .
- Root node 0 becomes .
- Updated tree array: .
The priority of drops automatically, naturally redirecting future sampling capacity toward remaining high-error transitions.
Code
import mathimport randomfrom typing import Any, List, Optional, Sequence, Tuple
class SumTree: """Binary Sum-Tree data structure storing priorities in leaves and sums in parents."""
def __init__(self, capacity: int) -> None: if capacity <= 0: raise ValueError("Capacity must be positive.") self.capacity = capacity # Array representation of binary tree: 2 * capacity - 1 nodes self.tree: List[float] = [0.0] * (2 * capacity - 1) self.data: List[Any] = [None] * capacity self.write_ptr: int = 0 self.size: int = 0
@property def total_priority(self) -> float: """Total sum of all priorities stored at the root node.""" return self.tree[0]
def update(self, tree_idx: int, priority: float) -> None: """Update leaf priority and propagate change up to root in O(log N).""" delta = priority - self.tree[tree_idx] self.tree[tree_idx] = priority
parent = (tree_idx - 1) // 2 while parent >= 0: self.tree[parent] += delta if parent == 0: break parent = (parent - 1) // 2
def add(self, priority: float, data: Any) -> int: """Add transition data with priority in O(log N). Overwrites oldest entry.""" leaf_idx = self.write_ptr + self.capacity - 1 self.data[self.write_ptr] = data self.update(leaf_idx, priority)
self.write_ptr = (self.write_ptr + 1) % self.capacity if self.size < self.capacity: self.size += 1 return leaf_idx
def get(self, value: float) -> Tuple[int, float, Any]: """Sample leaf matching prefix sum query value in O(log N).""" parent = 0 while True: left_child = 2 * parent + 1 right_child = left_child + 1
if left_child >= len(self.tree): # Leaf reached leaf_idx = parent break
if value <= self.tree[left_child]: parent = left_child else: value -= self.tree[left_child] parent = right_child
data_idx = leaf_idx - self.capacity + 1 return leaf_idx, self.tree[leaf_idx], self.data[data_idx]
class PrioritizedReplayBuffer: """Prioritized Experience Replay buffer using a SumTree."""
def __init__( self, capacity: int, alpha: float = 0.6, beta_start: float = 0.4, epsilon: float = 1e-5, ) -> None: self.tree = SumTree(capacity) self.alpha = alpha self.beta = beta_start self.epsilon = epsilon self.max_priority = 1.0
def push(self, data: Any) -> None: """Insert new transition with initial maximal priority.""" priority = (self.max_priority + self.epsilon) ** self.alpha self.tree.add(priority, data)
def sample( self, batch_size: int, rng: Optional[random.Random] = None, ) -> Tuple[List[Any], List[int], List[float]]: """Stratified sample of batch_size transitions with IS weights.""" sampler = rng if rng is not None else random batch_data: List[Any] = [] tree_indices: List[int] = [] priorities: List[float] = []
total_p = self.tree.total_priority segment_len = total_p / batch_size
for i in range(batch_size): low = segment_len * i high = segment_len * (i + 1) query_val = sampler.uniform(low, high)
leaf_idx, p, item = self.tree.get(query_val) tree_indices.append(leaf_idx) priorities.append(p) batch_data.append(item)
# Importance sampling weights calculation n = self.tree.size probs = [p / total_p for p in priorities] raw_weights = [(n * prob) ** (-self.beta) for prob in probs] max_w = max(raw_weights) if raw_weights else 1.0 normalized_weights = [w / max_w for w in raw_weights]
return batch_data, tree_indices, normalized_weights
def update_priorities( self, tree_indices: Sequence[int], td_errors: Sequence[float] ) -> None: """Update priorities following network gradient step.""" for idx, error in zip(tree_indices, td_errors): p = (abs(error) + self.epsilon) ** self.alpha self.max_priority = max(self.max_priority, abs(error)) self.tree.update(idx, p)
if __name__ == "__main__": # Reproduce the worked numerical example: N=4, alpha=1.0, beta=0.5, epsilon=0.0 tree = SumTree(capacity=4)
# Insert 4 items with priorities [1.0, 4.0, 0.5, 2.0] errors = [1.0, 4.0, 0.5, 2.0] for i, err in enumerate(errors): tree.add(priority=err, data=f"e{i+1}")
total_priority = tree.total_priority print(f"Total Priority: {total_priority:.4f}") print(f"Tree Array: {[round(x, 2) for x in tree.tree]}")
# Compute sampling probabilities probs = [err / total_priority for err in errors] print(f"Probabilities: {[round(p, 4) for p in probs]}")
# Compute normalized IS weights (beta = 0.5, N = 4) beta = 0.5 n = 4 raw_w = [(n * p) ** (-beta) for p in probs] max_weight = max(raw_w) norm_w = [w / max_weight for w in raw_w] print(f"Normalized IS Weights: {[round(w, 4) for w in norm_w]}")
# Test prefix sum queries query_vals = [0.5, 3.0, 5.2, 6.5] for q in query_vals: leaf_idx, p, item = tree.get(q) print(f"Query {q:3.1f} -> {item} (leaf_idx={leaf_idx}, p={p:.1f})")
# Update e2 priority from 4.0 to 1.0 (leaf_idx = 4) print("\nUpdating e2 (leaf 4) priority from 4.0 to 1.0...") tree.update(4, 1.0) print(f"New Total Priority: {tree.total_priority:.4f}") print(f"Updated Tree Array: {[round(x, 2) for x in tree.tree]}")# Expected Output:# Total Priority: 7.5000# Tree Array: [7.5, 5.0, 2.5, 1.0, 4.0, 0.5, 2.0]# Probabilities: [0.1333, 0.5333, 0.0667, 0.2667]# Normalized IS Weights: [0.7071, 0.3536, 1.0, 0.5]# Query 0.5 -> e1 (leaf_idx=3, p=1.0)# Query 3.0 -> e2 (leaf_idx=4, p=4.0)# Query 5.2 -> e3 (leaf_idx=5, p=0.5)# Query 6.5 -> e4 (leaf_idx=6, p=2.0)# # Updating e2 (leaf 4) priority from 4.0 to 1.0...# New Total Priority: 4.5000# Updated Tree Array: [4.5, 2.0, 2.5, 1.0, 1.0, 0.5, 2.0]Watch Out For
The Zero Priority Lockout and Un-Annealed Beta Trap
Two common implementation traps plague Prioritized Experience Replay:
-
The Zero-Priority Lockout (): If you omit the regularization constant () and a transition has a TD error of exactly zero (), its leaf priority becomes . In a Sum-Tree, a transition with priority has an exact sampling probability of . It will never be replayed again. If policy updates later render that state surprising or valuable, the network will never discover the change because the transition is permanently locked out. The Fix: Always maintain (typically ) to guarantee a non-zero probability floor for all stored transitions.
-
Failing to Anneal : Prioritized sampling introduces significant non-uniform distribution shift. If persists throughout training, the gradient updates remain permanently biased toward high-error regions, distorting the value surface and preventing convergence to the true Bellman fixed point. The Fix: Start with partial compensation () during early exploratory stages, and linearly anneal up to over training so that unbiased convergence is fully restored as the policy stabilizes.
The Quick Version
- Focuses on Surprise: Replaces uniform replay by sampling transitions proportional to their TD error (), accelerating learning where the agent can improve the most.
- Logarithmic Complexity: Employs a binary Sum-Tree to execute non-uniform sampling and priority updates in time, avoiding costly linear scans.
- Importance Sampling Correction: Multiplies gradient steps by normalized IS weights , annealing to eliminate non-uniform sampling bias.
- Prevents Lockout: Adds a small positive constant to ensure zero-error transitions maintain a non-zero probability of future selection.