Skip to content

About

Tuning PID and neural-network controllers by differentiating through simulated plant dynamics with JAX

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Latest commit

 

History

23 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

jax_based_controller

Tuning PID and neural-network controllers with gradient descent, by differentiating through a simulation of the system they control.

How it works

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.
  • ControlSystem computes 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.

Results

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.

PID gains during training on the thermal plant Neural controller MSE during training on the Cournot plant

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.

Usage

pip install -r requirements.txt
python main.py --config config/thermal_pid.yaml

Each file in config/ defines one experiment: plant, controller, training settings and plots.

Project structure

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

About

Tuning PID and neural-network controllers by differentiating through simulated plant dynamics with JAX

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages