Fully Spiking Variational Autoencoder

Hiromichi KamataYusuke MukutaTatsuya Harada

article2022AAAI57 citations

Proposes the first fully spiking variational autoencoder by modeling the latent space with autoregressive Bernoulli processes, enabling energy-efficient neuromorphic image generation that matches or exceeds the quality of conventional artificial neural networks.

Listen

Deploying advanced generative artificial intelligence models to edge devices presents significant challenges due to heavy computational and energy demands. Spiking neural networks, which mimic the human brain by processing information through binary event-driven signals, offer an energy-efficient and ultra-fast computing alternative for specialized neuromorphic hardware. However, previous attempts to perform generative image modeling with these brain-inspired networks have produced low-quality visuals or relied on hybrid designs requiring conventional artificial neural networks, which prevents end-to-end execution on neuromorphic chips.

The article demonstrates the Fully Spiking Variational Autoencoder, the first generative image model constructed entirely out of spiking neural network layers. The primary objective is to prove that a fully spiking architecture can generate and reconstruct complex images at quality levels that match or exceed conventional deep learning models while preserving compatibility with neuromorphic hardware.

To overcome the restriction that spiking networks can only transmit binary signals, the authors developed an autoregressive Bernoulli spike sampling technique. This method replaces the continuous, floating-point calculations of traditional variational autoencoders with discrete, random sampling suitable for hardware-based random number generators. The authors evaluated the system across four standard visual benchmark datasets: MNIST, Fashion-MNIST, CIFAR-10, and CelebA. The model was trained using an optimized discrepancy loss designed specifically for spike trains and benchmarked directly against equivalent conventional neural network architectures.

The experimental findings show that the proposed fully spiking model matches or surpasses traditional networks across key generation metrics. First, the spiking model achieved superior Inception Scores across all evaluated datasets, generating sharper and less hazy images. Second, it delivered lower image reconstruction errors and improved Fréchet Inception Distance on datasets such as MNIST and Fashion-MNIST. Third, the model showed a major reduction in expensive computations, requiring roughly 14.8 times fewer multiplications per inference than standard architectures, despite a 6.8-fold increase in simpler addition operations. Finally, an optimal balance between expressive capability and sample fidelity was established at a spike train length of 16 timesteps.

These results establish that discrete, spike-based architectures can effectively handle complex generative tasks without suffering from common generative model failures like latent space collapse. Because neuromorphic hardware can execute these event-driven models with significant speedups and energy reductions, this architecture provides a practical pathway for deploying real-time, low-power image synthesis and anomaly detection directly to battery-constrained edge devices.

Future work should focus on physically deploying and validating the architecture on production neuromorphic chips, such as Loihi or TrueNorth, to confirm real-world power and latency advantages. Additionally, researchers should incorporate recent advances in deep hierarchical architectures to scale the model toward high-resolution image generation. While current empirical tests demonstrate strong performance on standard low-resolution benchmarks up to 64x64 pixels, stakeholder confidence in wider commercial adoption will depend on confirming these computational efficiencies on actual hardware at larger image scales.

arXiv: 2110.00375
Cover for Fully Spiking Variational Autoencoder

Abstract

Spiking neural networks (SNNs) can be run on neuromorphic devices with ultra-high speed and ultra-low energy consumption because of their binary and event-driven nature. Therefore, SNNs are expected to have various applications, including as generative models being running on edge devices to create high-quality images. In this study, we build a variational autoencoder (VAE) with SNN to enable image generation. VAE is known for its stability among generative models; recently, its quality advanced. In vanilla VAE, the latent space is represented as a normal distribution, and floating-point calculations are required in sampling. However, this is not possible in SNNs because all features must be binary time series data. Therefore, we constructed the latent space with an autoregressive SNN model, and randomly selected samples from its output to sample the latent variables. This allows the latent variables to follow the Bernoulli process and allows variational learning. Thus, we build the Fully Spiking Variational Autoencoder where all modules are constructed with SNN. To the best of our knowledge, we are the first to build a VAE only with SNN layers. We experimented with several datasets, and confirmed that it can generate images with the same or better quality compared to conventional ANNs. The code is available at https://github.com/kamata1729/FullySpikingVAE.

Table of Contents

  • Introduction
  • Related Work
  • Development of SNNs
  • Spike Neuron Model
  • Variational Autoencoder
  • Variational Reccurent Neural Network
  • Generative Models in SNN
  • Applying VAE on SNN
  • Autoregressive Model for Spike Train Modeling
  • Proposed Method
  • Overview of FSVAE
  • Autoregressive Bernoulli Spike Sampling
  • Spike to Image Decoding Using Membrane Potential
  • Loss Function
  • Experiments
  • Datasets
  • Network Architecture
  • Training Settings
  • Evaluation Metrics
  • Qualitative Evaluation
  • Conclusion
  • Acknowledgments
  • References

Knowls

  1. Knowl 1 — Fully Spiking Variational Autoencoder Architecture

    model/method

    The Fully Spiking Variational Autoencoder (FSVAE) is an end-to-end spiking neural network (SNN) generative model designed to operate entirely with binary spike trains across all components. An input image xx is converted into binary spike trains x1:T∈{0,1}Cin×H×W×Tx_{1:T} \in \{0, 1\}^{C_{in} \times H \times W \times T} over TT discrete timesteps via direct input encoding.

    The FSVAE architecture consists of:

    1. SNN Encoder: A sequence of convolutional layers with kernel size 33 and stride 22, threshold-dependent batch normalization (tdBN), and leaky integrate-and-fire (LIF) neurons that map x1:Tx_{1:T} into encoder spike features x1:TE∈{0,1}C×Tx^E_{1:T} \in \{0, 1\}^{C \times T}, with latent channel dimensionality C=128C = 128.
    2. Autoregressive Posterior SNN (fqf_q): Formed by three fully connected layers that take the concatenated vector (zq,t−1,xtE)∈{0,1}2C(z_{q,t-1}, x^E_t) \in \{0, 1\}^{2C} (with zq,0=0z_{q,0} = 0) and output latent spikes zq,t∈{0,1}Cz_{q,t} \in \{0, 1\}^C sequentially for t=1,…,Tt = 1, \dots, T.
    3. Autoregressive Prior SNN (fpf_p): Formed by three fully connected layers that take previous latent sample zp,t−1∈{0,1}Cz_{p,t-1} \in \{0, 1\}^C (with zp,0=0z_{p,0} = 0) and sequentially output prior latent spikes zp,t∈{0,1}Cz_{p,t} \in \{0, 1\}^C.
    4. SNN Decoder: Deconvolutional layers with tdBN and LIF neurons that convert sampled latent spikes z1:T∈{0,1}C×Tz_{1:T} \in \{0, 1\}^{C \times T} into output spike trains x^1:T∈{0,1}Cout×H×W×T\hat{x}_{1:T} \in \{0, 1\}^{C_{out} \times H \times W \times T}.
    5. Spike-to-Image Decoding Layer: A non-firing LIF neuron layer that converts the spike trains x^1:T\hat{x}_{1:T} into a continuous reconstructed image x^\hat{x} using the final membrane potential.
  2. Knowl 2 — Autoregressive Bernoulli Spike Sampling

    model/method

    In the Fully Spiking Variational Autoencoder (FSVAE), the latent space is modeled as a discrete Bernoulli process rather than a continuous Gaussian distribution to maintain full binary compatibility with neuromorphic hardware:

    q(z1:T∣x1:T)=∏t=1Tq(zt∣x≤t,z<t)=∏t=1TBer(πq,t)q(z_{1:T} \mid x_{1:T}) = \prod_{t=1}^T q(z_t \mid x_{\le t}, z_{<t}) = \prod_{t=1}^T \text{Ber}(\pi_{q,t})

    p(z1:T)=∏t=1Tp(zt∣z<t)=∏t=1TBer(πp,t)p(z_{1:T}) = \prod_{t=1}^T p(z_t \mid z_{<t}) = \prod_{t=1}^T \text{Ber}(\pi_{p,t})

    At each timestep t∈{1,…,T}t \in \{1, \dots, T\}, autoregressive spiking networks fqf_q and fpf_p produce binary spike output vectors:

    ζq,t:=fq(zq,t−1,xtE;Θq,t)∈{0,1}kC\zeta_{q,t} := f_q(z_{q,t-1}, x^E_t; \Theta_{q,t}) \in \{0, 1\}^{kC}

    ζp,t:=fp(zp,t−1;Θp,t)∈{0,1}kC\zeta_{p,t} := f_p(z_{p,t-1}; \Theta_{p,t}) \in \{0, 1\}^{kC}

    where CC is the latent dimension, k≥2k \ge 2 is an expansion integer factor (set to k=20k=20), xtE∈{0,1}Cx^E_t \in \{0, 1\}^C is the encoder spike output, and Θq,t,Θp,t\Theta_{q,t}, \Theta_{p,t} denote the membrane potentials of neurons in fqf_q and fpf_p.

    Binary latent spikes z⋅,t,c∈{0,1}z_{\cdot, t, c} \in \{0, 1\} for each dimension c∈{1,…,C}c \in \{1, \dots, C\} are generated by uniformly selecting one spike at random from the corresponding kk-element segment:

    zq,t,c=random_select(ζq,t[k(c−1):kc])z_{q,t,c} = \text{random\_select}(\zeta_{q,t}[k(c-1) : kc])

    zp,t,c=random_select(ζp,t[k(c−1):kc])z_{p,t,c} = \text{random\_select}(\zeta_{p,t}[k(c-1) : kc])

    This sampling operation corresponds to drawing from Bernoulli distributions with parameters:

    πq,t,c=1k∑j=k(c−1)kc−1ζq,t[j]\pi_{q,t,c} = \frac{1}{k} \sum_{j=k(c-1)}^{kc-1} \zeta_{q,t}[j]

    πp,t,c=1k∑j=k(c−1)kc−1ζp,t[j]\pi_{p,t,c} = \frac{1}{k} \sum_{j=k(c-1)}^{kc-1} \zeta_{p,t}[j]

  3. Knowl 3 — Spike-to-Image Decoding via Non-Firing Output Membrane Potential

    equation

    In the Fully Spiking Variational Autoencoder (FSVAE), the output layer converts reconstructed spike trains x^1:T∈{0,1}Cout×H×W×T\hat{x}_{1:T} \in \{0, 1\}^{C_{out} \times H \times W \times T} into a continuous reconstructed image x^∈RCout×H×W\hat{x} \in \mathbb{R}^{C_{out} \times H \times W} using non-firing leaky integrate-and-fire (LIF) neurons. The terminal membrane potential uToutu^{out}_T accumulates the output spikes over all TT timesteps according to:

    uTout=∑t=1TτoutT−tx^tu^{out}_T = \sum_{t=1}^T \tau_{out}^{T-t} \hat{x}_t

    where τout=0.8\tau_{out} = 0.8 is the membrane potential decay factor. The continuous reconstructed image x^\hat{x} is calculated by:

    x^=tanh⁡(uTout)\hat{x} = \tanh(u^{out}_T)

  4. Knowl 4 — Maximum Mean Discrepancy Loss with Postsynaptic Potential Kernel

    equation

    In the Fully Spiking Variational Autoencoder (FSVAE), the prior and posterior latent distributions are aligned using the Maximum Mean Discrepancy (MMD) with a postsynaptic potential (PSP) kernel function k(z1:T,z1:T′)=∑t=1TPSP(z≤t)PSP(z≤t′)k(z_{1:T}, z'_{1:T}) = \sum_{t=1}^T \text{PSP}(z_{\le t}) \text{PSP}(z'_{\le t}). The training objective function L\mathcal{L} is:

    L=MSE(x,x^)+∑t=1T∥PSP(πq,≤t)−PSP(πp,≤t)∥2\mathcal{L} = \text{MSE}(x, \hat{x}) + \sum_{t=1}^T \|\text{PSP}(\pi_{q, \le t}) - \text{PSP}(\pi_{p, \le t})\|^2

    where MSE(x,x^)\text{MSE}(x, \hat{x}) is the mean squared error reconstruction loss between the input image xx and the decoded image x^\hat{x}, and πq,t,πp,t∈[0,1]C\pi_{q,t}, \pi_{p,t} \in [0, 1]^C are the Bernoulli probability vectors of the posterior and prior at time tt. The first-order synaptic filter PSP(v≤t)\text{PSP}(v_{\le t}) is computed iteratively for any time series v1:Tv_{1:T} via:

    PSP(v≤t)=(1−1τsyn)PSP(v≤t−1)+1τsynvt\text{PSP}(v_{\le t}) = \left(1 - \frac{1}{\tau_{syn}}\right) \text{PSP}(v_{\le t-1}) + \frac{1}{\tau_{syn}} v_t

    with initial condition PSP(v≤0)=0\text{PSP}(v_{\le 0}) = 0, where τsyn\tau_{syn} is the synaptic time constant.

  5. Knowl 5 — FSVAE Training Protocol with Teacher Forcing and Scheduled Sampling

    algorithm

    The training routine for the Fully Spiking Variational Autoencoder uses teacher forcing and linear scheduled sampling to stabilize autoregressive latent recurrence and avoid posterior collapse.

    Input: Dataset DD, total epochs E=150E = 150, batch size B=250B = 250, learning rate η=0.001\eta = 0.001, weight decay λ=0.001\lambda = 0.001, timesteps T=16T = 16, channel expansion factor k=20k = 20.
    Output: Trained FSVAE parameters.
    Initialize encoder, prior fpf_p, posterior fqf_q, and decoder SNN parameters.
    for epoch = 1 to EE do
        ϵsched←0.3×(epoch/E)\epsilon_{sched} \leftarrow 0.3 \times (\text{epoch} / E)
        for each mini-batch x∈Dx \in D do
            Encode xx to spike train x1:Tx_{1:T} via direct input encoding.
            Pass x1:Tx_{1:T} through SNN encoder to obtain x1:TEx^E_{1:T}.
            Initialize zq,0←0z_{q,0} \leftarrow 0, zp,0←0z_{p,0} \leftarrow 0, PSPq←0\text{PSP}_q \leftarrow 0, PSPp←0\text{PSP}_p \leftarrow 0, LMMD←0\mathcal{L}_{MMD} \leftarrow 0.
            for t=1t = 1 to TT do
                Compute posterior output ζq,t=fq(zq,t−1,xtE)∈{0,1}kC\zeta_{q,t} = f_q(z_{q,t-1}, x^E_t) \in \{0, 1\}^{kC}.
                Compute πq,t,c=mean(ζq,t[k(c−1):kc])\pi_{q,t,c} = \text{mean}(\zeta_{q,t}[k(c-1):kc]) for c∈{1,…,C}c \in \{1, \dots, C\}.
                Sample zq,t,c=random_select(ζq,t[k(c−1):kc])z_{q,t,c} = \text{random\_select}(\zeta_{q,t}[k(c-1):kc]).
                Sample r∼Uniform(0,1)r \sim \text{Uniform}(0, 1).
                if r<ϵschedr < \epsilon_{sched} then
                    zin,t−1←zp,t−1z_{in, t-1} \leftarrow z_{p,t-1}
                else
                    zin,t−1←zq,t−1z_{in, t-1} \leftarrow z_{q,t-1}
                Compute prior output ζp,t=fp(zin,t−1)∈{0,1}kC\zeta_{p,t} = f_p(z_{in, t-1}) \in \{0, 1\}^{kC}.
                Compute πp,t,c=mean(ζp,t[k(c−1):kc])\pi_{p,t,c} = \text{mean}(\zeta_{p,t}[k(c-1):kc]) for c∈{1,…,C}c \in \{1, \dots, C\}.
                Sample zp,t,c=random_select(ζp,t[k(c−1):kc])z_{p,t,c} = \text{random\_select}(\zeta_{p,t}[k(c-1):kc]).
                Update PSPq←(1−1/τsyn)PSPq+(1/τsyn)πq,t\text{PSP}_q \leftarrow (1 - 1/\tau_{syn})\text{PSP}_q + (1/\tau_{syn})\pi_{q,t}.
                Update PSPp←(1−1/τsyn)PSPp+(1/τsyn)πp,t\text{PSP}_p \leftarrow (1 - 1/\tau_{syn})\text{PSP}_p + (1/\tau_{syn})\pi_{p,t}.
                LMMD←LMMD+∥PSPq−PSPp∥2\mathcal{L}_{MMD} \leftarrow \mathcal{L}_{MMD} + \|\text{PSP}_q - \text{PSP}_p\|^2.
            Pass zq,1:Tz_{q,1:T} through SNN decoder to obtain x^1:T\hat{x}_{1:T}.
            Compute reconstructed image x^=tanh⁡(∑t=1TτoutT−tx^t)\hat{x} = \tanh(\sum_{t=1}^T \tau_{out}^{T-t} \hat{x}_t).
            L←MSE(x,x^)+LMMD\mathcal{L} \leftarrow \text{MSE}(x, \hat{x}) + \mathcal{L}_{MMD}.
            Backpropagate L\mathcal{L} using surrogate gradient ∂o∂u=1asign(∣u−Vth∣<a/2)\frac{\partial o}{\partial u} = \frac{1}{a} \text{sign}(|u - V_{th}| < a/2).
            Update parameters using AdamW.
  6. Knowl 6 — Image Generation Performance Benchmark across Datasets

    data/table

    Performance comparison between a standard continuous ANN VAE and the proposed Fully Spiking Variational Autoencoder (FSVAE) on MNIST (32×3232\times 32), FashionMNIST (32×3232\times 32), CIFAR10 (32×3232\times 32), and CelebA (64×6464\times 64). Evaluated metrics include Mean Squared Error Reconstruction Loss, Inception Score (IS), Fréchet Inception Distance (Inception FID), and Fréchet Distance measured on the latent space of a pretrained dataset-specific Autoencoder (Autoencoder FID) over 5,000 generated samples.

    Dataset Model Reconstruction Loss ↓\downarrow Inception Score ↑\uparrow Inception FID ↓\downarrow Autoencoder FID ↓\downarrow
    MNIST ANN 0.048 5.947 112.5 17.09
    MNIST FSVAE (Ours) 0.031 6.209 97.06 35.54
    FashionMNIST ANN 0.050 4.252 123.7 18.08
    FashionMNIST FSVAE (Ours) 0.031 4.551 90.12 15.75
    CIFAR10 ANN 0.105 2.591 229.6 196.9
    CIFAR10 FSVAE (Ours) 0.066 2.945 175.5 133.9
    CelebA ANN 0.059 3.231 92.53 156.9
    CelebA FSVAE (Ours) 0.051 3.697 101.6 112.9

    FSVAE achieves lower reconstruction errors and higher Inception Scores across all four datasets compared to the ANN VAE baseline. For FashionMNIST and CIFAR10, FSVAE outperforms the ANN VAE across all four evaluation metrics.

  7. Knowl 7 — Inference Computational Complexity Comparison

    data/table

    Computational complexity measured in floating-point addition and multiplication operations required to perform inference on a single MNIST image (32×3232\times 32) with T=16T=16 timesteps:

    Model Addition Multiplication
    ANN 7.4×1097.4 \times 10^9 7.4×1097.4 \times 10^9
    FSVAE (Ours) 5.0×10105.0 \times 10^{10} 5.6×1085.6 \times 10^8

    FSVAE requires 6.8 times more additions than the ANN VAE due to temporal unrolling across 16 timesteps, but reduces the number of multiplications by a factor of 13.2 (a 14.8×14.8\times reduction in the spiking layers). Because multiplication requires significantly more power and circuit area than addition in neuromorphic hardware, this shift reduces overall energy consumption.

  8. Knowl 8 — Ablation on Latent Distance Metric and PSP Filtering

    empirical result

    In an ablation study on the CelebA dataset (64×6464\times 64), the effect of different latent distribution alignment objectives in the ELBO loss function was evaluated using Fréchet Inception Distance (FID):

    • Kullback-Leibler Divergence (KLD) (with ϵ=0.01\epsilon = 0.01 added to πq,t\pi_{q,t} and πp,t\pi_{p,t} to prevent numerical divergence): FID =114.3= 114.3.
    • Maximum Mean Discrepancy (MMD) without postsynaptic potential filtering: FID =106.0= 106.0.
    • Maximum Mean Discrepancy (MMD) with PSP Kernel: FID =101.6= 101.6.

    The combination of MMD and postsynaptic potential (PSP) filtering yielded the lowest FID score, confirming that temporal synaptic filtering effectively captures spike train dynamics while avoiding the divergence and runaway self-excitation issues of KL divergence.

  9. Knowl 9 — Sensitivity of FSVAE Generation to Timestep Count and Channel Multiplier

    empirical result

    On the MNIST dataset, the sample generation quality (measured by FID) of the Fully Spiking Variational Autoencoder varies with the number of timesteps TT and channel multiplier kk:

    • Timestep count TT: Evaluated across T∈[4,24]T \in [4, 24], FID reaches an optimal minimum at T=16T = 16 (FID ≈97.06\approx 97.06, compared to the ANN VAE baseline of 112.5112.5). Fewer timesteps (T<16T < 16) lack expressive capacity, whereas larger values (T>16T > 16) enlarge the latent space excessively, causing generation quality to degrade.
    • Channel expansion factor kk: Evaluated across k∈[5,30]k \in [5, 30], FID consistently outperforms the ANN VAE baseline across all tested values of kk, with the best generation FID achieved at k=20k = 20.

Coverage note — None omitted; all core contributions, including the FSVAE architecture, autoregressive Bernoulli spike sampling, membrane-potential decoding, PSP-MMD loss formulation, training algorithms, and benchmark experiments have been fully covered.

Citation

MLA
Kamata, H., et al. “Fully Spiking Variational Autoencoder”. arXiv, 2021, http://arxiv.org/abs/2110.00375v3.
APA
Kamata, H., Mukuta, Y., & Harada, T. (2021). Fully Spiking Variational Autoencoder. arXiv. http://arxiv.org/abs/2110.00375v3
Chicago
Kamata, H., Y. Mukuta, and T. Harada. 2021. “Fully Spiking Variational Autoencoder”. arXiv. http://arxiv.org/abs/2110.00375v3.
Harvard
Kamata, H., Mukuta, Y. and Harada, T. (2021) “Fully Spiking Variational Autoencoder”, arXiv [Preprint]. Available at: http://arxiv.org/abs/2110.00375v3.
Vancouver
1. Kamata H, Mukuta Y, Harada T (2021) Fully Spiking Variational Autoencoder. arXiv

BibTeX

@article{kamata2021fully,
  title = {Fully Spiking Variational Autoencoder},
  author = {Kamata, Hiromichi and Mukuta, Yusuke and Harada, Tatsuya},
  year = {2021},
  journal = {arXiv},
  url = {http://arxiv.org/abs/2110.00375v3},
  eprint = {2110.00375}
}
Metadata:arXiv

Access the Paper

This paper is available from its original source. Click below to access the PDF.

Open PDF