You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch自定义SNN部分可学习参数未更新问题排查求助

脉冲神经网络(SNN)训练仅部分参数更新的原因分析

我用PyTorch从零实现了一个脉冲神经网络(SNN),定义了7个可学习参数,网络架构为2x10x1,用于实现两数相加功能。但训练后发现只有参数a和th2得到更新,其余参数均无变化。具体实现代码如下:

import torch
import torch.nn as nn

class SNN(nn.Module):
    def __init__(self,window_length,beta,dims):
        super(SNN,self).__init__()
        self.window_length=window_length
        self.beta=beta
        self.input_size=dims[0]
        self.hidden_size=dims[1]
        self.output_size=dims[2]
        
        self.a=nn.Parameter(torch.ones(1,requires_grad=True))
        self.W0=nn.Parameter(torch.ones(self.input_size,requires_grad=True))
        self.W1=nn.Parameter(torch.ones(self.hidden_size,requires_grad=True))
        self.W2=nn.Parameter(torch.ones(self.output_size,requires_grad=True))
        self.th0=nn.Parameter(torch.ones(self.input_size,requires_grad=True))
        self.th1=nn.Parameter(torch.ones(self.hidden_size,requires_grad=True))
        self.th2=nn.Parameter(3*torch.ones(self.output_size,requires_grad=True))

    def Layer(self,x):
        #I have three of such methods (for my three layers)
        n_samples=x.shape[0]
        window_length=self.window_length
        layer_size=self.hidden_size
        
        mems_sur_layer_output=torch.zeros(n_samples,layer_size,window_length)
        sout_sur_layer_output=torch.zeros(n_samples,layer_size,window_length)
        
        for i in range(layer_size):            
            sout_sur=torch.zeros(n_samples,window_length)
            mems_sur=torch.zeros(n_samples,window_length)
            mem=torch.zeros(n_samples)
            mem_sur=torch.zeros(n_samples)
            
            for time_step in range(window_length):
                mems_sur[:,time_step]=(self.beta*mem_sur+self.W0[i]*torch.sum(x[:,:,time_step],dim=[1]))
                #W0 is replaced with W1 and W2 in the two other layers
                
                x1=mems[:,time_step]-self.th0[i]  # 此处存在笔误,应为mems_sur[:,time_step]
                #th0 is replaced with th1 and th2 in the two other layers
                s_sur=torch.sigmoid(x1)
                sout_sur[:,time_step]=s_sur

                mem_sur=mems_sur[:,time_step]*(1-s_sur)
        
            mems_sur_layer_output[:,i,:]=mems_sur
            sout_sur_layer_output[:,i,:]=sout_sur   

        return sout_sur_layer_output

    def Rate_Decode(self,x):
        return x/self.a

    def forward(self,x):
        output=self.Layer(x)
        #Layer() is called 3 times with the variable 'output' being the output of each call and the input of the next.
        return self.Rate_Decode(output)

可能的原因分析

  • 参数未实际参与计算:
    你提到有三个对应不同层的Layer方法,分别使用W0/W1/W2和th0/th1/th2,但当前forward方法只调用了一次Layer,且现有Layer方法硬编码使用W0和th0。如果实际代码中没有正确调用另外两个层的方法,W1、W2、th1根本没参与前向传播,自然不会产生梯度更新。

  • 代码笔误导致无效计算:
    在Layer方法的时间步循环中,x1=mems[:,time_step]-self.th0[i]存在笔误——mems是未被赋值的初始零张量,应该是mems_sur[:,time_step]。这个错误会导致s_sur的计算和当前层的膜电位无关,前向输出完全脱离输入和参数W0、th0的控制,梯度无法正确回传到这些参数。

  • 循环结构破坏梯度追踪:
    你用torch.zeros初始化了mems_sur_layer_output、sout_sur_layer_output等张量,然后通过循环逐个赋值。这种操作在PyTorch中会破坏计算图的连续性,导致梯度无法回溯到前面的参数。建议改用可微分的张量拼接或者预分配带梯度的张量(比如用torch.empty后填充,或直接用张量运算代替循环)。

  • 激活函数梯度消失:
    你使用sigmoid作为脉冲发放的近似函数,当输入x1的绝对值较大时,sigmoid的梯度会趋近于0。如果初始参数设置(比如W0全1、th0全1)导致x1的取值落在梯度饱和区,W0、th0等参数的梯度会几乎为0,无法得到有效更新。

  • 初始化与任务匹配问题:
    你的任务是两数相加,初始参数全1的设置可能让前面的层直接进入饱和状态,输出无法区分不同输入,导致梯度无法传递到前面的参数。只有最后一层的th2和缩放参数a能通过损失函数获得有效梯度。

内容的提问来源于stack exchange,提问作者Ahmed Zeid

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 13:43:14