Skip to content
AI360Xpert
Beta

Semi-Gradient SARSA in RL

Semi-gradient SARSA extends on-policy temporal-difference control to continuous state spaces by approximating action-values with parameterized function approximators. By updating weights using sampled (S, A, R, S', A') transitions, agents learn optimal policies while generalizing across continuous dimensions.

Semi-gradient SARSA updates parameterized action-values on continuous transitions with detached bootstrap gradients.
Semi-gradient SARSA updates parameterized action-values on continuous transitions with detached bootstrap gradients.

Why Does This Exist?

In discrete tabular reinforcement learning, SARSA evaluates and improves policies by maintaining a separate table cell Q(s,a)Q(s, a) for every state-action pair. However, real-world control domains—such as balancing a pole on a cart, driving a car up a mountain, or controlling a multi-jointed robotic manipulator—operate in continuous state spaces.

In continuous domains, tabular methods collapse completely:

  1. Infinite State Cardinality: Continuous variables (like velocity θ˙∈R\dot{\theta} \in \mathbb{R} or position x∈Rx \in \mathbb{R}) mean the state space S\mathcal{S} is uncountably infinite. A lookup table cannot even be allocated in memory.
  2. Control Demands Action-Values: While state-value methods like semi-gradient TD(0) can estimate V(s)V(s) in continuous spaces, state values alone are insufficient for model-free control. Without an environmental transition model p(s′,r∣s,a)p(s', r \mid s, a), an agent in state ss cannot determine which action aa leads to the highest-value successor states. The agent must evaluate action-values q(s,a)q(s, a).

Semi-Gradient SARSA solves the continuous control challenge by extending function approximation directly to action-values:

q^(s,a,w)≈qπ(s,a)\hat{q}(s, a, \mathbf{w}) \approx q_\pi(s, a)

where w∈Rd\mathbf{w} \in \mathbb{R}^d is a shared parameter vector. By tracking sequential on-policy transitions (St,At,Rt+1,St+1,At+1)(S_t, A_t, R_{t+1}, S_{t+1}, A_{t+1}), the algorithm updates its action-value surface online, enabling generalized policy evaluation and ε\varepsilon-greedy policy improvement across continuous state landscapes.

Think of It Like This

A Racecar Driver Tuning Cornering Technique on an Unfamiliar Track

Imagine an expert racecar driver learning how to navigate an unfamiliar racecourse with a high-performance vehicle:

  • The Continuous State Space: At any instant mid-corner, the car’s state is defined by continuous physical telemetry: current vehicle speed (e.g., 84.3 mph), lateral g-force, steering angle, and distance to the apex. The driver does not maintain an exhaustive spreadsheet of every possible decimal speed and steering angle.
  • Action-Value Evaluation (q(s,a)q(s, a)): The driver considers distinct tactical inputs: tapping the brakes (A1A_1), feathering the throttle (A2A_2), or holding steady (A3A_3). The driver's mental model predicts the quality of taking throttle action A2A_2 given current telemetry SS.
  • The On-Policy SARSA Transition:
    1. Entering the turn at state StS_t, the driver chooses action AtA_t (feather throttle).
    2. The car reacts: tires grip, generating immediate reward Rt+1R_{t+1} (forward velocity without sliding), carrying the car to new telemetry state St+1S_{t+1}.
    3. Looking ahead to the next section of asphalt, the driver selects their next intended maneuver At+1A_{t+1} (shift to full throttle for corner exit).
    4. The driver updates the value of their initial feather-throttle decision based on the immediate grip (Rt+1R_{t+1}) combined with the anticipated exit speed of the actual next maneuver (γq^(St+1,At+1)\gamma \hat{q}(S_{t+1}, A_{t+1})).
  • The Semi-Gradient Aspect: The driver treats the projected exit speed as a fixed benchmark target during this adjustment. They do not attempt to differentiate the entire mathematical vehicle physics model backward through time.

Where the analogy breaks down: A human driver’s mental model relies on biological neural networks and kinesthetic intuition. In linear Semi-Gradient SARSA, the action-value function is parameterized by an explicit linear feature stack wTx(s,a)\mathbf{w}^T \mathbf{x}(s, a) optimized via stochastic gradient descent over an episodic Markov decision process.

How It Actually Works

Parameterized Action-Values and Episodic Semi-Gradient Control

Semi-gradient SARSA combines function approximation with the on-policy temporal-difference control loop.

1. Action-Value Function Approximation

The action-value function is approximated by a parameterized function:

q^(s,a,w)≈qπ(s,a)\hat{q}(s, a, \mathbf{w}) \approx q_\pi(s, a)

In the classical linear case, the approximation is the inner product of parameter vector w∈Rd\mathbf{w} \in \mathbb{R}^d and a joint state-action feature vector x(s,a)∈Rd\mathbf{x}(s, a) \in \mathbb{R}^d:

q^(s,a,w)≐wTx(s,a)=∑i=1dwixi(s,a)\hat{q}(s, a, \mathbf{w}) \doteq \mathbf{w}^T \mathbf{x}(s, a) = \sum_{i=1}^d w_i x_i(s, a)

2. Constructing State-Action Features via Action-Stacking

For discrete actions A={a1,a2,…,am}\mathcal{A} = \{a_1, a_2, \dots, a_m\} and continuous state features x(s)∈Rk\mathbf{x}(s) \in \mathbb{R}^k, the canonical representation uses action-stacking (block coding). The global feature vector has dimension d=k×md = k \times m:

x(s,a1)=[x(s)0⋮0],x(s,a2)=[0x(s)⋮0],…,x(s,am)=[00⋮x(s)]\mathbf{x}(s, a_1) = \begin{bmatrix} \mathbf{x}(s) \\ \mathbf{0} \\ \vdots \\ \mathbf{0} \end{bmatrix}, \quad \mathbf{x}(s, a_2) = \begin{bmatrix} \mathbf{0} \\ \mathbf{x}(s) \\ \vdots \\ \mathbf{0} \end{bmatrix}, \quad \dots, \quad \mathbf{x}(s, a_m) = \begin{bmatrix} \mathbf{0} \\ \mathbf{0} \\ \vdots \\ \mathbf{x}(s) \end{bmatrix}

This structure guarantees that each action maintains its own dedicated subspace within the weight vector w=[wa1T,wa2T,…,wamT]T\mathbf{w} = [\mathbf{w}_{a_1}^T, \mathbf{w}_{a_2}^T, \dots, \mathbf{w}_{a_m}^T]^T, allowing independent action-value surfaces to evolve without destructive cross-action parameter interference.

3. The On-Policy Semi-Gradient SARSA Update

At time step tt, the agent visits continuous state StS_t, executes action AtA_t, observes scalar reward Rt+1R_{t+1}, and transitions to St+1S_{t+1}. It then samples its next action At+1A_{t+1} using its behavior policy (such as ε\varepsilon-greedy with respect to q^(⋅,⋅,wt)\hat{q}(\cdot, \cdot, \mathbf{w}_t)).

The one-step on-policy SARSA target is:

Ut≐Rt+1+γq^(St+1,At+1,wt)U_t \doteq R_{t+1} + \gamma \hat{q}(S_{t+1}, A_{t+1}, \mathbf{w}_t)

where γ∈[0,1]\gamma \in [0, 1] is the discount factor. If St+1S_{t+1} is a terminal state, the target reduces strictly to Ut≐Rt+1U_t \doteq R_{t+1}.

The semi-gradient SGD update rule shifts parameters in the direction of the error gradient:

wt+1=wt+α[Ut−q^(St,At,wt)]∇wq^(St,At,wt)\mathbf{w}_{t+1} = \mathbf{w}_t + \alpha \left[ U_t - \hat{q}(S_t, A_t, \mathbf{w}_t) \right] \nabla_{\mathbf{w}} \hat{q}(S_t, A_t, \mathbf{w}_t)

For linear models, ∇wq^(St,At,wt)=x(St,At)\nabla_{\mathbf{w}} \hat{q}(S_t, A_t, \mathbf{w}_t) = \mathbf{x}(S_t, A_t), yielding:

wt+1=wt+α[Rt+1+γq^(St+1,At+1,wt)−q^(St,At,wt)]x(St,At)\mathbf{w}_{t+1} = \mathbf{w}_t + \alpha \left[ R_{t+1} + \gamma \hat{q}(S_{t+1}, A_{t+1}, \mathbf{w}_t) - \hat{q}(S_t, A_t, \mathbf{w}_t) \right] \mathbf{x}(S_t, A_t)

where α>0\alpha > 0 is the learning rate.

Why is it called "Semi-Gradient"?

The update is a semi-gradient because the target Ut=Rt+1+γq^(St+1,At+1,w)U_t = R_{t+1} + \gamma \hat{q}(S_{t+1}, A_{t+1}, \mathbf{w}) depends on w\mathbf{w}. A true gradient of the squared Bellman error 12[Ut(w)−q^(St,At,w)]2\frac{1}{2}[U_t(\mathbf{w}) - \hat{q}(S_t, A_t, \mathbf{w})]^2 would include an additional term −γ∇wq^(St+1,At+1,w)-\gamma \nabla_{\mathbf{w}} \hat{q}(S_{t+1}, A_{t+1}, \mathbf{w}). Semi-gradient methods intentionally omit this term, treating the target as an exogenous scalar. In linear function approximation, semi-gradient SARSA is provably stable and converges to a bounded region near the optimal TD fixed point.

4. The Complete Episodic Control Loop

Algorithm: Episodic Semi-Gradient SARSA for Estimating q̂ ≈ q*───────────────────────────────────────────────────────────────────Input: a parameterized function q̂(s, a, w) with differentiable weights wParameters: step size α > 0, exploration rate ε > 0, discount γ ∈ [0, 1]Initialize: weight vector w arbitrarily (e.g., w = 0)
Loop for each episode:  S ← initial continuous state  Choose A ~ ε-greedy based on q̂(S, ·, w)  Loop for each step of episode:    Take action A, observe reward R and next state S'    If S' is terminal:      w ← w + α [R - q̂(S, A, w)] ∇_w q̂(S, A, w)      Break (end episode)    Choose A' ~ ε-greedy based on q̂(S', ·, w)    w ← w + α [R + γ q̂(S', A', w) - q̂(S, A, w)] ∇_w q̂(S, A, w)    S ← S', A ← A'

Worked numerical example

Consider a 2-feature linear representation for an action on a single transition step (St,At)→Rt+1(St+1,At+1)(S_t, A_t) \xrightarrow{R_{t+1}} (S_{t+1}, A_{t+1}).

System Setup

  • Discount factor: γ=0.9\gamma = 0.9
  • Learning rate: α=0.2\alpha = 0.2
  • Current parameter weights: wt=[1.0,0.5]T\mathbf{w}_t = [1.0, 0.5]^T
  • Active state-action feature vector: x(St,At)=[1.0,2.0]T\mathbf{x}(S_t, A_t) = [1.0, 2.0]^T
  • Environmental reward observed: Rt+1=+2.0R_{t+1} = +2.0
  • Next state-action feature vector: x(St+1,At+1)=[0.5,1.0]T\mathbf{x}(S_{t+1}, A_{t+1}) = [0.5, 1.0]^T

Step 1: Compute Current Action-Value Prediction

q^(St,At,wt)=wtTx(St,At)=(1.0×1.0)+(0.5×2.0)=1.0+1.0=2.0000\hat{q}(S_t, A_t, \mathbf{w}_t) = \mathbf{w}_t^T \mathbf{x}(S_t, A_t) = (1.0 \times 1.0) + (0.5 \times 2.0) = 1.0 + 1.0 = 2.0000

Step 2: Compute Next-Action Value Prediction

q^(St+1,At+1,wt)=wtTx(St+1,At+1)=(1.0×0.5)+(0.5×1.0)=0.5+0.5=1.0000\hat{q}(S_{t+1}, A_{t+1}, \mathbf{w}_t) = \mathbf{w}_t^T \mathbf{x}(S_{t+1}, A_{t+1}) = (1.0 \times 0.5) + (0.5 \times 1.0) = 0.5 + 0.5 = 1.0000

Step 3: Compute On-Policy SARSA Target

Ut=Rt+1+γq^(St+1,At+1,wt)=2.0+(0.9×1.0000)=2.0+0.9000=2.9000U_t = R_{t+1} + \gamma \hat{q}(S_{t+1}, A_{t+1}, \mathbf{w}_t) = 2.0 + (0.9 \times 1.0000) = 2.0 + 0.9000 = 2.9000

Step 4: Compute Semi-Gradient TD Error

δt=Ut−q^(St,At,wt)=2.9000−2.0000=+0.9000\delta_t = U_t - \hat{q}(S_t, A_t, \mathbf{w}_t) = 2.9000 - 2.0000 = +0.9000

Step 5: Update the Parameter Weight Vector

For a linear model, the gradient is the feature vector itself: ∇wq^(St,At,wt)=[1.0,2.0]T\nabla_{\mathbf{w}} \hat{q}(S_t, A_t, \mathbf{w}_t) = [1.0, 2.0]^T.

wt+1=wt+αδtx(St,At)=[1.00.5]+0.2×0.9×[1.02.0]=[1.00.5]+0.18×[1.02.0]=[1.0+0.180.5+0.36]=[1.18000.8600]\mathbf{w}_{t+1} = \mathbf{w}_t + \alpha \delta_t \mathbf{x}(S_t, A_t) = \begin{bmatrix} 1.0 \\ 0.5 \end{bmatrix} + 0.2 \times 0.9 \times \begin{bmatrix} 1.0 \\ 2.0 \end{bmatrix} = \begin{bmatrix} 1.0 \\ 0.5 \end{bmatrix} + 0.18 \times \begin{bmatrix} 1.0 \\ 2.0 \end{bmatrix} = \begin{bmatrix} 1.0 + 0.18 \\ 0.5 + 0.36 \end{bmatrix} = \begin{bmatrix} 1.1800 \\ 0.8600 \end{bmatrix}

Step 6: Verify Action-Value Shift

Re-evaluating (St,At)(S_t, A_t) under the updated weights wt+1=[1.18,0.86]T\mathbf{w}_{t+1} = [1.18, 0.86]^T:

q^(St,At,wt+1)=(1.18×1.0)+(0.86×2.0)=1.18+1.72=2.9000\hat{q}(S_t, A_t, \mathbf{w}_{t+1}) = (1.18 \times 1.0) + (0.86 \times 2.0) = 1.18 + 1.72 = 2.9000

The action-value adjusted from 2.00002.0000 to exactly match the target 2.90002.9000, and the updated weights carry forward to guide subsequent decisions across the continuous environment.

Code

Below is a self-contained Python script implementing Linear Semi-Gradient SARSA with action-stacking feature representation, ε\varepsilon-greedy exploration, and assertions validating both single-step mechanics and continuous navigation convergence:

import mathimport randomfrom typing import Callable, List, Tuple

class LinearSemiGradientSARSA:    """Episodic semi-gradient SARSA control with linear function approximation."""
    def __init__(        self,        num_features: int,        actions: List[int],        alpha: float = 0.2,        gamma: float = 0.9,        epsilon: float = 0.1,    ) -> None:        self.num_features = num_features        self.actions = actions        self.alpha = alpha        self.gamma = gamma        self.epsilon = epsilon        self.weights: List[float] = [0.0] * num_features
    def q_value(self, feature_vector: List[float]) -> float:        """Compute q_hat(s, a, w) = w^T x(s, a)."""        return sum(w * x for w, x in zip(self.weights, feature_vector))
    def select_action(        self,        features_for_action: Callable[[int], List[float]],        explore: bool = True,    ) -> Tuple[int, List[float]]:        """Select action using epsilon-greedy policy over candidate feature vectors."""        if explore and random.random() < self.epsilon:            chosen_action = random.choice(self.actions)            return chosen_action, features_for_action(chosen_action)
        best_a = self.actions[0]        best_features = features_for_action(best_a)        best_q = self.q_value(best_features)
        for a in self.actions[1:]:            features = features_for_action(a)            q = self.q_value(features)            if q > best_q:                best_q = q                best_a = a                best_features = features
        return best_a, best_features
    def update(        self,        curr_features: List[float],        reward: float,        next_features: List[float],        is_terminal: bool,    ) -> Tuple[float, float]:        """Perform a single 1-step semi-gradient SARSA parameter update."""        q_curr = self.q_value(curr_features)        q_next = 0.0 if is_terminal else self.q_value(next_features)
        # On-policy TD target: U = R + gamma * q_hat(S', A', w)        target = reward + self.gamma * q_next        delta = target - q_curr
        # Gradient of linear q_hat w.r.t w is curr_features        for i in range(len(self.weights)):            self.weights[i] += self.alpha * delta * curr_features[i]
        return delta, target

def run_sarsa_demonstration() -> None:    # Part 1: Worked Numerical Walkthrough Verification    agent = LinearSemiGradientSARSA(num_features=2, actions=[0, 1], alpha=0.2, gamma=0.9)    agent.weights = [1.0, 0.5]
    x_curr = [1.0, 2.0]    x_next = [0.5, 1.0]    reward = 2.0
    q_curr = agent.q_value(x_curr)    q_next = agent.q_value(x_next)    print("--- Single-Step Worked Numerical Walkthrough ---")    print(f"Initial q(S_t, A_t):     {q_curr:.4f}")    print(f"Next q(S_t+1, A_t+1):   {q_next:.4f}")
    delta, target = agent.update(x_curr, reward, x_next, is_terminal=False)    print(f"SARSA Target:           {target:.4f}")    print(f"TD Error delta:         {delta:.4f}")    print(f"Updated weights w:      [{agent.weights[0]:.4f}, {agent.weights[1]:.4f}]")
    assert abs(q_curr - 2.0) < 1e-6    assert abs(q_next - 1.0) < 1e-6    assert abs(target - 2.9) < 1e-6    assert abs(delta - 0.9) < 1e-6    assert abs(agent.weights[0] - 1.18) < 1e-6    assert abs(agent.weights[1] - 0.86) < 1e-6
    # Part 2: Continuous 1D Navigation Control Task    # State s in [0.0, 10.0]. Action 0: move left (-1.0), Action 1: move right (+1.0)    # Terminal Goal: reach s >= 9.5 (Reward = +10.0, step cost = -0.1)    # Action-stacking feature representation: 2 base features * 2 actions = 4 features    def make_stacked_features(state: float, action: int) -> List[float]:        base_features = [state / 10.0, 1.0]        if action == 0:            return base_features + [0.0, 0.0]        else:            return [0.0, 0.0] + base_features
    control_agent = LinearSemiGradientSARSA(        num_features=4, actions=[0, 1], alpha=0.1, gamma=0.95, epsilon=0.1    )
    random.seed(42)    # Train across 100 episodes    for ep in range(100):        s = random.uniform(1.0, 7.0)        a, feats = control_agent.select_action(lambda act: make_stacked_features(s, act))        done = False        steps = 0        while not done and steps < 40:            steps += 1            s_next = max(0.0, min(10.0, s + (1.0 if a == 1 else -1.0)))            done = s_next >= 9.5            r = 10.0 if done else -0.1
            if done:                control_agent.update(feats, r, [], is_terminal=True)            else:                a_next, feats_next = control_agent.select_action(                    lambda act: make_stacked_features(s_next, act)                )                control_agent.update(feats, r, feats_next, is_terminal=False)                s, a, feats = s_next, a_next, feats_next
    # Evaluate learned policy at midpoint s = 5.0    q_left = control_agent.q_value(make_stacked_features(5.0, 0))    q_right = control_agent.q_value(make_stacked_features(5.0, 1))
    print("\n--- Continuous Navigation Control Results ---")    print(f"Midpoint s=5.0 Q(move_left):  {q_left:.4f}")    print(f"Midpoint s=5.0 Q(move_right): {q_right:.4f}")    print(f"Learned weights w:            {[round(w, 4) for w in control_agent.weights]}")
    # Moving right toward goal should have substantially higher value than moving left    assert q_right > q_left + 1.0

if __name__ == "__main__":    run_sarsa_demonstration()
# -> Expected output:# -> --- Single-Step Worked Numerical Walkthrough ---# -> Initial q(S_t, A_t):     2.0000# -> Next q(S_t+1, A_t+1):   1.0000# -> SARSA Target:           2.9000# -> TD Error delta:         0.9000# -> Updated weights w:      [1.1800, 0.8600]# -> # -> --- Continuous Navigation Control Results ---# -> Midpoint s=5.0 Q(move_left):  6.1970# -> Midpoint s=5.0 Q(move_right): 7.7725# -> Learned weights w:            [2.8601, 4.767, 5.976, 4.7845]

Watch Out For

Policy Chattering and the Danger of Off-Policy Max Backups

Two prominent stability traps threaten control algorithms when transitioning from tabular to function approximation:

1. Policy Chattering at Decision Boundaries: In tabular control, once Q(s,a1)>Q(s,a2)Q(s, a_1) > Q(s, a_2), the greedy policy remains solidly locked on a1a_1. In function approximation, updating the value of an action in a neighboring state alters weights across the shared feature space. Near decision boundaries, small parameter adjustments cause the dominant action to flip-flop rapidly between a1a_1 and a2a_2 on alternate steps—a pathology called policy chattering. The Fix: Use conservative learning rates (α\alpha), decayed exploration schedules, and smooth feature representations (such as overlapping tile codings or soft action distributions like softmax policies) to prevent erratic boundary oscillations.

2. The Off-Policy Max-Operator Trap (The Deadly Triad): Practitioners frequently attempt to replace SARSA's on-policy target with an off-policy Q-learning target: U=R+γmax⁡aq^(S′,a,w)U = R + \gamma \max_a \hat{q}(S', a, \mathbf{w}). When combined with function approximation and bootstrapping, the max⁡\max operator forms the notorious Deadly Triad. Because off-policy distribution mismatch interacts destructively with bootstrapping, linear Q-learning can spiral into unbounded parameter divergence! The Fix: Stick with on-policy Semi-Gradient SARSA for guaranteed linear stability. If off-policy control is mandatory, use target networks to freeze bootstrap parameters or employ Gradient-TD methods that explicitly minimize projected Bellman errors.

The Quick Version

  • Continuous Action-Value Control: Semi-gradient SARSA scales temporal-difference control to continuous state spaces by parameterizing action-values as q^(s,a,w)≈qπ(s,a)\hat{q}(s, a, \mathbf{w}) \approx q_\pi(s, a).
  • Action-Stacked Features: Discrete actions are typically handled by stacking state features into disjoint blocks x(s,a)\mathbf{x}(s, a), allowing independent linear action-value surfaces to evolve without destructive parameter interference.
  • On-Policy Bootstrapping: Updates evaluate the actual next action At+1∼π(⋅∣St+1)A_{t+1} \sim \pi(\cdot \mid S_{t+1}), ensuring that targets remain on-policy and avoiding the divergent instabilities of the Deadly Triad.
  • Semi-Gradient Detachment: The algorithm ignores the target's mathematical dependency on w\mathbf{w}, treating Rt+1+γq^(St+1,At+1,wt)R_{t+1} + \gamma \hat{q}(S_{t+1}, A_{t+1}, \mathbf{w}_t) as a fixed scalar to ensure stable SGD weight updates.