Forward Functions#

SpikingJelly 的 前向传播函数 实现了 SNN 的多步前向传播逻辑。


SpikingJelly's forward functions provide multi-step forward propagation logic for SNNs.

spikingjelly.activation_based.functional.forward.multi_step_forward(x_seq, single_step_module)[源代码]#

API Language - 中文 | English


  • 中文

在单步模块 single_step_module 上使用多步前向传播。函数内部将执行一个for循环, 执行 T 次单步前向传播。若 single_step_module 为多个模块,则每个时间步都会按顺序依次执行这些模块。

参数:
  • x_seq (Tensor) -- shape=[T, batch_size, ...] 的输入tensor

  • single_step_module (Union[nn.Module, list[nn.Module], tuple[nn.Module], nn.Sequential, Callable]) -- 一个或多个单步模块

返回:

shape=[T, batch_size, ...] 的输出tensor

返回类型:

Tensor

抛出:

Exception -- 任何底层模块在某个时间步前向传播时抛出的异常都会原样向上传播


  • English

Applies multi-step forward on single_step_module. The function runs a for loop to execute single-step forward for T times. If single_step_module contains multiple modules, they are applied sequentially at each time-step.

参数:
  • x_seq (Tensor) -- the input tensor with shape=[T, batch_size, ...]

  • single_step_module (Union[nn.Module, list[nn.Module], tuple[nn.Module], nn.Sequential, Callable]) -- one or many single-step modules

返回:

the output tensor with shape=[T, batch_size, ...]

返回类型:

Tensor

抛出:

Exception -- Any exception raised by an underlying module at any time step is propagated unchanged

spikingjelly.activation_based.functional.forward.t_last_multi_step_forward(x_seq, single_step_module)[源代码]#

API Language - 中文 | English


  • 中文

在单步模块 single_step_module 上使用多步前向传播。

此函数适用于时间维位于最后一维的序列张量,即 shape=[batch_size, ..., T]。 它会沿最后一维逐个时间步取出切片,并在每个时间步顺序执行单步模块。

参数:
  • x_seq (Tensor) -- shape=[batch_size, ..., T] 的输入tensor

  • single_step_module (Union[nn.Module, list[nn.Module], tuple[nn.Module], nn.Sequential, Callable]) -- 一个或多个单步模块

返回:

shape=[batch_size, ..., T] 的输出tensor

返回类型:

Tensor

抛出:

Exception -- 任何底层模块在某个时间步前向传播时抛出的异常都会原样向上传播


  • English

Apply multi-step forward on single_step_module.

This helper is intended for sequence tensors whose time axis is the last dimension, i.e. shape=[batch_size, ..., T]. It slices along the last dimension and applies the single-step module(s) at each time step.

参数:
  • x_seq (Tensor) -- the input tensor with shape=[batch_size, ..., T]

  • single_step_module (Union[nn.Module, list[nn.Module], tuple[nn.Module], nn.Sequential, Callable]) -- one or many single-step modules

返回:

the output tensor with shape=[batch_size, ..., T]

返回类型:

Tensor

抛出:

Exception -- Any exception raised by an underlying module at any time step is propagated unchanged

spikingjelly.activation_based.functional.forward.chunk_multi_step_forward(split_size, x_seq, multi_step_module)[源代码]#

API Language - 中文 | English


  • 中文

shape = [T, *] 的输入 x_seq 拆分成多个 shape = [split_size, *] 的小tensor(若 T % split_size != 0,最后一个tensor的 shape[0] 会小于 split_size),然后逐个输入到 multi_step_module 中,再沿着 dim=0 将输出重新拼接,因此输出的首维长度仍为 T

chunk_multi_step_forward 可以在使用很大的 T 进行不带梯度的推理(例如ANN2SNN)时使用,能够减少内存消耗量。

参数:
  • split_size (int) -- 分割的尺寸

  • x_seq (Tensor) -- 输入

  • multi_step_module (Module) -- 一个使用多步传播模式的网络

返回:

输出

返回类型:

Tensor

抛出:

Exception -- 任何 multi_step_module 在某个分块上的前向传播异常都会原样向上传播


  • English

Splits the input x_seq with shape = [T, *] to many tensor chunks with shape = [split_size, *] (if T % split_size != 0, shape[0] of the last tensor chunk will be smaller than split_size), and sends chunks to multi_step_module, then concatenates the outputs back along dim=0, so the output keeps the original leading length T.

chunk_multi_step_forward can be used for inference with a large T (e.g., ANN2SNN) to reduce the memory consumption.

参数:
  • split_size (int) -- the split size

  • x_seq (Tensor) -- the input tensor

  • multi_step_module (nn.Module) -- a network in multi-step mode

返回:

the output tensor

返回类型:

Tensor

抛出:

Exception -- Any exception raised by multi_step_module on a chunk is propagated unchanged


  • 代码示例 | Example

import torch
import torch.nn as nn
from spikingjelly.activation_based import neuron, layer, functional

net = nn.Sequential(
    layer.Linear(8, 4),
    neuron.IFNode(step_mode="m"),
    layer.Linear(4, 2),
    neuron.IFNode(step_mode="m"),
)

x_seq = torch.rand([1024, 8])
with torch.no_grad():
    y_seq = functional.chunk_multi_step_forward(16, x_seq, net)
    print(y_seq.shape)
    # torch.Size([1024, 2])
spikingjelly.activation_based.functional.forward.seq_to_ann_forward(x_seq, stateless_module)[源代码]#

API Language - 中文 | English


  • 中文

使用无状态层进行多步前向传播。输入 x_seq 的时间和批量维度将被展平,得到 [T*batch_size, ...] 形状的张量;随后,输入到无状态层中;最后,将输出张量恢复到序列形式 [T, batch_size, ...]

x_seq 也可以是 tensor tuple,例如同时输入池化值与池化索引。此时每个 tensor 的形状均为 shape=[T, batch_size, ...],且 Tbatch_size 必须相同;每个 tensor 的时间和批量维度会被分别展平,展平后的 tensor 作为 位置参数一起输入到第一个无状态层中。若给出多个无状态层,则后续的层依次接收 前一层的输出作为单个参数。因此第一个无状态层必须能够接收与tuple长度相同数量的 位置参数;torch.nn.Sequential 只接收单个输入,故不能作为第一层接收长度 大于1的tuple。

参数:
  • x_seq (Union[Tensor, tuple[Tensor, ...]]) -- shape=[T, batch_size, ...] 的输入tensor,或多个此类tensor组成的tuple

  • stateless_module (Union[Module, list, tuple, Sequential, Callable]) -- 单个或多个无状态网络层

返回:

shape=[T, batch_size, ...] 的输出tensor;若底层模块返回 tensor tuple,则分别恢复每个 tensor 的时间维和批量维

返回类型:

Union[Tensor, tuple[Tensor, ...]]

抛出:
  • ValueError -- 当tuple中tensor的 [T, batch_size] 前两维不一致时抛出

  • Exception -- 任何底层无状态模块在前向传播时抛出的异常都会原样向上传播


  • English

Applied forward on stateless modules. Flatten the time and batch dimensions of x_seq so that shape=[T*batch_size, ...], feed the reshaped tensor to the stateless module(s), and reshape the output back to the sequence form shape=[T, batch_size, ...].

x_seq can also be a tuple of tensors, e.g., pooled values together with pooling indices. In this case, every tensor must have shape=[T, batch_size, ...] with the same T and batch_size; the time and batch dimensions of each tensor are flattened separately, and the flattened tensors are fed to the first stateless module as positional arguments. If several stateless modules are given, each subsequent module receives the previous module's output as a single argument. The first stateless module must therefore accept as many positional arguments as the tuple holds; torch.nn.Sequential accepts a single input only, so it can not be the first module for a tuple of more than one tensor.

参数:
  • x_seq (Union[Tensor, tuple[Tensor, ...]]) -- the input tensor with shape=[T, batch_size, ...], or a tuple of such tensors

  • stateless_module (Union[Module, list, tuple, Sequential, Callable]) -- one or many stateless modules

返回:

the output tensor with shape=[T, batch_size, ...]; if the underlying module returns a tuple of tensors, each tensor has its time and batch dimensions restored

返回类型:

Union[Tensor, tuple[Tensor, ...]]

抛出:
  • ValueError -- if the tensors in x_seq do not share the same [T, batch_size] leading dimensions

  • Exception -- Any exception raised by an underlying stateless module is propagated unchanged

spikingjelly.activation_based.functional.forward.t_last_seq_to_ann_forward(x_seq, stateless_module)[源代码]#

API Language - 中文 | English


  • 中文

使用无状态层进行多步前向传播。

备注

SpikingJelly中默认序列数据形状为 shape=[T, batch_size, ...]。 但此函数是用于另一种格式,即 shape=[batch_size, ..., T]。 此函数使用 torch.vmap 沿最后一维执行单步前向传播。listtuple 中的模块会被逐项调用;nn.Sequential 作为容器调用。

备注

不能用于BN层,因为BN层的running mean/var是输入依赖的。 对于BN层,只需要输入被当作是 shape = [N, C, ..] 即可并行计算,需要用户手动实现。

参数:
  • x_seq (Tensor) -- shape=[batch_size, ..., T] 的输入tensor

  • stateless_module (Union[Module, list, tuple, Sequential, Callable]) -- 单个或多个无状态网络层

返回:

shape=[batch_size, ..., T] 的输出tensor;若底层模块返回 tensor tuple,则分别恢复每个 tensor 的时间维

返回类型:

Union[Tensor, tuple[Tensor, ...]]

抛出:

Exception -- 任何底层无状态模块在前向传播时抛出的异常都会原样向上传播


  • English

Applied forward on stateless modules.

Note

The default shape of sequence data in SpikingJelly is shape=[T, batch_size, ...]. However, this function is used for the other data format where shape=[batch_size, ..., T]. This function uses torch.vmap to apply the single-step forward pass over the last dimension. Modules in a list or tuple are called in order, while an nn.Sequential is called as a container.

Note

This function can not be applied to wrap BN because its running mean/var depends on inputs. The BN can be computed in parallel as long as the input is regarded as shape = [N, C, ..], which can be implemented by user manually.

参数:
  • x_seq (Tensor) -- the input tensor with shape=[batch_size, ..., T]

  • stateless_module (Union[Module, list, tuple, Sequential, Callable]) -- one or many stateless modules

返回:

the output tensor with shape=[batch_size, ..., T]; if the underlying module returns a tuple of tensors, each tensor has its time dimension restored

返回类型:

Union[Tensor, tuple[Tensor, ...]]

抛出:

Exception -- Any exception raised by an underlying stateless module is propagated unchanged