PyTorch自定义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

