Tuning PID and neural-network controllers with gradient descent, by differentiating through a simulation of the system they control.
Each training epoch simulates the plant over a sequence of timesteps with random disturbances and measures the mean squared error from the target. JAX then computes the gradient of that error with respect to the controller parameters, through every timestep of the simulation. The same training loop tunes both a classic PID controller (three gains) and a small neural network.
error ──► controller (PID or neural net) ──► control signal ──► plant + disturbance ──► output
▲ │
└──────────────────────────────────────── target − output ◄──────────────────────────────┘
The training step is in control_system.py. simulate_epoch runs the simulation and returns the MSE, and JAX differentiates it end to end:
mse_and_gradients = jax.value_and_grad(self.simulate_epoch)
mse, gradients = mse_and_gradients(params, disturbances, epoch_key)
params = jax.tree.map(
lambda param, grad: param - self.learning_rate * grad,
params,
gradients
)Because the update uses jax.tree.map, it works on any parameter structure: the three PID gains and the weight matrices of the neural network go through the same code.
Architecture. Controllers and plants never know about each other. All interaction goes through a coordinator, ControlSystem:
- Controllers receive the error, its derivative and its integral, and return a control signal.
- Plants receive a control signal and a disturbance, and return the regulated output.
ControlSystemcomputes the error terms, runs the simulation and does the gradient updates.
A new plant or controller is added by implementing BasePlant or BaseController, without changing any other code.
Both controller types learn to regulate all three plants: a draining bathtub, a Cournot competition market model and a heated room.
MSE from the first to the last epoch:
| Plant | PID | Neural net |
|---|---|---|
| Bathtub | 1.14 → 0.051 | 12.4 → 0.023 |
| Cournot | 0.027 → 0.0032 | 0.51 → 0.018 |
| Thermal | 258 → 39.9 | 899 → 80.9 |
MSE is in each plant's own units, so compare along a row, not down a column. The number of epochs differs between runs; see the lab report.
Left: on the thermal plant, the integral gain quickly takes over. This is what you would expect: holding a temperature against constant heat loss depends on accumulated error. Here the plain PID also ends with half the error of the neural controller.
Right: the neural controller on the Cournot plant initially could not learn at all. The plant clips production to [0, 1], and clipping has zero gradient at the boundary, so large control signals cut off all learning. Scaling the network output down kept the plant in the range where gradients flow, and training converged after about 12 epochs.
All six runs, with parameters and analysis, are in the lab report.
pip install -r requirements.txt
python main.py --config config/thermal_pid.yamlEach file in config/ defines one experiment: plant, controller, training settings and plots.
main.py entry point: loads a config, builds controller and plant, trains
config/ one YAML file per experiment
src/
consys/control_system.py simulation loop and gradient-descent training
controllers/ BaseController, PIDController, NeuralNetController
plants/ BasePlant, BathtubPlant, CournotPlant, ThermalPlant
utils/ config loading, activations, weight initialisation
visualization/plotter.py MSE and PID-gain plots
docs/lab_report.md all six runs with parameters and analysis

