Training CartPole on WebGPU

This page trains a neural network to balance a pole on a cart. Training starts when the page loads and runs in WebGPU compute shaders on your GPU. If your browser doesn't support WebGPU, the same algorithm runs on the CPU.

Live TrainingLoading…
Cart and pole driven by the latest policy
t = 0
Updates
0
Environment Steps
0
Episodes
0
Mean Length
–
Mean Episode Length, Last 20 Episodes–
Waiting for the first finished episode…
Policy Loss–
Waiting for the first update…
Value Loss–
Waiting for the first update…
Starting the trainer…

At first, the policy picks left and right with equal probability, and the pole falls after about 20 steps. In my tests, the mean episode length passed 475 steps after 230 to 800 updates. That is the threshold Gymnasium uses to call CartPole solved. The rest of this post explains the math behind each number in the figures.

CartPole as a Markov decision process

CartPole is a control task from Barto, Sutton, and Anderson (1983). A cart moves on a frictionless track, and a hinge attaches a pole to the cart. At each time step, you push the cart left or right with a fixed force. The goal is to keep the pole upright for as long as possible.

A Markov decision process (MDP) describes this kind of problem. Markov means that the next state depends only on the current state and action, not on earlier steps. The CartPole MDP has these parts:

Each step earns the same reward, so the return of an episode equals its length. More reward means more steps with the pole up. The trainer discounts later rewards with γ=0.99\gamma = 0.99, so the return from step tt is Gt=∑k≥0γkrt+kG_t = \sum_{k \ge 0} \gamma^k r_{t+k}.

One Transition
The cart in the first figure takes one of these steps every 20 ms. Pause, then step to read each value.

How the cart and pole move

The transition function ff comes from the equations of motion for a pole on a cart. The simulator uses the same form as Gymnasium's CartPole. It computes the angular acceleration of the pole first, then uses it to compute the acceleration of the cart:

θ¨t=gsin⁡θt−cos⁡θt⋅Ft+mpl θ˙t2sin⁡θtmc+mpl(43−mpcos⁡2θtmc+mp)x¨t=Ft+mpl(θ˙t2sin⁡θt−θ¨tcos⁡θt)mc+mp\ddot\theta_t = \frac{g\sin\theta_t - \cos\theta_t \cdot \dfrac{F_t + m_p l\,\dot\theta_t^2\sin\theta_t}{m_c + m_p}}{l\left(\dfrac{4}{3} - \dfrac{m_p\cos^2\theta_t}{m_c + m_p}\right)} \qquad \ddot x_t = \frac{F_t + m_p l\left(\dot\theta_t^2\sin\theta_t - \ddot\theta_t\cos\theta_t\right)}{m_c + m_p}

The constants are the Gymnasium defaults:

The simulator then moves each state variable forward by τ=0.02 s\tau = 0.02\,\text{s} with an explicit Euler step. Each position changes by its old velocity, and each velocity changes by the new acceleration:

xt+1=xt+τx˙tx˙t+1=x˙t+τx¨tθt+1=θt+τθ˙tθ˙t+1=θ˙t+τθ¨tx_{t+1} = x_t + \tau\dot x_t \qquad \dot x_{t+1} = \dot x_t + \tau\ddot x_t \qquad \theta_{t+1} = \theta_t + \tau\dot\theta_t \qquad \dot\theta_{t+1} = \dot\theta_t + \tau\ddot\theta_t

The agent never sees these equations. It observes the four numbers in sts_t and the reward, and it learns the effect of each push from experience.

Accelerations for This Step
The simulator evaluates both equations once per step, then applies Euler updates.

The policy and value networks

The trainer uses two small neural networks. The policy network πθ\pi_\theta maps a state to a probability for each action. The value network VϕV_\phi maps a state to an estimate of the discounted return from that state.

Each network has one hidden layer of 32 units with a tanh activation:

πθ(⋅∣s)=softmax⁡(W2tanh⁡(W1s+b1)+b2)Vϕ(s)=w2⊤tanh⁡(U1s+c1)+c2\pi_\theta(\cdot \mid s) = \operatorname{softmax}\left(W_2 \tanh(W_1 s + b_1) + b_2\right) \qquad V_\phi(s) = w_2^\top \tanh(U_1 s + c_1) + c_2

The policy network has 226 parameters and the value network has 193, so each update adjusts 419 weights. The output weights of the policy start near zero, so both actions start with a probability near 0.5.

The agent samples each action from πθ\pi_\theta instead of always taking the more likely action. Sampling makes the agent try both actions, and it needs both to learn which one works better.

Network Outputs
Both networks read the same four numbers. Only the policy chooses the action.

Learning from advantages

The policy improves when actions that did better than expected become more likely. The advantage A^t\hat A_t measures how much better an action did. It is the return after ata_t minus the prediction Vϕ(st)V_\phi(s_t).

The trainer estimates the advantage with generalized advantage estimation (GAE). GAE starts from the temporal difference error δt\delta_t, which compares one real reward plus the predicted value of the next state with the predicted value of the current state:

δt=rt+γ(1−dt+1)Vϕ(st+1)−Vϕ(st)A^t=δt+γλ(1−dt+1)A^t+1\delta_t = r_t + \gamma (1 - d_{t+1}) V_\phi(s_{t+1}) - V_\phi(s_t) \qquad \hat A_t = \delta_t + \gamma\lambda (1 - d_{t+1}) \hat A_{t+1}

With λ=0.95\lambda = 0.95, A^t\hat A_t is a weighted sum of the errors over the next steps. The factor 1−dt+11 - d_{t+1} stops the sum at the end of an episode. The target for the value network is R^t=A^t+Vϕ(st)\hat R_t = \hat A_t + V_\phi(s_t).

Each update minimizes two losses over a batch of B=512B = 512 samples:

Lπ(θ)=−1B∑tA^tlog⁡πθ(at∣st)−βHLV(ϕ)=12B∑t(Vϕ(st)−R^t)2L_\pi(\theta) = -\frac{1}{B}\sum_t \hat A_t \log \pi_\theta(a_t \mid s_t) - \beta \mathcal{H} \qquad L_V(\phi) = \frac{1}{2B}\sum_t \left(V_\phi(s_t) - \hat R_t\right)^2

The gradient of LπL_\pi increases log⁡πθ(at∣st)\log \pi_\theta(a_t \mid s_t) when A^t\hat A_t is positive and decreases it when A^t\hat A_t is negative. The entropy term H\mathcal{H}, with β=0.01\beta = 0.01, keeps the policy from becoming certain too early. Adam updates the policy with a learning rate of 0.003 and the value network with 0.01.

The policy loss doesn't measure progress the way a supervised loss does. Its scale depends on the advantages, which change as the value network improves. Use the episode length chart to measure progress. The value loss shows how well VϕV_\phi predicts returns. In my runs, it rose as episodes got longer, because the discounted returns grew from under 20 to almost 100.

Latest Update
The advantage and log probability belong to the first sample of the batch. The losses average all 512 samples.

Running the trainer on the GPU

An update has two kinds of independent work. Each environment steps on its own, and the gradient of each parameter is a separate sum over the batch. The trainer runs one WGSL compute shader for each kind:

  1. The rollout kernel runs one thread per environment. Each of the 32 threads runs 16 steps. At each step, it evaluates both networks, samples an action, steps the physics, and resets the environment if the episode ended. Then it walks back through its 16 steps to compute A^t\hat A_t and the gradients for each sample.
  2. The update kernel runs one thread per parameter. Each of the 419 threads sums its gradient over the 512 samples and applies the Adam step to its parameter.

The environment states, the random number generator states, and the Adam moments stay in GPU buffers between updates. After the updates for each animation frame, the page reads back the weights, the lengths of finished episodes, and the losses. The cart and the equations on this page use those weights on the CPU.

The GPU and CPU versions use the same PCG random number generator, so I can check one against the other. From the same weights and seeds, they produce the same episodes. After 3 updates, the weights differ by at most 7.6×10−77.6 \times 10^{-7} because of float32 rounding on the GPU. On an Apple M5 Max, a GPU update took 0.3 ms in Chrome, including one readback for every 4 updates. A CPU update took 2.1 ms in Node.

The page limits training to 30 updates per second so you can watch the policy learn. At that rate, a run that solves in 300 updates takes 10 s.

What this trainer leaves out

The core loop has three steps:

  1. Act with the current policy.
  2. Measure how much better each action did than the value network predicted.
  3. Move probability toward the better actions.

CartPole is small enough that this loop fits in two compute shaders.

The trainer leaves out these standard techniques: