Skip to content
AI360Xpert
Beta

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.

Prioritized Experience Replay uses a binary Sum-Tree to sample transitions proportional to their TD error in O(log N) time, correcting distribution shift with importance sampling weights.
Prioritized Experience Replay uses a binary Sum-Tree to sample transitions proportional to their TD error in O(log N) time, correcting distribution shift with importance sampling weights.

Why Does This Exist?

In standard uniform experience replay, every stored transition has an identical probability of selection:

P(i)=1∣D∣P(i) = \frac{1}{|\mathcal{D}|}

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 (∣δ∣≈0|\delta| \approx 0). In contrast, rare, critical transitions—such as colliding with an obstacle, reaching a sparse goal, or uncovering a trap—carry massive TD errors (∣δ∣≫0|\delta| \gg 0).

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 ∣δi∣|\delta_i|. 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 (P(i)P(i) 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 P(i)∝piαP(i) \propto p_i^\alpha 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 i=(si,ai,ri,si′,di)i = (s_i, a_i, r_i, s'_i, d_i), its priority pip_i is determined by the absolute magnitude of its TD error:

pi≐∣δi∣+ϵp_i \doteq |\delta_i| + \epsilon

Where:

  • δi=ri+γ(1−di)max⁡a′Q(si′,a′;θ−)−Q(si,ai;θ)\delta_i = r_i + \gamma (1 - d_i) \max_{a'} Q(s'_i, a'; \theta^-) - Q(s_i, a_i; \theta) is the TD error.
  • ϵ>0\epsilon > 0 is a small positive regularization constant (e.g., ϵ=10−5\epsilon = 10^{-5}) 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 ii from a buffer of size NN is defined by:

P(i)≐piα∑k=1NpkαP(i) \doteq \frac{p_i^\alpha}{\sum_{k=1}^N p_k^\alpha}

Where:

  • α∈[0,1]\alpha \in [0, 1] dictates the degree of prioritization.
  • When α=0\alpha = 0, P(i)=1NP(i) = \frac{1}{N} (pure uniform sampling).
  • When α=1\alpha = 1, sampling is strictly proportional to priority.
  • In practice, α≈0.6\alpha \approx 0.6 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 ii in mini-batch B\mathcal{B} is weighted by an importance sampling weight wiw_i:

wi≐(1N⋅1P(i))β=(N⋅P(i))−βw_i \doteq \left( \frac{1}{N} \cdot \frac{1}{P(i)} \right)^\beta = (N \cdot P(i))^{-\beta}

Where β∈[0,1]\beta \in [0, 1] controls the degree of bias correction. In practice, β\beta is annealed linearly from an initial value β0≈0.4\beta_0 \approx 0.4 up to 1.01.0 at the end of training. Setting β=1.0\beta = 1.0 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:

wi←wimax⁡k∈Bwkw_i \leftarrow \frac{w_i}{\max_{k \in \mathcal{B}} w_k}

The weighted loss minimized by gradient descent is:

LPER(θ)≐1∣B∣∑i∈Bwi⋅δi2\mathcal{L}_{\text{PER}}(\theta) \doteq \frac{1}{|\mathcal{B}|} \sum_{i \in \mathcal{B}} w_i \cdot \delta_i^2

4. The Binary Sum-Tree Data Structure

A naive implementation that calculates P(i)P(i) and samples across N=106N = 10^6 items requires an O(N)O(N) linear pass per sample, making PER computationally intractable.

To achieve O(log⁡N)O(\log N) sampling and priority updates, PER employs a Binary Sum-Tree:

  • The tree is stored as a flat array of size 2N−12N - 1.
  • The NN leaf nodes store the individual priorities piαp_i^\alpha.
  • Each parent node stores the sum of its two children: Nodek=LeftChild+RightChild\text{Node}_k = \text{LeftChild} + \text{RightChild}.
  • The root node (index 0) contains the total sum of all priorities: ptotal=∑k=1Npkαp_{\text{total}} = \sum_{k=1}^N p_k^\alpha.

Stratified Sampling in O(log⁡N)O(\log N): To draw a mini-batch of size BB, the range [0,ptotal][0, p_{\text{total}}] is divided into BB equal segments. A value vv is drawn uniformly from each segment:

v∼[kBptotal,k+1Bptotal),k∈{0,…,B−1}v \sim \left[ \frac{k}{B} p_{\text{total}}, \frac{k+1}{B} p_{\text{total}} \right), \quad k \in \{0, \dots, B-1\}

Starting at the root, the algorithm walks down the tree: if v≤LeftChildv \le \text{LeftChild}, descend left; otherwise, subtract LeftChild\text{LeftChild} from vv and descend right. A leaf is reached in exactly ⌈log⁡2N⌉\lceil \log_2 N \rceil comparisons.

Priority Update in O(log⁡N)O(\log N): After computing new TD errors on the sampled mini-batch, each leaf's priority is updated. The change Δ=pnew−pold\Delta = p_{\text{new}} - p_{\text{old}} propagates upward through parent pointers to the root in O(log⁡N)O(\log N) operations.


Worked Numerical Example

Consider a buffer of N=4N = 4 transitions with α=1.0\alpha = 1.0, ϵ=0.0\epsilon = 0.0, and β=0.5\beta = 0.5.

1. Transition Setup & Sum-Tree Construction

Four transitions arrive with the following TD error magnitudes:

  • e1e_1: ∣δ1∣=1.0  ⟹  p1=1.0|\delta_1| = 1.0 \implies p_1 = 1.0
  • e2e_2: ∣δ2∣=4.0  ⟹  p2=4.0|\delta_2| = 4.0 \implies p_2 = 4.0 (High TD error)
  • e3e_3: ∣δ3∣=0.5  ⟹  p3=0.5|\delta_3| = 0.5 \implies p_3 = 0.5 (Low TD error)
  • e4e_4: ∣δ4∣=2.0  ⟹  p4=2.0|\delta_4| = 2.0 \implies p_4 = 2.0

Total Priority: ptotal=1.0+4.0+0.5+2.0=7.5p_{\text{total}} = 1.0 + 4.0 + 0.5 + 2.0 = 7.5

Sum-Tree Array Structure (2N−1=72N - 1 = 7 nodes):

  • Leaf nodes (indices 3, 4, 5, 6): tree[3]=1.0\text{tree}[3] = 1.0, tree[4]=4.0\text{tree}[4] = 4.0, tree[5]=0.5\text{tree}[5] = 0.5, tree[6]=2.0\text{tree}[6] = 2.0.
  • Parent of leaves 3 & 4 (index 1): tree[1]=1.0+4.0=5.0\text{tree}[1] = 1.0 + 4.0 = 5.0.
  • Parent of leaves 5 & 6 (index 2): tree[2]=0.5+2.0=2.5\text{tree}[2] = 0.5 + 2.0 = 2.5.
  • Root node (index 0): tree[0]=5.0+2.5=7.5\text{tree}[0] = 5.0 + 2.5 = 7.5.
  • Array: [7.5,5.0,2.5,1.0,4.0,0.5,2.0][7.5, 5.0, 2.5, 1.0, 4.0, 0.5, 2.0].

2. Sampling Probabilities

P(i)=piptotalP(i) = \frac{p_i}{p_{\text{total}}}

  • P(e1)=1.07.5=215≈0.1333P(e_1) = \frac{1.0}{7.5} = \frac{2}{15} \approx 0.1333
  • P(e2)=4.07.5=815≈0.5333P(e_2) = \frac{4.0}{7.5} = \frac{8}{15} \approx 0.5333 (Sampled 4x more often than e1e_1)
  • P(e3)=0.57.5=115≈0.0667P(e_3) = \frac{0.5}{7.5} = \frac{1}{15} \approx 0.0667
  • P(e4)=2.07.5=415≈0.2667P(e_4) = \frac{2.0}{7.5} = \frac{4}{15} \approx 0.2667

3. Importance Sampling Weights (β=0.5,N=4\beta = 0.5, N = 4)

Raw weight formula: wiraw=(4⋅P(i))−0.5=14⋅P(i)w_i^{\text{raw}} = (4 \cdot P(i))^{-0.5} = \frac{1}{\sqrt{4 \cdot P(i)}}:

  • e1e_1: w1raw=(4×0.1333)−0.5=(8/15)−0.5≈1.3693w_1^{\text{raw}} = (4 \times 0.1333)^{-0.5} = (8/15)^{-0.5} \approx 1.3693
  • e2e_2: w2raw=(4×0.5333)−0.5=(32/15)−0.5≈0.6847w_2^{\text{raw}} = (4 \times 0.5333)^{-0.5} = (32/15)^{-0.5} \approx 0.6847
  • e3e_3: w3raw=(4×0.0667)−0.5=(4/15)−0.5≈1.9365w_3^{\text{raw}} = (4 \times 0.0667)^{-0.5} = (4/15)^{-0.5} \approx 1.9365 (wmax⁡w_{\max})
  • e4e_4: w4raw=(4×0.2667)−0.5=(16/15)−0.5≈0.9682w_4^{\text{raw}} = (4 \times 0.2667)^{-0.5} = (16/15)^{-0.5} \approx 0.9682

Normalize by wmax⁡=1.9365w_{\max} = 1.9365:

  • w1=1.36931.9365=48=0.5≈0.7071w_1 = \frac{1.3693}{1.9365} = \sqrt{\frac{4}{8}} = \sqrt{0.5} \approx 0.7071
  • w2=0.68471.9365=432=0.125≈0.3536w_2 = \frac{0.6847}{1.9365} = \sqrt{\frac{4}{32}} = \sqrt{0.125} \approx 0.3536
  • w3=1.93651.9365=1.0000w_3 = \frac{1.9365}{1.9365} = 1.0000
  • w4=0.96821.9365=416=0.5000w_4 = \frac{0.9682}{1.9365} = \sqrt{\frac{4}{16}} = 0.5000

Notice how e2e_2, which is sampled most frequently (P=53.3%P = 53.3\%), is discounted the most (w2=0.3536w_2 = 0.3536). This down-weighting prevents frequent updates on e2e_2 from dominating and destabilizing the overall gradient.

4. Post-SGD Priority Update

Suppose the network trains on e2e_2 and its TD error drops from 4.04.0 to 1.01.0:

  • Delta: Δ=1.0−4.0=−3.0\Delta = 1.0 - 4.0 = -3.0.
  • Leaf 4 becomes 1.01.0.
  • Parent node 1 becomes 5.0+(−3.0)=2.05.0 + (-3.0) = 2.0.
  • Root node 0 becomes 7.5+(−3.0)=4.57.5 + (-3.0) = 4.5.
  • Updated tree array: [4.5,2.0,2.5,1.0,1.0,0.5,2.0][4.5, 2.0, 2.5, 1.0, 1.0, 0.5, 2.0].

The priority of e2e_2 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:

  1. The Zero-Priority Lockout (ϵ=0\epsilon = 0): If you omit the regularization constant (ϵ=0\epsilon = 0) and a transition has a TD error of exactly zero (δ=0\delta = 0), its leaf priority becomes 00. In a Sum-Tree, a transition with priority 00 has an exact sampling probability of 00. 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 ϵ>0\epsilon > 0 (typically ϵ∈[10−6,10−5]\epsilon \in [10^{-6}, 10^{-5}]) to guarantee a non-zero probability floor for all stored transitions.

  2. Failing to Anneal β→1.0\beta \to 1.0: Prioritized sampling introduces significant non-uniform distribution shift. If β<1.0\beta < 1.0 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 (β0=0.4\beta_0 = 0.4) during early exploratory stages, and linearly anneal β\beta up to 1.01.0 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 (P(i)∝(∣δi∣+ϵ)αP(i) \propto (|\delta_i| + \epsilon)^\alpha), 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 O(log⁡N)O(\log N) time, avoiding costly O(N)O(N) linear scans.
  • Importance Sampling Correction: Multiplies gradient steps by normalized IS weights wi=(N⋅P(i))−β/max⁡kwkw_i = (N \cdot P(i))^{-\beta} / \max_k w_k, annealing β→1.0\beta \to 1.0 to eliminate non-uniform sampling bias.
  • Prevents Lockout: Adds a small positive constant ϵ>0\epsilon > 0 to ensure zero-error transitions maintain a non-zero probability of future selection.