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 nnimport jaximport 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.
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.
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.