Skip to content
GitHub

The pilot

Learn

The pilot

How the network on the front page learned to fly a drone with nothing to imitate, how much of it it needs, and how the browser runs the network sparx trained.

The drone

The drone is a bar with a rotor at each end, in a vertical plane: mass 1 kg, arms of 25 cm, at most 12 N per rotor, with some air drag. Its state is where it is, how fast it moves, its angle and its spin. Every 10 ms the physics takes one semi-implicit Euler step: the two thrusts push along the drone's up axis, their difference turns it, and gravity pulls it down. A drone that hits the edge of its box bounces off at 30% of its speed.

The pilot reads seven numbers each step: the way to its target (clipped to 1.5 m, so a far target reads as "that way"), its velocity, the sine and cosine of its angle, and its spin. They enter a dense layer as currents. Two layers of 64 LIF neurons follow, then two leaky integrators whose membranes set the rotors' thrust through a sigmoid. Hovering is a membrane of zero. Nothing else carries memory: the membranes are the network's whole state.

The code
import flax.linen as nn
import jax
import jax.numpy as jnp
from sparx.nn import LI, LIF
pilot = nn.Sequential([
nn.Dense(64), LIF(tau=3.0, reset="zero"), # 7 readings in: the way to the target,
nn.Dense(64), LIF(tau=3.0, reset="zero"), # velocity, attitude and spin
nn.Dense(2), LI(tau=5.0), # 2 membranes out: the rotors' thrust
])
def step(params, carried, readings):
"""10 ms of the network, its membranes carried in the `state` collection."""
out, mutated = pilot.apply({"params": params, "state": carried}, readings[None], mutable=["state"])
return out[0], mutated["state"]
params = pilot.init(jax.random.key(0), jnp.zeros((1, 1, 7)))["params"]
membranes, carried = step(params, {}, jnp.zeros((256, 7))) # 256 drones, every neuron at rest
# Training scans step() and the drone's physics over 2 s of flight and takes jax.grad of the
# distance to the target: through the spikes by their surrogate, and through the physics.

Learning by gradients through physics

The physics is written in JAX, so a flight is differentiable end to end: through the drone, and through the spikes by their surrogate. Each training step flies 256 drones for 2 s from random states, 40% of them at any angle, in boxes from a phone's shape to an ultrawide's, chasing targets that jump every 1.2 s on average, through gusts. The loss is their mean distance to their targets, plus small penalties on tilt and spin, and on neurons firing outside a band of rates. Adam follows its gradient.

01201,0002,0003,0004,000training stepdistance (m)
Mean distance to the target over training flights, which keep moving the target and pushing the drone, so it never reaches zero. 4,000 steps took 10 minutes on a 4-vCPU container.

There is no expert to copy and no reward to estimate: just the derivative of the distance with respect to every weight, through two seconds of flight. The gradient is clipped to a norm of 1, which keeps an occasional exploding one from undoing the rest.

How well it flies

On 1,000 flights of 6 s to a fixed target from random starts, half of them at any angle, it arrived within 15 cm and stayed there on 100%, in a median of 0.93 s. That includes all 252 starts more than a quarter turn from upright. On 1,000 more, thrown at about 4 m/s and spun at about 15 rad/s, it arrived on 100%. After 6 s the median drone was 1.6 cm from its target.

Silencing neurons

On the front page you can silence neurons and watch the flying get worse. Here is the same experiment in sparx: silence neurons at random from both layers, five times for each count, and fly 500 flights each time.

0501000163248648096112neurons silenced, of 128flights arriving (%)
Each dot is one random set of silenced neurons; the line is their mean.

The network is not uniformly redundant. Most sets of 8 or 16 silenced neurons leave it flying, but one set of 8 grounded it entirely: some neurons carry something the rest cannot replace. By 48 the drone mostly falls, and past 64 it never arrives.

The browser runs sparx's network

The front page loads the trained weights, 4,802 numbers, and steps the network and the physics in TypeScript, in double precision, with the same operations in the same order as sparx. When this page was built, it flew sparx's two recorded float64 flights, 1,400 steps with targets, gusts and spins, and fired the same 38,093 spikes, every one at the same step. The drone's state stayed within 4e-13 of sparx's.

sparx exports the same network through NIR, the format neuromorphic hardware and other simulators read: NIR export.

All of Learn