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.
Files changed (366) hide show
  1. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/PKG-INFO +10 -3
  2. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/README.md +2 -0
  3. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/pyproject.toml +33 -3
  4. spikingjelly-2.0.0.dev1/spikingjelly/__init__.py +7 -0
  5. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/__init__.py +20 -17
  6. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/converter.py +69 -29
  7. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/delay.py +45 -32
  8. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/bert_sst2_transformer_td_equivalent.py +4 -2
  9. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/cnn_mnist.py +14 -18
  10. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/imagenet_resnet18_ltb.py +2 -18
  11. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/resnet18_cifar10.py +2 -5
  12. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/operators.py +560 -153
  13. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/ann2snn/qcfs.py +458 -0
  14. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/__init__.py +12 -0
  15. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/base.py +39 -0
  16. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/ann2snn/recipes/local_threshold_balancing.py +325 -0
  17. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/ann2snn/recipes/qwen2.py +1360 -0
  18. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/rate_coding.py +332 -397
  19. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/spikezip_qann.py +313 -326
  20. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/sta_transformer.py +83 -540
  21. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/ann2snn/recipes/step_mode_adapters.py +636 -0
  22. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/recipes/transformer_td_equivalent.py +37 -65
  23. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/sample_models/cifar10_resnet.py +6 -6
  24. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/utils.py +34 -1
  25. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/base.py +421 -81
  26. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/__init__.py +38 -0
  27. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/auto_cuda/__init__.py +1 -0
  28. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/auto_cuda/base.py +31 -54
  29. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/auto_cuda/cfunction.py +14 -7
  30. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/auto_cuda/generator.py +19 -16
  31. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/cuda_utils.py +60 -151
  32. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/__init__.py +1 -0
  33. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/cuda_code.py +27 -0
  34. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step/__init__.py +21 -0
  35. 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
  36. {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
  37. {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
  38. {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
  39. {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
  40. {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
  41. {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
  42. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/multi_step/runtime.py +84 -0
  43. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/single_step/__init__.py +1 -0
  44. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/single_step/base.py +254 -0
  45. {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
  46. {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
  47. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_kernel/surrogate_registry.py +44 -0
  48. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/neuron_linear.py +708 -0
  49. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/cuda_kernel/spike_linear.py +933 -0
  50. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/spike_op.py +23 -11
  51. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/cuda_kernel/tensor_cache.py +23 -24
  52. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/distributed/__init__.py +31 -18
  53. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/distributed/adapters/__init__.py +10 -0
  54. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/distributed/adapters/base.py +16 -19
  55. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/adapters/cifar10dvs_vgg.py +48 -0
  56. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/adapters/spikformer.py +48 -0
  57. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/analysis.py +160 -0
  58. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/api.py +310 -0
  59. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/config.py +114 -0
  60. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/data_parallel.py +54 -0
  61. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/execution.py +227 -0
  62. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/fsdp.py +105 -0
  63. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/mesh.py +259 -0
  64. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/distributed/metrics.py +1 -1
  65. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/optimizer.py +95 -0
  66. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/__init__.py +21 -0
  67. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/cifar10dvs_vgg.py +148 -0
  68. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/memopt.py +183 -0
  69. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/partition.py +166 -0
  70. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/runtime.py +486 -0
  71. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/pipeline/spikformer.py +190 -0
  72. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/planner.py +490 -0
  73. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/runtime.py +160 -0
  74. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/__init__.py +41 -0
  75. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/channel.py +435 -0
  76. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/cifar10dvs_vgg.py +146 -0
  77. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/debug.py +63 -0
  78. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/linear.py +392 -0
  79. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/spikformer.py +223 -0
  80. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/state.py +166 -0
  81. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/distributed/tensor_parallel/utils.py +50 -0
  82. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/distributed/topology.py +22 -40
  83. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/encoding.py +51 -93
  84. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/A2C.py +0 -1
  85. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DQN_state.py +1 -1
  86. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/agent.py +10 -13
  87. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/experience.py +0 -1
  88. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/train.py +0 -1
  89. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/NoisySAN/hybrid_td3_cuda_norm.py +0 -1
  90. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/NoisySAN/test_hybrid_td3_cpu.py +0 -1
  91. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/Spiking_A2C.py +5 -31
  92. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/Spiking_DQN_state.py +9 -29
  93. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/Spiking_PPO.py +5 -30
  94. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/cifar10_r11_enabling_spikebased_backpropagation.py +2 -4
  95. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/common/multiprocessing_env.py +0 -4
  96. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/lava_mnist.py +1 -1
  97. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/train.py +0 -1
  98. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/train_distributed.py +22 -23
  99. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/mstdp.py +0 -1
  100. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/mstdpet.py +0 -1
  101. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/speechcommands.py +0 -8
  102. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/spiking_lstm_sequential_mnist.py +1 -13
  103. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/spiking_lstm_text.py +1 -31
  104. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/stdp_trace.py +0 -1
  105. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/__init__.py +3 -0
  106. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/conv_bn_fusion.py +24 -30
  107. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/forward.py +94 -72
  108. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/functional/layer.py +200 -0
  109. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/functional/learning.py +602 -0
  110. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/misc.py +11 -11
  111. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/net_config.py +37 -56
  112. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/functional/neuron.py +2994 -0
  113. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/online_learning.py +46 -73
  114. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/lava_exchange.py +114 -134
  115. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/attention.py +58 -86
  116. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/bn.py +43 -15
  117. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/container.py +67 -49
  118. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/dropout.py +75 -68
  119. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/layer/misc.py +629 -0
  120. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/online_learning.py +25 -30
  121. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/stateless_wrapper.py +331 -208
  122. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/learning.py +182 -232
  123. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/lynxi_exchange.py +12 -18
  124. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/memopt/checkpointing.py +46 -15
  125. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/memopt/compress.py +4 -11
  126. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/memopt/pipeline.py +82 -84
  127. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/parametric_lif_net.py +16 -17
  128. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/snas_net.py +7 -9
  129. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/spike_dhs.py +78 -48
  130. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/spiking_vggws_ottt.py +13 -34
  131. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/train_classify.py +121 -96
  132. 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
  133. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/train_imagenet_example.py +1 -1
  134. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/tv_ref_classify/utils.py +55 -60
  135. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/monitor.py +76 -84
  136. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/__init__.py +1 -0
  137. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/adapt.py +87 -140
  138. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/base_node.py +249 -204
  139. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/dsr.py +35 -65
  140. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/few_spike.py +25 -42
  141. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/flexsn.py +351 -300
  142. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/neuron/ilif.py +410 -0
  143. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/integrate_and_fire.py +525 -376
  144. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/neuron/lif.py +397 -0
  145. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/lif_variants.py +174 -201
  146. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/mpbn.py +152 -225
  147. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/noisy.py +76 -175
  148. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/nonlinear_if.py +113 -58
  149. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/online_learning.py +61 -202
  150. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/plif.py +73 -147
  151. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/psn.py +55 -43
  152. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/neuron/spikezip.py +223 -0
  153. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/nir_exchange/from_nir.py +9 -0
  154. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/nir_exchange/to_nir.py +9 -0
  155. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/ac.py +17 -28
  156. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/analytical_energy/core.py +0 -1
  157. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/base.py +104 -52
  158. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/compute_energy.py +23 -19
  159. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/flop.py +6 -12
  160. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/lemaire_addressing.py +8 -30
  161. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/mac.py +3 -9
  162. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/memory_access.py +10 -25
  163. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/memory_residency.py +16 -142
  164. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/add_counter.py +8 -26
  165. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/base_counter.py +17 -11
  166. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/cmp_counter.py +2 -2
  167. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/config.py +11 -43
  168. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/core.py +149 -221
  169. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/memory_residency_counter.py +0 -15
  170. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/mul_counter.py +1 -15
  171. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/op_counter/neuromc/utils.py +59 -0
  172. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuron_state.py +26 -6
  173. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/spikesim/__init__.py +0 -4
  174. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/spikesim/core.py +9 -34
  175. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/spikesim/counter.py +4 -28
  176. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/synop.py +12 -140
  177. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/__init__.py +20 -0
  178. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/api.py +50 -7
  179. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/capability.py +167 -1
  180. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/precision/config.py +46 -0
  181. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/convert.py +22 -2
  182. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/precision/float8_attention.py +175 -0
  183. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/float8_base.py +24 -19
  184. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/precision/float8_conv.py +128 -0
  185. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/precision/float8_te.py +595 -0
  186. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/float8_torchao.py +19 -43
  187. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/policy.py +7 -14
  188. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/precision/runtime.py +7 -8
  189. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/profiler.py +63 -55
  190. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/quantize.py +43 -7
  191. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/rnn.py +119 -191
  192. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/surrogate.py +135 -296
  193. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/compress.py +3 -3
  194. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/__init__.py +4 -4
  195. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/custom_ops.py +87 -88
  196. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/hop.py +2 -2
  197. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/kernel.py +4 -4
  198. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/template.py +36 -63
  199. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/wrapper.py +5 -6
  200. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/fp8_capability.py +412 -0
  201. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/__init__.py +15 -0
  202. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/activation_aware_if.py +281 -0
  203. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/ilif.py +724 -0
  204. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/integrate_and_fire.py +1225 -0
  205. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/lif.py +1333 -0
  206. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/plif.py +1337 -0
  207. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/stbif.py +470 -0
  208. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/neuron_kernel/utils.py +209 -0
  209. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/surrogate_kernel.py +3 -4
  210. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/torch2triton/__init__.py +9 -0
  211. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/torch2triton/graph2triton.py +14 -30
  212. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/torch2triton/torch2graph.py +2 -0
  213. spikingjelly-2.0.0.dev1/spikingjelly/activation_based/triton_kernel/triton_utils.py +458 -0
  214. spikingjelly-2.0.0.dev1/spikingjelly/configure.py +177 -0
  215. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/asl_dvs.py +16 -14
  216. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/base.py +114 -151
  217. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/bullying10k.py +22 -18
  218. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/cifar10_dvs.py +21 -22
  219. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/dvs128_gesture.py +28 -23
  220. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/dvs_lip.py +2 -1
  221. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/es_imagenet.py +26 -183
  222. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/hardvs.py +12 -9
  223. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/n_caltech101.py +11 -12
  224. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/n_mnist.py +12 -13
  225. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/nav_gesture.py +42 -68
  226. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/shd.py +63 -169
  227. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/speechcommands.py +7 -6
  228. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/utils.py +58 -415
  229. spikingjelly-2.0.0.dev1/spikingjelly/logger.py +38 -0
  230. spikingjelly-2.0.0.dev1/spikingjelly/timing_based/encoding.py +134 -0
  231. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/timing_based/examples/tempotron_mnist.py +0 -3
  232. spikingjelly-2.0.0.dev1/spikingjelly/timing_based/neuron.py +192 -0
  233. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly.egg-info/PKG-INFO +10 -3
  234. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly.egg-info/SOURCES.txt +58 -28
  235. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly.egg-info/requires.txt +10 -2
  236. spikingjelly-2.0.0.dev1/test/test_configure.py +127 -0
  237. spikingjelly-2.0.0.dev1/test/test_dataset_builders.py +144 -0
  238. spikingjelly-2.0.0.dev1/test/test_dataset_utils.py +197 -0
  239. spikingjelly-2.0.0.dev1/test/test_logging_policy.py +264 -0
  240. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/factories.py +0 -243
  241. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/recipes/local_threshold_balancing.py +0 -461
  242. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/recipes/step_mode_adapters.py +0 -1255
  243. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/rules.py +0 -380
  244. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/ann2snn/threshold.py +0 -90
  245. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/__init__.py +0 -42
  246. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/example.py +0 -35
  247. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/neuron_kernel/__init__.py +0 -4
  248. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/ss_neuron_kernel/__init__.py +0 -35
  249. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/auto_cuda/ss_neuron_kernel/ss_neuron_kernel_base.py +0 -522
  250. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/__init__.py +0 -37
  251. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/common.py +0 -470
  252. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/helpers.py +0 -51
  253. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/integrate_and_fire.py +0 -794
  254. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/lif.py +0 -864
  255. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/cuda_kernel/neuron_kernel/plif.py +0 -823
  256. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/adapters/cifar10dvs_vgg.py +0 -57
  257. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/adapters/spikformer.py +0 -61
  258. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/api.py +0 -210
  259. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/dtensor.py +0 -3434
  260. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/planner.py +0 -38
  261. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/distributed/runtime.py +0 -264
  262. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/layer/misc.py +0 -378
  263. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/neuron/lif.py +0 -708
  264. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/neuron/spikezip.py +0 -256
  265. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/neuron_cupy.py +0 -3180
  266. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/neuron_cupy_lite.py +0 -2867
  267. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/op_counter/neuromc/utils.py +0 -179
  268. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/precision/config.py +0 -83
  269. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/neuron_kernel/__init__.py +0 -11
  270. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/neuron_kernel/integrate_and_fire.py +0 -679
  271. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/neuron_kernel/lif.py +0 -736
  272. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/neuron_kernel/plif.py +0 -784
  273. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/spikezip_kernel.py +0 -226
  274. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/torch2triton/__init__.py +0 -9
  275. spikingjelly-2.0.0.dev0/spikingjelly/activation_based/triton_kernel/triton_utils.py +0 -188
  276. spikingjelly-2.0.0.dev0/spikingjelly/configure.py +0 -85
  277. spikingjelly-2.0.0.dev0/spikingjelly/timing_based/__init__.py +0 -0
  278. spikingjelly-2.0.0.dev0/spikingjelly/timing_based/encoding.py +0 -386
  279. spikingjelly-2.0.0.dev0/spikingjelly/timing_based/examples/__init__.py +0 -0
  280. spikingjelly-2.0.0.dev0/spikingjelly/timing_based/neuron.py +0 -362
  281. spikingjelly-2.0.0.dev0/spikingjelly/timing_based/orig_encoding.py +0 -57
  282. spikingjelly-2.0.0.dev0/spikingjelly/timing_based/orig_neuron.py +0 -127
  283. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/LICENSE +0 -0
  284. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/LICENSES/translations/LICENSE.de +0 -0
  285. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/LICENSES/translations/LICENSE.en +0 -0
  286. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/LICENSES/translations/LICENSE.fr +0 -0
  287. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/LICENSES/translations/LICENSE.hi +0 -0
  288. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/setup.cfg +0 -0
  289. {spikingjelly-2.0.0.dev0/spikingjelly → spikingjelly-2.0.0.dev1/spikingjelly/activation_based}/__init__.py +0 -0
  290. {spikingjelly-2.0.0.dev0/spikingjelly/activation_based → spikingjelly-2.0.0.dev1/spikingjelly/activation_based/ann2snn/examples}/__init__.py +0 -0
  291. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/imagenet_vit_sta.py +0 -0
  292. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/examples/roberta_spikezip_qann_synthetic.py +0 -0
  293. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/modules.py +0 -0
  294. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/sample_models/__init__.py +0 -0
  295. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/ann2snn/sample_models/mnist_cnn.py +0 -0
  296. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/__init__.py +0 -0
  297. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/actions.py +0 -0
  298. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/common/__init__.py +0 -0
  299. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/common/runfile.py +0 -0
  300. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/common/utils.py +0 -0
  301. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/common/wrappers.py +0 -0
  302. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/common/wrappers_simple.py +0 -0
  303. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/ptan/ignite.py +0 -0
  304. {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
  305. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/utils/atari_wrappers.py +0 -0
  306. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/utils/common.py +0 -0
  307. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/DSQN/utils/model.py +0 -0
  308. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/ILC-SAN/core_cuda.py +0 -0
  309. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/ILC-SAN/hybrid_td3_cuda_norm.py +0 -0
  310. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/ILC-SAN/ilcsan.py +0 -0
  311. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/ILC-SAN/replay_buffer_norm.py +0 -0
  312. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/ILC-SAN/test_hybrid_td3_cpu.py +0 -0
  313. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/NoisySAN/core_cuda.py +0 -0
  314. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/NoisySAN/noisysan.py +0 -0
  315. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/NoisySAN/replay_buffer_norm.py +0 -0
  316. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/PPO.py +0 -0
  317. {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
  318. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/classify_dvsg.py +0 -0
  319. {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
  320. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/conv_fashion_mnist.py +0 -0
  321. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/lif_fc_mnist.py +0 -0
  322. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/lynxi_fmnist_inference.py +0 -0
  323. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/data_module.py +0 -0
  324. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/lightning_callbacks.py +0 -0
  325. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/lightning_modules.py +0 -0
  326. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/loss.py +0 -0
  327. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/memopt/models.py +0 -0
  328. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/examples/rsnn_sequential_fmnist.py +0 -0
  329. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/functional/loss.py +0 -0
  330. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/layer/__init__.py +0 -0
  331. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/memopt/__init__.py +0 -0
  332. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/__init__.py +0 -0
  333. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/sew_resnet.py +0 -0
  334. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/spikformer.py +0 -0
  335. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/spiking_resnet.py +0 -0
  336. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/spiking_vgg.py +0 -0
  337. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/tv_ref_classify/__init__.py +0 -0
  338. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/tv_ref_classify/presets.py +0 -0
  339. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/tv_ref_classify/sampler.py +0 -0
  340. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/model/tv_ref_classify/transforms.py +0 -0
  341. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/neuron/inter_layer_connection.py +0 -0
  342. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/nir_exchange/__init__.py +0 -0
  343. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/__init__.py +0 -0
  344. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/_sparse_memory.py +0 -0
  345. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/analytical_energy/__init__.py +0 -0
  346. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/__init__.py +0 -0
  347. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/mux_counter.py +0 -0
  348. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/neuromc/sqrt_counter.py +0 -0
  349. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/spikesim/config.py +0 -0
  350. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/op_counter/spikesim/formulas.py +0 -0
  351. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/__init__.py +0 -0
  352. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/dummy.py +0 -0
  353. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/activation_based/triton_kernel/flexsn/info.py +0 -0
  354. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/__init__.py +0 -0
  355. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/datasets/transform.py +0 -0
  356. {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/examples → spikingjelly-2.0.0.dev1/spikingjelly/timing_based}/__init__.py +0 -0
  357. {spikingjelly-2.0.0.dev0/spikingjelly/activation_based/examples/common → spikingjelly-2.0.0.dev1/spikingjelly/timing_based/examples}/__init__.py +0 -0
  358. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/__init__.py +0 -0
  359. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/_utils.py +0 -0
  360. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/bar3d.py +0 -0
  361. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/feature_map.py +0 -0
  362. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/heatmap.py +0 -0
  363. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly/visualizing/spikes.py +0 -0
  364. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly.egg-info/dependency_links.txt +0 -0
  365. {spikingjelly-2.0.0.dev0 → spikingjelly-2.0.0.dev1}/spikingjelly.egg-info/top_level.txt +0 -0
  366. {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.dev0
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.dev0"
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
@@ -0,0 +1,7 @@
1
+ """SpikingJelly package.
2
+
3
+ Import the package logger explicitly from :mod:`spikingjelly.logger` when it is
4
+ needed; it is not re-exported from this package namespace.
5
+ """
6
+
7
+ __all__ = []
@@ -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`,以及 ``HookFactory``、``NeuronFactory``、
13
- ``ReLURule``、``ThresholdOptimizer`` 等可扩展组件,并附带 ``download_url``
14
- 工具函数。兼容名 :class:`Converter` 等价于 :class:`FXConverter`,
15
- :class:`ConversionRecipe` 等价于 :class:`FXConversionRecipe`。
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`, extensible building blocks —
26
- :class:`HookFactory`, :class:`NeuronFactory`, :class:`ReLURule` and
27
- :class:`ThresholdOptimizer` — and a ``download_url`` helper for fetching
28
- pretrained models. The compatibility name :class:`Converter` is equivalent to
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 warnings
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 = self._resolve_device(ann)
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 = fx.symbolic_trace(ann).to(self.device)
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 = self._resolve_device(ann)
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
- return converted.to(self.device)
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
@@ -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, Optional, Tuple, Union
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
- try:
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
- scale = _scale_view(post_scaler, x)
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
- _reset_snn(model)
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
- _reset_snn(model)
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
- _reset_snn(model)
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
- delay_start = math.ceil(delay) if math.isfinite(delay) else time_steps
279
- if time_steps < delay_start + _MIN_READOUT_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
- return max(time_steps - _MIN_READOUT_STEPS - 1, 0)
288
- return min(delay_start, time_steps - 1)
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
- *x.size()[:-1],
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
- *hidden_states.size()[:-1],
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
- for m in reset_modules:
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
- # loss_function = nn.CrossEntropyLoss()
319
- # optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=5e-4)
320
- # for epoch in range(epochs):
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 (img, label) in train_data_loader:
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(), 'SJ-mnist-cnn_model-sample.pth')
329
- # print('Epoch: %d' % epoch)
330
- # acc = val(model, device, train_data_loader)
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(