Skip to content

Latest commit

 

History

2 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

spikefit

Fit spiking neural network training into the VRAM you actually have.

CI Python License: MIT

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"))

The memory problem

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.

Measured

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.

What it costs

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.

What it buys: batch size

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.

Correctness

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.

The tuner

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.

Using it

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)      # disabled

state 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.

Works with your model

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.

Limitations

  • Measured on one GPU. An RTX 5060 Laptop. The √T scaling 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.checkpoint restores 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_bytes will be wrong for them. The tuner measures, so it will not be.

Development

pytest              # 94 tests, 98% coverage on a GPU machine
ruff check .
mypy                # strict

Tests 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.

Related reading

  • Chen et al., Training Deep Nets with Sublinear Memory Cost, 2016. The √T result.
  • Griewank & Walther, Algorithm 799: revolve, TOMS 2000. Optimal checkpointing schedules.
  • Neftci et al., Surrogate Gradient Learning in Spiking Neural Networks, IEEE SPM 2019.

License

MIT. See LICENSE.

About

Gradient checkpointing and VRAM auto-tuning for spiking neural network training. O(sqrt(T)) activation memory with bit-identical gradients.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages