spikingjelly 2.0.0.dev0__tar.gz → 2.0.0.dev1__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/PKG-INFO +10 -3
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/README.md +2 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/pyproject.toml +33 -3
- spikingjelly-2.0.0.dev1/spikingjelly/__init__.py +7 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/__init__.py +20 -17
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/converter.py +69 -29
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/delay.py +45 -32
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/bert_sst2_transformer_td_equivalent.py +4 -2
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/cnn_mnist.py +14 -18
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/imagenet_resnet18_ltb.py +2 -18
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/resnet18_cifar10.py +2 -5
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/operators.py +560 -153
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/ann2snn/qcfs.py +458 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/__init__.py +12 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/base.py +39 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/ann2snn/recipes/local_threshold_balancing.py +325 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/ann2snn/recipes/qwen2.py +1360 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/rate_coding.py +332 -397
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/spikezip_qann.py +313 -326
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/sta_transformer.py +83 -540
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/ann2snn/recipes/step_mode_adapters.py +636 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/transformer_td_equivalent.py +37 -65
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/sample_models/cifar10_resnet.py +6 -6
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/utils.py +34 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/base.py +421 -81
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/__init__.py +38 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/auto_cuda/__init__.py +1 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/auto_cuda/base.py +31 -54
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/auto_cuda/cfunction.py +14 -7
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/auto_cuda/generator.py +19 -16
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/cuda_utils.py +60 -151
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/__init__.py +1 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/cuda_code.py +27 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step/__init__.py +21 -0
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/neuron_kernel/common.py → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step/base.py +25 -178
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step}/eif.py +40 -190
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/neuron_kernel → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step}/integrate_and_fire.py +26 -37
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step}/izhikevich.py +113 -187
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/neuron_kernel → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step}/lif.py +35 -49
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/neuron_kernel → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step}/plif.py +33 -141
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step}/qif.py +40 -190
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step/runtime.py +84 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/single_step/__init__.py +1 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/single_step/base.py +254 -0
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/ss_neuron_kernel → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/single_step}/integrate_and_fire.py +79 -36
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/ss_neuron_kernel → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/single_step}/lif.py +89 -37
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/surrogate_registry.py +44 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_linear.py +708 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/spike_linear.py +933 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/spike_op.py +23 -11
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/tensor_cache.py +23 -24
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/distributed/__init__.py +31 -18
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/distributed/adapters/__init__.py +10 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/distributed/adapters/base.py +16 -19
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/adapters/cifar10dvs_vgg.py +48 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/adapters/spikformer.py +48 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/analysis.py +160 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/api.py +310 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/config.py +114 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/data_parallel.py +54 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/execution.py +227 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/fsdp.py +105 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/mesh.py +259 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/distributed/metrics.py +1 -1
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/optimizer.py +95 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/__init__.py +21 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/cifar10dvs_vgg.py +148 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/memopt.py +183 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/partition.py +166 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/runtime.py +486 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/spikformer.py +190 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/planner.py +490 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/runtime.py +160 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/__init__.py +41 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/channel.py +435 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/cifar10dvs_vgg.py +146 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/debug.py +63 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/linear.py +392 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/spikformer.py +223 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/state.py +166 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/utils.py +50 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/distributed/topology.py +22 -40
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/encoding.py +51 -93
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/A2C.py +0 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DQN_state.py +1 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/agent.py +10 -13
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/experience.py +0 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/train.py +0 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/NoisySAN/hybrid_td3_cuda_norm.py +0 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/NoisySAN/test_hybrid_td3_cpu.py +0 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/Spiking_A2C.py +5 -31
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/Spiking_DQN_state.py +9 -29
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/Spiking_PPO.py +5 -30
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/cifar10_r11_enabling_spikebased_backpropagation.py +2 -4
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/common/multiprocessing_env.py +0 -4
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/lava_mnist.py +1 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/train.py +0 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/train_distributed.py +22 -23
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/mstdp.py +0 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/mstdpet.py +0 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/speechcommands.py +0 -8
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/spiking_lstm_sequential_mnist.py +1 -13
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/spiking_lstm_text.py +1 -31
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/stdp_trace.py +0 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/__init__.py +3 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/conv_bn_fusion.py +24 -30
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/forward.py +94 -72
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/functional/layer.py +200 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/functional/learning.py +602 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/misc.py +11 -11
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/net_config.py +37 -56
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/functional/neuron.py +2994 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/online_learning.py +46 -73
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/lava_exchange.py +114 -134
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/attention.py +58 -86
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/bn.py +43 -15
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/container.py +67 -49
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/dropout.py +75 -68
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/layer/misc.py +629 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/online_learning.py +25 -30
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/stateless_wrapper.py +331 -208
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/learning.py +182 -232
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/lynxi_exchange.py +12 -18
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/memopt/checkpointing.py +46 -15
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/memopt/compress.py +4 -11
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/memopt/pipeline.py +82 -84
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/parametric_lif_net.py +16 -17
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/snas_net.py +7 -9
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/spike_dhs.py +78 -48
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/spiking_vggws_ottt.py +13 -34
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/train_classify.py +121 -96
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/model/train_flexsn_inductor_example.py → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/model/train_flexsn_compile_example.py +12 -11
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/train_imagenet_example.py +1 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/tv_ref_classify/utils.py +55 -60
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/monitor.py +76 -84
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/__init__.py +1 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/adapt.py +87 -140
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/base_node.py +249 -204
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/dsr.py +35 -65
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/few_spike.py +25 -42
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/flexsn.py +351 -300
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/neuron/ilif.py +410 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/integrate_and_fire.py +525 -376
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/neuron/lif.py +397 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/lif_variants.py +174 -201
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/mpbn.py +152 -225
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/noisy.py +76 -175
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/nonlinear_if.py +113 -58
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/online_learning.py +61 -202
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/plif.py +73 -147
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/psn.py +55 -43
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/neuron/spikezip.py +223 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/nir_exchange/from_nir.py +9 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/nir_exchange/to_nir.py +9 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/ac.py +17 -28
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/analytical_energy/core.py +0 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/base.py +104 -52
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/compute_energy.py +23 -19
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/flop.py +6 -12
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/lemaire_addressing.py +8 -30
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/mac.py +3 -9
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/memory_access.py +10 -25
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/memory_residency.py +16 -142
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/add_counter.py +8 -26
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/base_counter.py +17 -11
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/cmp_counter.py +2 -2
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/config.py +11 -43
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/core.py +149 -221
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/memory_residency_counter.py +0 -15
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/mul_counter.py +1 -15
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/op_counter/neuromc/utils.py +59 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuron_state.py +26 -6
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/spikesim/__init__.py +0 -4
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/spikesim/core.py +9 -34
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/spikesim/counter.py +4 -28
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/synop.py +12 -140
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/__init__.py +20 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/api.py +50 -7
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/capability.py +167 -1
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/precision/config.py +46 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/convert.py +22 -2
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/precision/float8_attention.py +175 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/float8_base.py +24 -19
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/precision/float8_conv.py +128 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/precision/float8_te.py +595 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/float8_torchao.py +19 -43
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/policy.py +7 -14
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/runtime.py +7 -8
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/profiler.py +63 -55
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/quantize.py +43 -7
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/rnn.py +119 -191
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/surrogate.py +135 -296
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/compress.py +3 -3
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/__init__.py +4 -4
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/custom_ops.py +87 -88
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/hop.py +2 -2
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/kernel.py +4 -4
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/template.py +36 -63
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/wrapper.py +5 -6
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/fp8_capability.py +412 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/__init__.py +15 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/activation_aware_if.py +281 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/ilif.py +724 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/integrate_and_fire.py +1225 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/lif.py +1333 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/plif.py +1337 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/stbif.py +470 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/utils.py +209 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/surrogate_kernel.py +3 -4
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/torch2triton/__init__.py +9 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/torch2triton/graph2triton.py +14 -30
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/torch2triton/torch2graph.py +2 -0
- spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/triton_utils.py +458 -0
- spikingjelly-2.0.0.dev1/spikingjelly/configure.py +177 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/asl_dvs.py +16 -14
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/base.py +114 -151
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/bullying10k.py +22 -18
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/cifar10_dvs.py +21 -22
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/dvs128_gesture.py +28 -23
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/dvs_lip.py +2 -1
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/es_imagenet.py +26 -183
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/hardvs.py +12 -9
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/n_caltech101.py +11 -12
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/n_mnist.py +12 -13
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/nav_gesture.py +42 -68
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/shd.py +63 -169
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/speechcommands.py +7 -6
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/utils.py +58 -415
- spikingjelly-2.0.0.dev1/spikingjelly/logger.py +38 -0
- spikingjelly-2.0.0.dev1/spikingjelly/timing_based/encoding.py +134 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/timing_based/examples/tempotron_mnist.py +0 -3
- spikingjelly-2.0.0.dev1/spikingjelly/timing_based/neuron.py +192 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly.egg-info/PKG-INFO +10 -3
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly.egg-info/SOURCES.txt +58 -28
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly.egg-info/requires.txt +10 -2
- spikingjelly-2.0.0.dev1/test/test_configure.py +127 -0
- spikingjelly-2.0.0.dev1/test/test_dataset_builders.py +144 -0
- spikingjelly-2.0.0.dev1/test/test_dataset_utils.py +197 -0
- spikingjelly-2.0.0.dev1/test/test_logging_policy.py +264 -0
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/factories.py +0 -243
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/recipes/local_threshold_balancing.py +0 -461
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/recipes/step_mode_adapters.py +0 -1255
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/rules.py +0 -380
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/threshold.py +0 -90
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/__init__.py +0 -42
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/example.py +0 -35
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/neuron_kernel/__init__.py +0 -4
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/ss_neuron_kernel/__init__.py +0 -35
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/ss_neuron_kernel/ss_neuron_kernel_base.py +0 -522
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/__init__.py +0 -37
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/common.py +0 -470
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/helpers.py +0 -51
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/integrate_and_fire.py +0 -794
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/lif.py +0 -864
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/plif.py +0 -823
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/adapters/cifar10dvs_vgg.py +0 -57
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/adapters/spikformer.py +0 -61
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/api.py +0 -210
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/dtensor.py +0 -3434
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/planner.py +0 -38
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/runtime.py +0 -264
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/layer/misc.py +0 -378
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/neuron/lif.py +0 -708
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/neuron/spikezip.py +0 -256
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/neuron_cupy.py +0 -3180
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/neuron_cupy_lite.py +0 -2867
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/op_counter/neuromc/utils.py +0 -179
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/precision/config.py +0 -83
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/neuron_kernel/__init__.py +0 -11
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/neuron_kernel/integrate_and_fire.py +0 -679
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/neuron_kernel/lif.py +0 -736
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/neuron_kernel/plif.py +0 -784
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/spikezip_kernel.py +0 -226
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/torch2triton/__init__.py +0 -9
- spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/triton_utils.py +0 -188
- spikingjelly-2.0.0.dev0/spikingjelly/configure.py +0 -85
- spikingjelly-2.0.0.dev0/spikingjelly/timing_based/__init__.py +0 -0
- spikingjelly-2.0.0.dev0/spikingjelly/timing_based/encoding.py +0 -386
- spikingjelly-2.0.0.dev0/spikingjelly/timing_based/examples/__init__.py +0 -0
- spikingjelly-2.0.0.dev0/spikingjelly/timing_based/neuron.py +0 -362
- spikingjelly-2.0.0.dev0/spikingjelly/timing_based/orig_encoding.py +0 -57
- spikingjelly-2.0.0.dev0/spikingjelly/timing_based/orig_neuron.py +0 -127
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/LICENSE +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/LICENSES/translations/LICENSE.de +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/LICENSES/translations/LICENSE.en +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/LICENSES/translations/LICENSE.fr +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/LICENSES/translations/LICENSE.hi +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/setup.cfg +0 -0
- {spikingjelly-2.0.0.dev0/spikingjelly → spikingjelly-2.0.0.dev1/spikingjelly/activation_based}/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/ann2snn/examples}/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/imagenet_vit_sta.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/roberta_spikezip_qann_synthetic.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/modules.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/sample_models/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/sample_models/mnist_cnn.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/actions.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/common/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/common/runfile.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/common/utils.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/common/wrappers.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/common/wrappers_simple.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/ignite.py +0 -0
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/examples → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/examples/DSQN/utils}/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/utils/atari_wrappers.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/utils/common.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/utils/model.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/ILC-SAN/core_cuda.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/ILC-SAN/hybrid_td3_cuda_norm.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/ILC-SAN/ilcsan.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/ILC-SAN/replay_buffer_norm.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/ILC-SAN/test_hybrid_td3_cpu.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/NoisySAN/core_cuda.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/NoisySAN/noisysan.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/NoisySAN/replay_buffer_norm.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/PPO.py +0 -0
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/examples}/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/classify_dvsg.py +0 -0
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/examples/DSQN/utils → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/examples/common}/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/conv_fashion_mnist.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/lif_fc_mnist.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/lynxi_fmnist_inference.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/data_module.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/lightning_callbacks.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/lightning_modules.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/loss.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/models.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/rsnn_sequential_fmnist.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/loss.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/memopt/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/sew_resnet.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/spikformer.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/spiking_resnet.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/spiking_vgg.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/tv_ref_classify/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/tv_ref_classify/presets.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/tv_ref_classify/sampler.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/tv_ref_classify/transforms.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/inter_layer_connection.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/nir_exchange/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/_sparse_memory.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/analytical_energy/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/mux_counter.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/sqrt_counter.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/spikesim/config.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/spikesim/formulas.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/dummy.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/info.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/transform.py +0 -0
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/examples → spikingjelly-2.0.0.dev1/spikingjelly/timing_based}/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/examples/common → spikingjelly-2.0.0.dev1/spikingjelly/timing_based/examples}/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/__init__.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/_utils.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/bar3d.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/feature_map.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/heatmap.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/spikes.py +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly.egg-info/dependency_links.txt +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly.egg-info/top_level.txt +0 -0
- {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/test/test_visualizing.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: spikingjelly
|
|
3
|
-
Version: 2.0.0.
|
|
3
|
+
Version: 2.0.0.dev1
|
|
4
4
|
Summary: A deep learning framework for SNNs built on PyTorch.
|
|
5
5
|
Author: PKU MLG, PCL, and other contributors
|
|
6
6
|
Author-email: fwei@pku.edu.cn, chyq@pku.edu.cn, yfhuang@pku.edu.cn
|
|
@@ -20,12 +20,11 @@ License-File: LICENSES/translations/LICENSE.hi
|
|
|
20
20
|
Requires-Dist: torch>=2.6.0
|
|
21
21
|
Requires-Dist: torchvision
|
|
22
22
|
Requires-Dist: torchaudio
|
|
23
|
-
Requires-Dist: einops
|
|
24
23
|
Requires-Dist: h5py
|
|
24
|
+
Requires-Dist: loguru<1,>=0.7
|
|
25
25
|
Requires-Dist: matplotlib
|
|
26
26
|
Requires-Dist: numpy
|
|
27
27
|
Requires-Dist: packaging
|
|
28
|
-
Requires-Dist: pydantic
|
|
29
28
|
Requires-Dist: requests
|
|
30
29
|
Requires-Dist: SciencePlots
|
|
31
30
|
Requires-Dist: scipy
|
|
@@ -43,6 +42,12 @@ Requires-Dist: nirtorch; extra == "nir"
|
|
|
43
42
|
Provides-Extra: lightning
|
|
44
43
|
Requires-Dist: lightning; extra == "lightning"
|
|
45
44
|
Requires-Dist: jsonargparse[signatures]; extra == "lightning"
|
|
45
|
+
Provides-Extra: fp8-torchao
|
|
46
|
+
Requires-Dist: torchao<0.18,>=0.13; extra == "fp8-torchao"
|
|
47
|
+
Provides-Extra: fp8-te
|
|
48
|
+
Requires-Dist: transformer-engine[pytorch]<3,>=2.16; extra == "fp8-te"
|
|
49
|
+
Provides-Extra: qwen
|
|
50
|
+
Requires-Dist: transformers==5.13.0; extra == "qwen"
|
|
46
51
|
Dynamic: license-file
|
|
47
52
|
|
|
48
53
|
# SpikingJelly
|
|
@@ -208,6 +213,8 @@ generations, `MINOR` adds backward-compatible functionality, and `PATCH` fixes
|
|
|
208
213
|
bugs. Python package pre-release spelling is used for V2 development releases,
|
|
209
214
|
for example `2.0.0.dev0`, `2.0.0a1`, `2.0.0b1`, and `2.0.0rc1`.
|
|
210
215
|
|
|
216
|
+
See [CHANGELOG.md](./CHANGELOG.md) for the V2 release changelog.
|
|
217
|
+
|
|
211
218
|
<details>
|
|
212
219
|
<summary>Compatibility, migration, and older docs</summary>
|
|
213
220
|
|
|
@@ -161,6 +161,8 @@ generations, `MINOR` adds backward-compatible functionality, and `PATCH` fixes
|
|
|
161
161
|
bugs. Python package pre-release spelling is used for V2 development releases,
|
|
162
162
|
for example `2.0.0.dev0`, `2.0.0a1`, `2.0.0b1`, and `2.0.0rc1`.
|
|
163
163
|
|
|
164
|
+
See [CHANGELOG.md](./CHANGELOG.md) for the V2 release changelog.
|
|
165
|
+
|
|
164
166
|
<details>
|
|
165
167
|
<summary>Compatibility, migration, and older docs</summary>
|
|
166
168
|
|
|
@@ -10,7 +10,7 @@ include = ["spikingjelly*"]
|
|
|
10
10
|
|
|
11
11
|
[project]
|
|
12
12
|
name = "spikingjelly"
|
|
13
|
-
version = "2.0.0.
|
|
13
|
+
version = "2.0.0.dev1"
|
|
14
14
|
description = "A deep learning framework for SNNs built on PyTorch."
|
|
15
15
|
readme = {file = "README.md", content-type = "text/markdown"}
|
|
16
16
|
authors = [
|
|
@@ -24,12 +24,11 @@ dependencies = [
|
|
|
24
24
|
"torch>=2.6.0",
|
|
25
25
|
"torchvision",
|
|
26
26
|
"torchaudio",
|
|
27
|
-
"einops",
|
|
28
27
|
"h5py",
|
|
28
|
+
"loguru>=0.7,<1",
|
|
29
29
|
"matplotlib",
|
|
30
30
|
"numpy",
|
|
31
31
|
"packaging",
|
|
32
|
-
"pydantic",
|
|
33
32
|
"requests",
|
|
34
33
|
"SciencePlots",
|
|
35
34
|
"scipy",
|
|
@@ -50,13 +49,27 @@ cupy12 = ["cupy-cuda12x"]
|
|
|
50
49
|
triton = ["triton>=3.3.1"]
|
|
51
50
|
nir = ["nir", "nirtorch"]
|
|
52
51
|
lightning = ["lightning", "jsonargparse[signatures]"]
|
|
52
|
+
fp8-torchao = [
|
|
53
|
+
"torchao>=0.13,<0.18",
|
|
54
|
+
]
|
|
55
|
+
fp8-te = [
|
|
56
|
+
"transformer-engine[pytorch]>=2.16,<3",
|
|
57
|
+
]
|
|
58
|
+
qwen = [
|
|
59
|
+
"transformers==5.13.0",
|
|
60
|
+
]
|
|
53
61
|
|
|
54
62
|
[dependency-groups]
|
|
55
63
|
dev = [
|
|
56
64
|
"pytest",
|
|
65
|
+
"pytest-cov",
|
|
57
66
|
"pydata-sphinx-theme>=0.16.1",
|
|
58
67
|
"sphinx>=8.0.0,<9",
|
|
59
68
|
]
|
|
69
|
+
llm-benchmark = [
|
|
70
|
+
"datasets>=5,<6",
|
|
71
|
+
"lm-eval[hf]==0.4.12",
|
|
72
|
+
]
|
|
60
73
|
|
|
61
74
|
[project.urls]
|
|
62
75
|
Homepage = "https://github.com/fangwei123456/spikingjelly"
|
|
@@ -73,3 +86,20 @@ quote-style = "double"
|
|
|
73
86
|
|
|
74
87
|
[tool.ruff.lint]
|
|
75
88
|
ignore = ["E402", "E731", "E741", "F403"]
|
|
89
|
+
|
|
90
|
+
[tool.ruff.lint.isort]
|
|
91
|
+
known-first-party = ["benchmark", "spikingjelly", "test"]
|
|
92
|
+
|
|
93
|
+
[tool.pytest.ini_options]
|
|
94
|
+
pythonpath = ["."]
|
|
95
|
+
testpaths = ["test"]
|
|
96
|
+
|
|
97
|
+
[tool.coverage.run]
|
|
98
|
+
branch = true
|
|
99
|
+
omit = [
|
|
100
|
+
"*/examples/*",
|
|
101
|
+
]
|
|
102
|
+
source = ["spikingjelly"]
|
|
103
|
+
|
|
104
|
+
[tool.coverage.report]
|
|
105
|
+
show_missing = true
|
|
@@ -9,10 +9,10 @@ r"""
|
|
|
9
9
|
|
|
10
10
|
ANN 到 SNN 的转换模块。提供 FX graph 路径的 :class:`FXConverter` /
|
|
11
11
|
:class:`FXConversionRecipe`、module tree 路径的 :class:`ModuleConverter` /
|
|
12
|
-
:class:`ModuleConversionRecipe`,以及 ``
|
|
13
|
-
|
|
14
|
-
|
|
15
|
-
:class:`
|
|
12
|
+
:class:`ModuleConversionRecipe`,以及 ``download_url`` 工具函数。兼容名
|
|
13
|
+
:class:`Converter` 等价于
|
|
14
|
+
:class:`FXConverter`,:class:`ConversionRecipe` 等价于
|
|
15
|
+
:class:`FXConversionRecipe`。
|
|
16
16
|
|
|
17
17
|
----
|
|
18
18
|
|
|
@@ -22,30 +22,31 @@ ANN 到 SNN 的转换模块。提供 FX graph 路径的 :class:`FXConverter` /
|
|
|
22
22
|
|
|
23
23
|
ANN-to-SNN conversion module. Provides the FX graph :class:`FXConverter` /
|
|
24
24
|
:class:`FXConversionRecipe`, the module-tree :class:`ModuleConverter` /
|
|
25
|
-
:class:`ModuleConversionRecipe`,
|
|
26
|
-
|
|
27
|
-
:class:`
|
|
28
|
-
|
|
29
|
-
:class:`FXConverter`, and :class:`ConversionRecipe` is equivalent to
|
|
30
|
-
:class:`FXConversionRecipe`.
|
|
25
|
+
:class:`ModuleConversionRecipe`, and a ``download_url`` helper for fetching
|
|
26
|
+
pretrained models. The compatibility name
|
|
27
|
+
:class:`Converter` is equivalent to :class:`FXConverter`, and
|
|
28
|
+
:class:`ConversionRecipe` is equivalent to :class:`FXConversionRecipe`.
|
|
31
29
|
"""
|
|
32
30
|
|
|
33
31
|
from .converter import Converter, FXConverter, ModuleConverter
|
|
34
32
|
from .delay import estimate_delay_start
|
|
35
|
-
from .factories import HookFactory, NeuronFactory
|
|
36
33
|
from .modules import ChannelVoltageScaler
|
|
34
|
+
from .qcfs import SignedQCFSSequenceEncoder
|
|
37
35
|
from .recipes import (
|
|
38
36
|
ConversionRecipe,
|
|
39
37
|
FXConversionRecipe,
|
|
40
38
|
LocalThresholdBalancingRecipe,
|
|
41
39
|
ModuleConversionRecipe,
|
|
40
|
+
Qwen2SNNCalibration,
|
|
41
|
+
Qwen2SNNConfig,
|
|
42
|
+
Qwen2SNNModel,
|
|
43
|
+
Qwen2SNNRecipe,
|
|
42
44
|
RateCodingRecipe,
|
|
43
45
|
SpikeZIPTFQANNRecipe,
|
|
44
46
|
STATransformerRecipe,
|
|
45
47
|
TransformerTDEquivalentRecipe,
|
|
48
|
+
calibrate_qwen2_snn,
|
|
46
49
|
)
|
|
47
|
-
from .rules import ReLURule
|
|
48
|
-
from .threshold import ThresholdOptimizer
|
|
49
50
|
from .utils import download_url
|
|
50
51
|
|
|
51
52
|
__all__ = [
|
|
@@ -55,16 +56,18 @@ __all__ = [
|
|
|
55
56
|
"ConversionRecipe",
|
|
56
57
|
"FXConversionRecipe",
|
|
57
58
|
"ModuleConversionRecipe",
|
|
59
|
+
"Qwen2SNNCalibration",
|
|
60
|
+
"Qwen2SNNConfig",
|
|
61
|
+
"Qwen2SNNModel",
|
|
62
|
+
"Qwen2SNNRecipe",
|
|
63
|
+
"calibrate_qwen2_snn",
|
|
58
64
|
"RateCodingRecipe",
|
|
59
65
|
"LocalThresholdBalancingRecipe",
|
|
60
66
|
"SpikeZIPTFQANNRecipe",
|
|
61
67
|
"STATransformerRecipe",
|
|
62
68
|
"TransformerTDEquivalentRecipe",
|
|
63
69
|
"ChannelVoltageScaler",
|
|
70
|
+
"SignedQCFSSequenceEncoder",
|
|
64
71
|
"estimate_delay_start",
|
|
65
72
|
"download_url",
|
|
66
|
-
"ReLURule",
|
|
67
|
-
"NeuronFactory",
|
|
68
|
-
"HookFactory",
|
|
69
|
-
"ThresholdOptimizer",
|
|
70
73
|
]
|
|
@@ -1,4 +1,6 @@
|
|
|
1
|
-
import
|
|
1
|
+
import threading
|
|
2
|
+
import time
|
|
3
|
+
import types
|
|
2
4
|
from typing import Optional, Union
|
|
3
5
|
|
|
4
6
|
import torch
|
|
@@ -10,6 +12,45 @@ from spikingjelly.activation_based.ann2snn.recipes import (
|
|
|
10
12
|
ModuleConversionRecipe,
|
|
11
13
|
TransformerTDEquivalentRecipe,
|
|
12
14
|
)
|
|
15
|
+
from spikingjelly.logger import logger
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
_FX_TRACE_LOCK = threading.RLock()
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def _resolve_device(
|
|
22
|
+
ann: nn.Module, configured_device: Optional[Union[torch.device, str]]
|
|
23
|
+
) -> torch.device:
|
|
24
|
+
if configured_device is not None:
|
|
25
|
+
return torch.device(configured_device)
|
|
26
|
+
parameter = next(ann.parameters(), None)
|
|
27
|
+
return parameter.device if parameter is not None else torch.device("cpu")
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _symbolic_trace(root: nn.Module) -> fx.GraphModule:
|
|
31
|
+
with _FX_TRACE_LOCK:
|
|
32
|
+
original_reshape = torch.reshape
|
|
33
|
+
|
|
34
|
+
def proxy_aware_reshape(input, shape):
|
|
35
|
+
if isinstance(input, fx.Proxy):
|
|
36
|
+
return input.tracer.create_proxy(
|
|
37
|
+
"call_function", original_reshape, (input, shape), {}
|
|
38
|
+
)
|
|
39
|
+
return original_reshape(input, shape)
|
|
40
|
+
|
|
41
|
+
# Torch 2.6/2.7 validates reshape's shape before dispatching FX Proxy inputs.
|
|
42
|
+
torch.reshape = proxy_aware_reshape
|
|
43
|
+
try:
|
|
44
|
+
tracer = fx.Tracer()
|
|
45
|
+
graph = tracer.trace(root)
|
|
46
|
+
finally:
|
|
47
|
+
torch.reshape = original_reshape
|
|
48
|
+
|
|
49
|
+
# Torch 2.6/2.7 FX codegen cannot register repeated PEP 604 union types.
|
|
50
|
+
for node in graph.nodes:
|
|
51
|
+
if isinstance(node.type, types.UnionType):
|
|
52
|
+
node.type = None
|
|
53
|
+
return fx.GraphModule(tracer.root, graph)
|
|
13
54
|
|
|
14
55
|
|
|
15
56
|
class FXConverter:
|
|
@@ -79,14 +120,6 @@ class FXConverter:
|
|
|
79
120
|
"FXConverter/Converter requires an FXConversionRecipe. "
|
|
80
121
|
"Use ModuleConverter for ModuleConversionRecipe instances."
|
|
81
122
|
)
|
|
82
|
-
if recipe == "transformer_spike_equivalent":
|
|
83
|
-
warnings.warn(
|
|
84
|
-
"The 'transformer_spike_equivalent' recipe string is deprecated; "
|
|
85
|
-
"use 'transformer_td_equivalent' instead.",
|
|
86
|
-
DeprecationWarning,
|
|
87
|
-
stacklevel=3,
|
|
88
|
-
)
|
|
89
|
-
return TransformerTDEquivalentRecipe()
|
|
90
123
|
if recipe == "transformer_td_equivalent":
|
|
91
124
|
return TransformerTDEquivalentRecipe()
|
|
92
125
|
if recipe == "rate_coding":
|
|
@@ -107,14 +140,6 @@ class FXConverter:
|
|
|
107
140
|
f"instance, but got {type(recipe).__name__}."
|
|
108
141
|
)
|
|
109
142
|
|
|
110
|
-
def _resolve_device(self, ann: nn.Module) -> torch.device:
|
|
111
|
-
if self.device is not None:
|
|
112
|
-
return torch.device(self.device)
|
|
113
|
-
try:
|
|
114
|
-
return next(ann.parameters()).device
|
|
115
|
-
except StopIteration:
|
|
116
|
-
return torch.device("cpu")
|
|
117
|
-
|
|
118
143
|
def convert(self, ann: nn.Module) -> nn.Module:
|
|
119
144
|
r"""
|
|
120
145
|
**API Language** - :ref:`中文 <Converter.convert-cn>` | :ref:`English <Converter.convert-en>`
|
|
@@ -154,20 +179,31 @@ class FXConverter:
|
|
|
154
179
|
"""
|
|
155
180
|
configured_device = self.device
|
|
156
181
|
original_training_modes: dict[nn.Module, bool] = {}
|
|
182
|
+
start_time = time.perf_counter()
|
|
183
|
+
target_device = None
|
|
157
184
|
try:
|
|
158
185
|
original_training_modes = {
|
|
159
186
|
module: module.training for module in ann.modules()
|
|
160
187
|
}
|
|
161
|
-
self.device =
|
|
188
|
+
self.device = _resolve_device(ann, self.device)
|
|
189
|
+
target_device = self.device
|
|
162
190
|
with torch.no_grad():
|
|
163
191
|
self.recipe.validate(self)
|
|
164
192
|
ann = self.recipe.before_trace(self, ann)
|
|
165
|
-
fx_model =
|
|
193
|
+
fx_model = _symbolic_trace(ann).to(self.device)
|
|
166
194
|
fx_model = self.recipe.after_trace(self, fx_model)
|
|
167
195
|
fx_model = self.recipe.insert_observers(self, fx_model)
|
|
168
196
|
fx_model = self.recipe.calibrate(self, fx_model)
|
|
169
197
|
fx_model = self.recipe.replace(self, fx_model)
|
|
170
198
|
fx_model = self.recipe.finalize(self, fx_model)
|
|
199
|
+
logger.info(
|
|
200
|
+
"Conversion completed: converter={} recipe={} device={} modules={} elapsed_ms={:.3f}",
|
|
201
|
+
type(self).__name__,
|
|
202
|
+
type(self.recipe).__name__,
|
|
203
|
+
target_device,
|
|
204
|
+
sum(1 for _ in fx_model.modules()),
|
|
205
|
+
(time.perf_counter() - start_time) * 1000.0,
|
|
206
|
+
)
|
|
171
207
|
return fx_model
|
|
172
208
|
finally:
|
|
173
209
|
for module, training in original_training_modes.items():
|
|
@@ -239,14 +275,6 @@ class ModuleConverter:
|
|
|
239
275
|
self.recipe = recipe
|
|
240
276
|
self.device = device
|
|
241
277
|
|
|
242
|
-
def _resolve_device(self, ann: nn.Module) -> torch.device:
|
|
243
|
-
if self.device is not None:
|
|
244
|
-
return torch.device(self.device)
|
|
245
|
-
try:
|
|
246
|
-
return next(ann.parameters()).device
|
|
247
|
-
except StopIteration:
|
|
248
|
-
return torch.device("cpu")
|
|
249
|
-
|
|
250
278
|
def convert(self, ann: nn.Module) -> nn.Module:
|
|
251
279
|
r"""
|
|
252
280
|
**API Language** - :ref:`中文 <ModuleConverter.convert-cn>` | :ref:`English <ModuleConverter.convert-en>`
|
|
@@ -279,11 +307,14 @@ class ModuleConverter:
|
|
|
279
307
|
"""
|
|
280
308
|
configured_device = self.device
|
|
281
309
|
original_training_modes: dict[nn.Module, bool] = {}
|
|
310
|
+
start_time = time.perf_counter()
|
|
311
|
+
target_device = None
|
|
282
312
|
try:
|
|
283
313
|
original_training_modes = {
|
|
284
314
|
module: module.training for module in ann.modules()
|
|
285
315
|
}
|
|
286
|
-
self.device =
|
|
316
|
+
self.device = _resolve_device(ann, self.device)
|
|
317
|
+
target_device = self.device
|
|
287
318
|
with torch.no_grad():
|
|
288
319
|
self.recipe.validate(self)
|
|
289
320
|
converted = self.recipe.convert_module(self, ann)
|
|
@@ -293,7 +324,16 @@ class ModuleConverter:
|
|
|
293
324
|
"a torch.nn.Module, got "
|
|
294
325
|
f"{type(converted).__name__}."
|
|
295
326
|
)
|
|
296
|
-
|
|
327
|
+
converted = converted.to(self.device)
|
|
328
|
+
logger.info(
|
|
329
|
+
"Conversion completed: converter={} recipe={} device={} modules={} elapsed_ms={:.3f}",
|
|
330
|
+
type(self).__name__,
|
|
331
|
+
type(self.recipe).__name__,
|
|
332
|
+
target_device,
|
|
333
|
+
sum(1 for _ in converted.modules()),
|
|
334
|
+
(time.perf_counter() - start_time) * 1000.0,
|
|
335
|
+
)
|
|
336
|
+
return converted
|
|
297
337
|
finally:
|
|
298
338
|
for module, training in original_training_modes.items():
|
|
299
339
|
module.training = training
|
{spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/delay.py
RENAMED
|
@@ -1,8 +1,9 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
3
|
import math
|
|
4
|
+
import time
|
|
4
5
|
import warnings
|
|
5
|
-
from typing import Dict, Iterable, List,
|
|
6
|
+
from typing import Dict, Iterable, List, Tuple, Union
|
|
6
7
|
|
|
7
8
|
import torch
|
|
8
9
|
import torch.nn as nn
|
|
@@ -12,19 +13,15 @@ from spikingjelly.activation_based.ann2snn.modules import (
|
|
|
12
13
|
ChannelVoltageScaler,
|
|
13
14
|
VoltageScaler,
|
|
14
15
|
)
|
|
16
|
+
from spikingjelly.activation_based.functional.net_config import reset_net
|
|
15
17
|
from spikingjelly.activation_based.neuron.base_node import BaseNode
|
|
18
|
+
from spikingjelly.logger import logger
|
|
16
19
|
|
|
17
20
|
|
|
18
21
|
Scaler = Union[VoltageScaler, ChannelVoltageScaler]
|
|
19
22
|
_MIN_READOUT_STEPS = 4
|
|
20
23
|
|
|
21
24
|
|
|
22
|
-
def _reset_snn(model: nn.Module) -> None:
|
|
23
|
-
for module in model.modules():
|
|
24
|
-
if hasattr(module, "reset"):
|
|
25
|
-
module.reset()
|
|
26
|
-
|
|
27
|
-
|
|
28
25
|
def _extract_batch_input(batch):
|
|
29
26
|
if isinstance(batch, torch.Tensor):
|
|
30
27
|
return batch
|
|
@@ -42,18 +39,6 @@ def _as_runtime_tensor(value, x: torch.Tensor) -> torch.Tensor:
|
|
|
42
39
|
return torch.as_tensor(value, device=x.device, dtype=x.dtype)
|
|
43
40
|
|
|
44
41
|
|
|
45
|
-
def _scale_view(scaler: Scaler, x: torch.Tensor) -> torch.Tensor:
|
|
46
|
-
if isinstance(scaler, ChannelVoltageScaler):
|
|
47
|
-
return scaler._view_scale(x).to(device=x.device, dtype=x.dtype)
|
|
48
|
-
return scaler.scale.to(device=x.device, dtype=x.dtype)
|
|
49
|
-
|
|
50
|
-
|
|
51
|
-
def _channel_dim(scaler: Scaler) -> Optional[int]:
|
|
52
|
-
if isinstance(scaler, ChannelVoltageScaler):
|
|
53
|
-
return scaler.channel_dim
|
|
54
|
-
return None
|
|
55
|
-
|
|
56
|
-
|
|
57
42
|
def _compute_delay_ratio(
|
|
58
43
|
module: BaseNode,
|
|
59
44
|
post_scaler: Scaler,
|
|
@@ -65,21 +50,22 @@ def _compute_delay_ratio(
|
|
|
65
50
|
if not torch.isfinite(v_threshold).all() or (v_threshold <= 0).any():
|
|
66
51
|
raise ValueError("Delay estimation requires finite positive v_threshold.")
|
|
67
52
|
|
|
68
|
-
|
|
69
|
-
v_init = module.get_reset_value("v")
|
|
70
|
-
except (KeyError, AttributeError):
|
|
71
|
-
v_init = getattr(module, "v_reset", 0.0)
|
|
53
|
+
v_init = module.get_reset_value("v")
|
|
72
54
|
if v_init is None:
|
|
73
55
|
v_init = 0.0
|
|
74
56
|
v_init = _as_runtime_tensor(v_init, x)
|
|
75
57
|
|
|
76
|
-
|
|
58
|
+
if isinstance(post_scaler, ChannelVoltageScaler):
|
|
59
|
+
scale = post_scaler._view_scale(x).to(device=x.device, dtype=x.dtype)
|
|
60
|
+
channel_dim = post_scaler.channel_dim
|
|
61
|
+
else:
|
|
62
|
+
scale = post_scaler.scale.to(device=x.device, dtype=x.dtype)
|
|
63
|
+
channel_dim = None
|
|
77
64
|
x_nonnegative = torch.clamp(x.detach(), min=0)
|
|
78
65
|
original_activation = x_nonnegative * scale / v_threshold
|
|
79
66
|
required_charge = torch.clamp(v_threshold - v_init, min=0)
|
|
80
67
|
original_required_charge = required_charge * scale / v_threshold
|
|
81
68
|
|
|
82
|
-
channel_dim = _channel_dim(post_scaler)
|
|
83
69
|
if channel_dim is None or scale.dim() == 0:
|
|
84
70
|
max_mean_activation = original_activation.mean()
|
|
85
71
|
reduced_required_charge = original_required_charge.mean()
|
|
@@ -210,8 +196,13 @@ def estimate_delay_start(
|
|
|
210
196
|
if num_batches <= 0:
|
|
211
197
|
raise ValueError("num_batches must be positive.")
|
|
212
198
|
|
|
199
|
+
start_time = time.perf_counter()
|
|
213
200
|
paths = _find_scaler_neuron_scaler_paths(model)
|
|
214
201
|
if not paths:
|
|
202
|
+
logger.info(
|
|
203
|
+
"Delay estimated: matched_paths=0 batches=0 raw_delay_start=0 delay_start=0 readout_clamped=False elapsed_ms={:.3f}",
|
|
204
|
+
(time.perf_counter() - start_time) * 1000.0,
|
|
205
|
+
)
|
|
215
206
|
return 0
|
|
216
207
|
|
|
217
208
|
original_device = None
|
|
@@ -225,6 +216,7 @@ def estimate_delay_start(
|
|
|
225
216
|
original_training_modes = {module: module.training for module in model.modules()}
|
|
226
217
|
ratios: Dict[BaseNode, List[float]] = {module: [] for module, _ in paths}
|
|
227
218
|
handles = []
|
|
219
|
+
processed_batches = 0
|
|
228
220
|
|
|
229
221
|
def make_hook(module: BaseNode, post_scaler: Scaler):
|
|
230
222
|
def hook(_module, inputs, _output):
|
|
@@ -245,7 +237,7 @@ def estimate_delay_start(
|
|
|
245
237
|
for module, post_scaler in paths:
|
|
246
238
|
handles.append(module.register_forward_hook(make_hook(module, post_scaler)))
|
|
247
239
|
|
|
248
|
-
|
|
240
|
+
reset_net(model)
|
|
249
241
|
with torch.no_grad():
|
|
250
242
|
for batch_idx, batch in enumerate(dataloader):
|
|
251
243
|
if batch_idx >= num_batches:
|
|
@@ -253,8 +245,9 @@ def estimate_delay_start(
|
|
|
253
245
|
x = _extract_batch_input(batch)
|
|
254
246
|
if not isinstance(x, torch.Tensor):
|
|
255
247
|
raise TypeError("The extracted model input must be a tensor.")
|
|
256
|
-
|
|
248
|
+
reset_net(model)
|
|
257
249
|
model(x.to(device, non_blocking=True))
|
|
250
|
+
processed_batches += 1
|
|
258
251
|
delay = 0.0
|
|
259
252
|
for values in ratios.values():
|
|
260
253
|
if values:
|
|
@@ -269,14 +262,14 @@ def estimate_delay_start(
|
|
|
269
262
|
finally:
|
|
270
263
|
for handle in handles:
|
|
271
264
|
handle.remove()
|
|
272
|
-
|
|
265
|
+
reset_net(model)
|
|
273
266
|
if original_device is not None:
|
|
274
267
|
model.to(original_device)
|
|
275
268
|
for module, training in original_training_modes.items():
|
|
276
269
|
module.training = training
|
|
277
270
|
|
|
278
|
-
|
|
279
|
-
if time_steps <
|
|
271
|
+
raw_delay_start = math.ceil(delay) if math.isfinite(delay) else time_steps
|
|
272
|
+
if time_steps < raw_delay_start + _MIN_READOUT_STEPS:
|
|
280
273
|
warnings.warn(
|
|
281
274
|
"estimate_delay_start: time_steps is too small to keep "
|
|
282
275
|
f"{_MIN_READOUT_STEPS} readout steps after the estimated delay; "
|
|
@@ -284,5 +277,25 @@ def estimate_delay_start(
|
|
|
284
277
|
RuntimeWarning,
|
|
285
278
|
stacklevel=2,
|
|
286
279
|
)
|
|
287
|
-
|
|
288
|
-
|
|
280
|
+
result = max(time_steps - _MIN_READOUT_STEPS - 1, 0)
|
|
281
|
+
logger.info(
|
|
282
|
+
"Delay estimated: matched_paths={} batches={} raw_delay_start={} delay_start={} readout_clamped={} elapsed_ms={:.3f}",
|
|
283
|
+
len(paths),
|
|
284
|
+
processed_batches,
|
|
285
|
+
raw_delay_start,
|
|
286
|
+
result,
|
|
287
|
+
True,
|
|
288
|
+
(time.perf_counter() - start_time) * 1000.0,
|
|
289
|
+
)
|
|
290
|
+
return result
|
|
291
|
+
result = min(raw_delay_start, time_steps - 1)
|
|
292
|
+
logger.info(
|
|
293
|
+
"Delay estimated: matched_paths={} batches={} raw_delay_start={} delay_start={} readout_clamped={} elapsed_ms={:.3f}",
|
|
294
|
+
len(paths),
|
|
295
|
+
processed_batches,
|
|
296
|
+
raw_delay_start,
|
|
297
|
+
result,
|
|
298
|
+
False,
|
|
299
|
+
(time.perf_counter() - start_time) * 1000.0,
|
|
300
|
+
)
|
|
301
|
+
return result
|
|
@@ -59,7 +59,8 @@ class FXFriendlyBertSelfAttention(nn.Module):
|
|
|
59
59
|
|
|
60
60
|
def transpose_for_scores(self, x: torch.Tensor) -> torch.Tensor:
|
|
61
61
|
x = x.view(
|
|
62
|
-
|
|
62
|
+
x.shape[0],
|
|
63
|
+
x.shape[1],
|
|
63
64
|
self.num_attention_heads,
|
|
64
65
|
self.attention_head_size,
|
|
65
66
|
)
|
|
@@ -80,7 +81,8 @@ class FXFriendlyBertSelfAttention(nn.Module):
|
|
|
80
81
|
attention_probs = self.dropout(attention_probs)
|
|
81
82
|
context_layer = torch.matmul(attention_probs, value_layer)
|
|
82
83
|
context_layer = context_layer.permute(0, 2, 1, 3).reshape(
|
|
83
|
-
|
|
84
|
+
hidden_states.shape[0],
|
|
85
|
+
hidden_states.shape[1],
|
|
84
86
|
self.all_head_size,
|
|
85
87
|
)
|
|
86
88
|
return context_layer
|
|
@@ -9,7 +9,7 @@ import torch
|
|
|
9
9
|
import torchvision
|
|
10
10
|
from tqdm import tqdm
|
|
11
11
|
|
|
12
|
-
from spikingjelly.activation_based import ann2snn
|
|
12
|
+
from spikingjelly.activation_based import ann2snn, functional
|
|
13
13
|
from spikingjelly.activation_based.ann2snn.sample_models import mnist_cnn
|
|
14
14
|
|
|
15
15
|
|
|
@@ -47,7 +47,6 @@ def val(net, device, data_loader, T=None):
|
|
|
47
47
|
total = 0.0
|
|
48
48
|
if T is not None:
|
|
49
49
|
corrects = np.zeros(T)
|
|
50
|
-
reset_modules = [m for m in net.modules() if hasattr(m, "reset")]
|
|
51
50
|
with torch.no_grad():
|
|
52
51
|
for batch, (img, label) in enumerate(tqdm(data_loader)):
|
|
53
52
|
img = img.to(device, non_blocking=True)
|
|
@@ -56,8 +55,7 @@ def val(net, device, data_loader, T=None):
|
|
|
56
55
|
out = net(img)
|
|
57
56
|
correct += (out.argmax(dim=1) == label).float().sum().item()
|
|
58
57
|
else:
|
|
59
|
-
|
|
60
|
-
m.reset()
|
|
58
|
+
functional.reset_net(net)
|
|
61
59
|
out = None
|
|
62
60
|
for t in range(T):
|
|
63
61
|
step = net(img)
|
|
@@ -157,8 +155,6 @@ def run_recipe_comparison(
|
|
|
157
155
|
ltb_accs = convert_and_eval(
|
|
158
156
|
ann2snn.LocalThresholdBalancingRecipe(
|
|
159
157
|
dataloader=calibration_data_loader,
|
|
160
|
-
time_steps=time_steps,
|
|
161
|
-
mode="99.9%",
|
|
162
158
|
),
|
|
163
159
|
device,
|
|
164
160
|
model,
|
|
@@ -293,9 +289,6 @@ def main(args):
|
|
|
293
289
|
transform=torchvision.transforms.ToTensor(),
|
|
294
290
|
download=True,
|
|
295
291
|
)
|
|
296
|
-
train_data_loader = torch.utils.data.DataLoader(
|
|
297
|
-
dataset=train_data_dataset, batch_size=batch_size, shuffle=True, drop_last=False
|
|
298
|
-
)
|
|
299
292
|
calibration_data_loader = torch.utils.data.DataLoader(
|
|
300
293
|
dataset=train_data_dataset,
|
|
301
294
|
batch_size=batch_size,
|
|
@@ -315,21 +308,24 @@ def main(args):
|
|
|
315
308
|
drop_last=False,
|
|
316
309
|
)
|
|
317
310
|
|
|
318
|
-
#
|
|
319
|
-
#
|
|
320
|
-
#
|
|
311
|
+
# To train the ANN checkpoint locally:
|
|
312
|
+
# model = mnist_cnn.CNN().to(device)
|
|
313
|
+
# loss_function = torch.nn.CrossEntropyLoss()
|
|
314
|
+
# optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=5e-4)
|
|
315
|
+
# train_data_loader = torch.utils.data.DataLoader(
|
|
316
|
+
# train_data_dataset, batch_size=batch_size, shuffle=True
|
|
317
|
+
# )
|
|
318
|
+
# for epoch in range(10):
|
|
321
319
|
# model.train()
|
|
322
|
-
# for
|
|
320
|
+
# for img, label in train_data_loader:
|
|
323
321
|
# optimizer.zero_grad()
|
|
324
322
|
# out = model(img.to(device))
|
|
325
323
|
# loss = loss_function(out, label.to(device))
|
|
326
324
|
# loss.backward()
|
|
327
325
|
# optimizer.step()
|
|
328
|
-
# torch.save(model.state_dict(),
|
|
329
|
-
# print(
|
|
330
|
-
#
|
|
331
|
-
# print('Validating Accuracy: %.3f' % (acc))
|
|
332
|
-
# print()
|
|
326
|
+
# torch.save(model.state_dict(), "SJ-mnist-cnn_model-sample.pth")
|
|
327
|
+
# print("Epoch: %d" % epoch)
|
|
328
|
+
# print("Validating Accuracy: %.3f" % val(model, device, train_data_loader))
|
|
333
329
|
|
|
334
330
|
if args.plot_mode_sweep:
|
|
335
331
|
run_legacy_mode_sweep(
|