Manages multi-turn state (KV cache, positions, per-conversation batching) for a converted SNN.
Note: this is the state-management mechanism; conversational quality under genuine spiking
requires training (see coverage-quality.md).
class TemporalSpikeProcessor(nn.Module):
def __init__(self, snn_model, T=16, max_context_length=512):
"""
Initialize the temporal spike processor.
Args:
snn_model: The converted SNN model
T: Number of timesteps for spike processing
max_context_length: Maximum sequence length
"""Process input through the SNN with temporal dynamics.
Parameters:
input_ids(torch.Tensor): Input token IDsattention_mask(torch.Tensor, optional): Attention maskuse_cache(bool): Whether to use KV cache for multi-turn**kwargs: Additional model arguments
Returns:
- Model output with logits and optional past key values
Reset the KV cache for new conversations.
Parameters:
batch_id(int, optional): Specific batch to reset
Get the last computed position IDs for the conversation.
Returns:
torch.Tensorof shape[1, seq_len]containing the most recent position IDs (a zero tensor if none have been computed yet).
Spiking-compatible attention mechanism.
class SpikeAttention(nn.Module):
def __init__(self, embed_dim, num_heads, T=16, causal=True):
"""
Initialize spike-based attention.
Args:
embed_dim: Embedding dimension
num_heads: Number of attention heads
T: Number of timesteps
causal: Whether to use causal attention
"""Spiking-compatible layer normalization.
class SpikeLayerNorm(nn.Module):
def __init__(self, normalized_shape, eps=1e-5):
"""
Initialize spike-compatible layer normalization.
Args:
normalized_shape: Input shape to normalize
eps: Small constant for numerical stability
"""Replace GELU activations with ReLU for SNN compatibility.
Parameters:
model(torch.nn.Module): Model to modify
Returns:
- Modified model with ReLU activations
Perform simplified ANN→SNN conversion. Returns a TemporalSpikeProcessor wrapping the
converted model.
Parameters:
model(torch.nn.Module): Source modeltimesteps(int): Number of SNN timestepsskip_gelu_replacement(bool): IfTrue, skip the GELU→ReLU substitution. This keeps the path closer to the ANN (and, combined with spiking off, faithful to it); setFalsefor neuromorphic deployment, where quality then depends on training, not conversion alone.real_spiking(bool): IfTrue,SpikeAttentionroutes Q/K/V through its LIF neurons and drops softmax (genuine spiking self-attention), so the T-timestep loop stops being a no-op. DefaultFalsereproduces the source model exactly; enabling it changes the model's outputs and should be measured (seespike_metrics.py) before it is relied on.
Returns:
TemporalSpikeProcessorwrapping the converted SNN model
Replace LayerNorm with SpikeLayerNorm.
Parameters:
model(torch.nn.Module): Model to modify
Returns:
- Modified model with spike-compatible normalization
Replace standard attention with SpikeAttention.
Parameters:
model(torch.nn.Module): Model to modifyspiking(bool): IfTrue, the installedSpikeAttentionroutes Q/K/V through its LIF neurons; defaultFalseconstructs the modules but bypasses the neurons (structural spiking only), matching the defaultsimplified_conversionbehaviour.
Returns:
- Modified model with spike-compatible attention
Apply surrogate gradient functions for SNN training.
Parameters:
model(torch.nn.Module): SNN modelalpha(float): Surrogate gradient scaling factor
Returns:
- Model with surrogate gradients
Calibrate spike timing for different timestep counts.
Parameters:
model(torch.nn.Module): SNN modeloriginal_T(int): Original timestep counttarget_T(int): Target timestep count
Returns:
- Calibrated model
Save the converted SNN model with metadata. (smollm2_converter version.)
Parameters:
model(torch.nn.Module): SNN model to savetokenizer: Associated tokenizerpath(str): Save directory
Returns:
- Success status
Note:
convert.pyexports a different function of the same name,save_snn_model(model, path, timesteps=None, simplified=True), which takes a file path and no tokenizer. Import the one matching the module you converted with.
Create calibration data for SNN conversion.
Parameters:
tokenizer: HuggingFace tokenizernum_samples(int): Number of calibration samplesmax_length(int): Maximum sequence length
Returns:
- Dictionary with calibration data
Test position ID handling at sequence boundaries.
Parameters:
model: SNN model to testtokenizer: Associated tokenizerargs: Test configuration
Returns:
- Test results
Test attention mask continuity across conversation turns.
Parameters:
model: SNN model to testtokenizer: Associated tokenizerargs: Test configuration
Returns:
- Test results
Test multi-turn conversation coherence.
Parameters:
model: SNN model to testtokenizer: Associated tokenizerargs: Test configuration
Returns:
- Test results
Simulate a multi-turn conversation for testing.
Parameters:
model: SNN modeltokenizer: Associated tokenizerturns(int): Number of conversation turnsdevice(str): Computing device
Returns:
- Conversation results
Main CLI tool for model conversion.
Usage:
python scripts/run_conversion.py [OPTIONS]Options:
--model_name: Model to convert (distilgpt2,HuggingFaceTB/SmolLM2-1.7B-Instruct)--output_dir: Output directory--timesteps: Number of SNN timesteps--simplified: Use simplified conversion--verify: Reload the saved weights into the base model and run a forward pass--quantize: Load with 8-bit quantization before converting (requiresbitsandbytes)--num_samples: Calibration samples (default 3)--calibration_batch_size: Calibration batch size (default 1)--optimize_for_torchscript: Also export a TorchScript artifact next to the model
Flags accepted for CLI compatibility but not applied by this runner (it warns when
they are set): --use_sparse, --use_delayed_spikes, --use_function_calling, and
--surrogate_function values other than the default.
Testing and validation tool.
Usage:
python tests/test_conversational_snn.py --model_name distilgpt2 [OPTIONS]Options:
--test_all: Run all tests--test_position_boundaries: Test position ID boundaries--test_attention_mask: Test attention mask continuity--test_multi_turn: Test multi-turn capabilities--test_energy: Test energy proxy (software profiling / spike-count telemetry; not hardware watt-hour measurements)
Supported Models:
distilgpt2: DistilGPT-2 (82M parameters)HuggingFaceTB/SmolLM2-1.7B-Instruct: SmolLM2 1.7B Instruct (1.7B parameters)
Conversion Parameters:
timesteps: 8-64 (recommended: 16)max_context_length: 512-2048 (recommended: 512)surrogate_function: atan, sigmoid, stbif_plus
GPU Memory Requirements:
- DistilGPT-2: 4-8 GB
- SmolLM2-1.7B-Instruct: 20 GB
CPU Requirements:
- Multi-core processor recommended
- 16-32 GB RAM
ImportError: SpikingJelly version compatibility
# Ensure SpikingJelly >= 0.0.0.0.14
pip install spikingjelly[cuda] -U --preCUDA Out of Memory: Insufficient GPU memory
# Reduce batch size or use CPU
device = 'cpu'Position ID Errors: Sequence length exceeds model limits
# Reduce max_context_length
max_context_length = 512from smollm2_converter import *
# Load model
model = AutoModelForCausalLM.from_pretrained("distilgpt2")
tokenizer = AutoTokenizer.from_pretrained("distilgpt2")
# Convert to SNN. This ALREADY returns a TemporalSpikeProcessor — do not wrap it again.
# Re-wrapping nests T x T timestep loops, applies the logit scaling twice, and makes the
# inner processor's KV cache grow by T positions per call until it exceeds the model's
# max_position_embeddings and raises "IndexError: index out of range in self".
processor = simplified_conversion(model, timesteps=16)
# Test conversation
result = simulate_conversation(processor, tokenizer, turns=3)# Full pipeline conversion
from convert import convert_model_to_spiking, create_calibration_data, was_simplified
from convert import save_snn_model as save_converted_model
from smollm2_converter import apply_surrogate_gradients
# Create calibration data
calib_data = create_calibration_data(tokenizer, num_samples=10)
# Convert with calibration. SpikingJelly's ann2snn requires a torch.fx-traceable model;
# HuggingFace causal LMs generally are not, so this falls back to the simplified path.
# Check which one you actually got before reporting results.
snn_model = convert_model_to_spiking(model, calib_data, timesteps=32)
print("simplified fallback used:", was_simplified(snn_model))
# Apply surrogate gradients
snn_model = apply_surrogate_gradients(snn_model, alpha=4.0)
# Save model. NOTE: the two modules have different signatures —
# convert.save_snn_model(model, path, timesteps=None, simplified=True)
# smollm2_converter.save_snn_model(model, tokenizer, path)
save_converted_model(snn_model, "./my_snn_model/snn_model.pt", timesteps=32)