| name | psvit-structured-pruning-spiking-vision |
| description | Structured pruning methodology for Spiking Vision Transformers (SViT) using uniform channel-wise filter pruning and sensitivity analysis for 22.4% memory saving |
| version | 1.0.0 |
| category | ai_collection |
| tags | ["deep-learning","neuromorphic","SNN","pruning","efficiency","vision"] |
| arxiv | 2606.03257v1 |
| paper_title | PSViT: A Methodology for Structurally Pruning Spiking Vision Transformers |
| authors | ["Rachmad Vidya Wicaksana Putra","Achyuta Muthuvelan","Alberto Marchisio","Muhammad Shafique"] |
| published | 2026-06-02T00:00:00.000Z |
| activation_keywords | ["spiking neural network","vision transformer pruning","structured pruning","neuromorphic efficiency","SViT","channel-wise pruning"] |
PSViT: Structured Pruning for Spiking Vision Transformers
Core Innovation
Structured pruning for Spiking Vision Transformers (SViT) enabling efficient acceleration on existing hardware (vs. unstructured pruning requiring specialized architectures).
Problem Addressed
Unstructured pruning limitations:
- Requires specialized hardware for sparsity patterns
- Not scalable for widespread deployment
- Hardware-dependent efficiency gains
Methodology
Three-Stage Pruning Pipeline
- Uniform channel-wise filter pruning: Structurally eliminate non-significant weights
- Sensitivity analysis: Evaluate pruning impact per layer on accuracy/size
- Fine-grained channel-wise pruning: Layer-specific pruning based on sensitivity
Key Advantages
- Hardware-agnostic: Works on standard computing architectures
- Structured sparsity: Regular patterns for efficient execution
- Accuracy preservation: Maintains high performance with fine-tuning
Performance Results
- Memory saving: 22.4% through single-shot pruning
- Accuracy:
- Without fine-tuning: 70.3% (3.0% drop)
- With fine-tuning: 72.8% (0.5% drop)
- Original SViT: 73.3%
- Dataset: ImageNet-1K
Implementation Pattern
import torch
class PSViTPruner:
def __init__(self, svit_model, target_reduction=0.224):
self.model = svit_model
self.target_reduction = target_reduction
def uniform_channel_pruning(self, layer, pruning_ratio):
"""Stage 1: Uniform channel-wise filter pruning"""
weight = layer.weight.data
num_channels = weight.shape[0]
prune_count = (num_channels * pruning_ratio)
channel_importance = torch.norm(weight, p=, dim=(, , ))
prune_channels = torch.argsort(channel_importance)[:prune_count]
layer.weight.data[prune_channels] =
prune_channels
():
sensitivities = {}
name, layer .model.named_modules():
(layer, ):
sensitivity = .measure_layer_sensitivity(layer)
sensitivities[name] = sensitivity
sensitivities
():
pruned_layers = {}
name, layer .model.named_modules():
(layer, ):
sensitivity = sensitivities[name]
pruning_ratio = .compute_pruning_ratio(sensitivity)
pruned_layers[name] = .uniform_channel_pruning(layer, pruning_ratio)
pruned_layers
():
original_accuracy = .evaluate_model()
.prune_temporarily(layer, ratio=)
pruned_accuracy = .evaluate_model()
.restore_layer(layer)
original_accuracy - pruned_accuracy