Fully asynchronous, event-driven (AED) neural networks in JAX + MPI.
Each MPI rank owns one layer (or a shard of one layer). Layers communicate by
passing sparse events — (neuron_idx, value) for MLP, (channel, x, y, value)
for CNN/ResNet — through MPI point-to-point messages, terminated by an
END_SIGNAL. Neurons fire and forward individually rather than waiting for a
whole layer to finish.
| Script | Model |
|---|---|
async_MLP_general.py |
AED multilayer perceptron |
async_CNN_general.py |
AED convolutional network (data + model parallelism) |
async_ResNet_general.py |
AED CNN with residual (skip) connections |
mpirun -n <ntasks> python async_MLP_general.py --config configs/MLP_config.yaml<ntasks> must be a multiple of len(layer_sizes); ntasks / len(layer_sizes)
is the number of data-parallel replicas. For inference from a checkpoint, set
mode: inference and rerun: <path/to/checkpoint.json> in the config.
One generic config per runner lives in configs/ — every supported parameter is
listed there with a comment.
firing_nb— top-k neurons that fire per layer per event. Lower means fewer events in the network: faster and more energy-efficient, at some accuracy cost.sync_rate— how many events a neuron must receive before it may fire again. Use1when benchmarking raw performance.restrict— soft-reset multiplier applied to a neuron's value after it fires.frame_size— event datasets only: fire the first hidden layer once per true time frame instead of using itssync_rate.
async_{MLP,CNN,ResNet}_general.py runners
configs/ one generic example config per runner
dataset_helpers/ dataset loaders (MNIST, N-MNIST, SHD, DVS, NCARS, CIFAR-10, iris)
forward_backward_pass/ event-driven inference, backprop, losses
other_helpers/ MPI partitioning, params/config handling, weight init, pooling
Results are written to network_results/<dataset>/training/<arch>/ as JSON.
Note that result JSON keys are space-separated (firing number,
synchronization rate, learning rate), not the snake_case config names.