Fit spiking neural network training into the VRAM you actually have.
Backpropagation through time keeps every intermediate activation, so memory
grows linearly with the simulation length T. And T is exactly the knob that
makes a spiking model work, which is why a network that trains happily at
T = 8 dies at T = 32 on the same card.
This library makes that trade cheaper, then tells you where the limit is before you spend nine hours discovering it.
pip install -e ".[dev]"from spikefit import AutoTuner, SpikingMLP, compare
# What actually fits, measured before committing to a run
tuner = AutoTuner(max_vram_gb=8.0)
print(tuner.tune(lambda: SpikingMLP(128, 512, 10), timesteps=64))
# What checkpointing costs and saves
print(compare(lambda: SpikingMLP(128, 512, 10), batch=64, timesteps=64,
device="cuda"))For batch B, simulation length T, L layers of width H, standard BPTT
stores
M_standard = O(B · T · L · H)
Split the timeline into K chunks. Keep only the membrane state at each chunk
boundary during the forward pass, and recompute the chunk's interior during the
backward pass. Peak memory becomes the boundaries plus the one chunk being
recomputed:
M_checkpointed = O(B · L · H · (K + T/K))
Minimise by differentiating with respect to K:
d/dK (K + T/K) = 1 - T/K² = 0 => K = √T
M_optimal = O(B · L · H · √T)
The saving grows as √T, so it gets better exactly where the problem gets
worse. optimal_chunk_count implements this, and a test verifies the choice
against an exhaustive search over all K for a range of T.
RTX 5060 Laptop, 8.55 GB, PyTorch 2.11 + CUDA 12.8. Three LIF layers of width
512, batch 128. Reproduce with python examples/benchmark.py.
| T | K | Standard | Checkpointed | Memory saved | Time |
|---|---|---|---|---|---|
| 16 | 4 | 55 MB | 36 MB | 34.2% | +13.9% |
| 32 | 4 | 90 MB | 44 MB | 50.4% | +43.8% |
| 64 | 8 | 157 MB | 49 MB | 69.0% | +50.7% |
| 128 | 10 | 292 MB | 60 MB | 79.4% | +57.6% |
| 256 | 16 | 563 MB | 74 MB | 86.8% | +49.7% |
The two columns behave exactly as the analysis predicts. Standard memory grows linearly: 16× the timesteps costs 10.2× the memory, and subtracting the roughly 25 MB of fixed overhead gives 17.9×, which is the 16× the model calls for. Checkpointed memory grows as √T: the same 16× costs 2.1× gross, or 4.5× net of overhead, against the 4× the square-root law predicts.
Every chunk is computed twice on the forward side. A backward pass costs roughly two forwards, so 3 units of work become 4, predicting about +33%. Measured overhead runs from +14% to +58%, so the estimate is optimistic. The extra comes from launching many more small kernels and from the recomputation not reusing cached workspaces. If you need the number for your model, measure it: that is what the profiler is for.
The memory table is the mechanism. This is the consequence.
AutoTuner was given a 2 GB budget and asked for the largest batch that fits,
with and without checkpointing:
| T | Without | With | Gain |
|---|---|---|---|
| 16 | 6,463 | 15,870 | 2.46× |
| 32 | 3,242 | 9,727 | 3.00× |
| 64 | 1,621 | 8,471 | 5.23× |
| 128 | 810 | 5,631 | 6.95× |
| 256 | 419 | 4,351 | 10.38× |
Read the "Without" column downwards: 6463, 3242, 1621, 810, 419. The batch
halves every time T doubles, which is O(B · T) showing up directly in
what you are allowed to run.
Going from T = 16 to T = 256 costs you 93.5% of your batch size without
checkpointing, and 72.6% with it. At T = 256 that is ten times the batch
for the same card.
A memory optimisation that changes the gradient is worse than none, because the model still trains, still converges to something, and is wrong in a way nothing reports. So that is the first thing tested and the strictest.
Gradients under checkpointing are bit-identical to standard BPTT, tested
across T ∈ {1, 2, 7, 16, 33} and K ∈ {1, 2, 3, 5, 16}:
max |grad difference| per tensor: 0.00e+00, 0.00e+00, 0.00e+00, 0.00e+00
Recomputation must reproduce the original forward pass exactly, which anything
stochastic breaks: dropout, or a neuron with a random threshold, would draw
different numbers the second time and produce gradients for a network that
never ran. torch.utils.checkpoint is used rather than a hand-rolled loop
precisely because it saves and restores RNG state.
Passing the gradient tests would not, on its own, prove anything is being
saved: a no-op would pass them too. So the retained activations are counted
directly, using saved_tensors_hooks:
| Saved tensors | Retained | |
|---|---|---|
| Standard | 316 | 1.73 MB |
| Checkpointed | 24 | 0.07 MB |
Counting graph nodes does not work here, which took a failing test to notice.
With use_reentrant=False the autograd graph has the same structure either
way; only the set of retained tensors differs.
Probing beats predicting. Activation memory has a clean closed form and
estimate_activation_bytes implements it, but peak VRAM also contains
parameters, gradients, optimiser state, cuDNN workspaces, the CUDA context, and
allocator fragmentation that depends on allocation order. The estimate gives
the right shape, which makes the search efficient. Only a real pass gives the
right number.
So AutoTuner seeds a bisection from the analytical model and then measures,
catching out-of-memory as data rather than as failure:
largest fitting configuration on cuda:
batch 9727
timesteps 32
chunks 4
peak 1.74 GB of 1.80 GB (97%)
probes 14
Verified by using the answer: at the recommended batch peak is 1800 MB against an 1800 MB budget; at 1.6× that batch it reaches 2866 MB and exceeds it.
The default safety_fraction=0.90 exists because a configuration measured at
exactly the limit will fail later. Fragmentation grows over a long run, and
evaluation batches, logging and checkpointing all want memory the probe never
asked for.
The decorator supplies a configured runner. It cannot rewrite your timestep loop, because only your model knows what its state is and how a step advances it:
from spikefit import spikefit_checkpoint
class MySNN(nn.Module):
@spikefit_checkpoint(n_chunks=8)
def forward(self, x, timesteps=64, checkpoint=None):
state = self.init_state(x)
return checkpoint(lambda t, s: self.step(x, t, s), timesteps, state)Or call it directly:
from spikefit import checkpoint_sequence
state = checkpoint_sequence(step_fn, timesteps=64, state=initial) # K = √T
state = checkpoint_sequence(step_fn, 64, initial, n_chunks=1) # disabledstate may be a tensor or a tuple or list of them. Non-tensor entries pass
through untouched and are not differentiated. Anything else is refused rather
than silently dropping a gradient.
SpikingMLP exists so the library can be tested and profiled on its own, not
because you should use it. checkpoint_sequence takes any step function, so
snnTorch and SpikingJelly modules work unchanged.
examples/with_snntorch.py wraps two snn.Leaky layers. Nothing in the
network changes; the timestep loop is handed over, and that is all:
snnTorch net, T=64, K=8, device=cuda
max |grad difference| vs standard BPTT: 0.00e+00
retained activations: 8.13 MB -> 0.28 MB (96.5% less)
Measured with snnTorch 1.0.0.
- Measured on one GPU. An RTX 5060 Laptop. The
√Tscaling is architecture-independent, but the time overhead and the fixed memory floor are not. - Time overhead exceeds theory. Between +14% and +58% against a predicted +33%. Kernel launch count and lost workspace reuse account for the gap, but that has not been profiled per-kernel.
- Recomputation must be deterministic. A stochastic step function will
produce wrong gradients.
torch.utils.checkpointrestores RNG state, which covers dropout, but a step reading external mutable state is on you. - No activation offloading. Moving activations to host memory is the other standard technique and is not implemented. It trades bandwidth rather than compute and suits a different bottleneck.
- No mixed precision interaction tested. AMP should compose, since this operates above the dtype, but it has not been verified.
- Analytical estimates assume dense layers. Convolutional or recurrent
spiking layers store different intermediates, so
estimate_activation_byteswill be wrong for them. The tuner measures, so it will not be.
pytest # 94 tests, 98% coverage on a GPU machine
ruff check .
mypy # strictTests needing a GPU are marked cuda and skip without one. Correctness,
scaling and the tuner's search logic are all tested on CPU, so a machine
without a GPU still checks everything that can be checked.
Coverage is 98% with a GPU and 92% without. The gap is the code that
actually reads peak VRAM: AutoTuner._probe and the profiler's synchronize
and reset-peak helpers. Mocking torch.cuda would raise the number without
testing anything, so CI enforces 90 and the bisection those functions feed is
tested separately against a synthetic memory model with a known exact answer.
If pytest fails to start with a PluginValidationError mentioning
nengo, an unrelated package in your environment ships an incompatible pytest
plugin. Run pytest -p no:nengo, or use a clean virtual environment.
- Chen et al., Training Deep Nets with Sublinear Memory Cost, 2016. The
√Tresult. - Griewank & Walther, Algorithm 799: revolve, TOMS 2000. Optimal checkpointing schedules.
- Neftci et al., Surrogate Gradient Learning in Spiking Neural Networks, IEEE SPM 2019.
MIT. See LICENSE.