基于LIF神经元的SNN-UNet迭代变慢且显存溢出问题求助
问题描述
我正在实现一个基于脉冲神经网络(SNN)的全卷积UNet。测试不含LIF神经元的普通模型时运行正常,但内存占用高,batch size必须小于4。换成LIF神经元后,仅跑5次迭代就报CUDA显存不足,而且每次迭代速度越来越慢。输入是4通道512×512图像,哪怕batch size=1、num_steps=1(无时间维度)还是会显存溢出,说明不是硬件性能不够,而是存在内存泄漏。试过加torch.cuda.empty_cache()和snnTorch的utils.reset(model)清理内存,没用。作为PyTorch新手,找不到问题所在,以下是实现代码:
# UNet parts class DoubleConv(nn.Module): """(convolution => [BNTT] => Spikes) * 2 + Dropout""" def __init__(self, in_channels, out_channels, beta, SG_func, mid_channels=None): super().__init__() if not mid_channels: mid_channels = out_channels self.double_conv = nn.Sequential( nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1), #snn.bntt.BatchNormTT2d(mid_channels, time_steps=num_steps), snn.Leaky(beta=beta, spike_grad=SG_func, init_hidden=True), nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1), #snn.bntt.BatchNormTT2d(out_channels, time_steps=num_steps), snn.Leaky(beta=beta, spike_grad=SG_func, init_hidden=True), nn.Dropout2d() ) def forward(self, x): UT.reset(self) out_spks= self.double_conv(x) return out_spks class Down(nn.Module): """Downscaling with maxpool then double conv""" def __init__(self, in_channels, out_channels, beta, SG_func): super().__init__() self.maxpool_conv = nn.Sequential( nn.MaxPool2d(2), DoubleConv(in_channels, out_channels, beta=beta, SG_func=SG_func) ) def forward(self, x): return self.maxpool_conv(x) class Up(nn.Module): """Upscaling then double conv""" def __init__(self, in_channels, out_channels, beta, SG_func, bilinear=True): super().__init__() # transposed conv option omiited, only bilinear used instead self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True) self.conv = DoubleConv(in_channels, out_channels, beta=beta, SG_func=SG_func, mid_channels=in_channels // 2) def forward(self, x1, x2): x1 = self.up(x1) # input is BxCxHxW diffY = x2.size()[2] - x1.size()[2] diffX = x2.size()[3] - x1.size()[3] x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]) x = torch.cat([x2, x1], dim=1) return self.conv(x) class OutConv(nn.Module): """Output conv with 1x1 kernel""" def __init__(self, in_channels, out_channels): super(OutConv, self).__init__() self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1) def forward(self, x): return self.conv(x) # full Unet model class UNet(nn.Module): def __init__(self, n_channels, n_classes, bilinear=True, beta=0.5, SG_func=surrogate.ATan.apply): super(UNet, self).__init__() self.n_channels = n_channels self.n_classes = n_classes self.bilinear = bilinear self.beta = beta self.SG_func = SG_func self.inc = DoubleConv(n_channels, 64, beta, SG_func) self.down1 = Down(64, 128, beta, SG_func) self.down2 = Down(128, 256, beta, SG_func) self.down3 = Down(256, 512, beta, SG_func) self.down4 = Down(512, 1024 // 2, beta, SG_func) self.up1 = Up(1024, 512 // 2, beta, SG_func) self.up2 = Up(512, 256 // 2, beta, SG_func) self.up3 = Up(256, 128 // 2,beta, SG_func) self.up4 = Up(128, 64, beta, SG_func) self.outc = OutConv(64, n_classes) def forward(self, x): spikes = [] torch.cuda.empty_cache() UT.reset(self) for step in range(num_steps): x1 = self.inc(x[step]) x2 = self.down1(x1) x3 = self.down2(x2) x4 = self.down3(x3) x5 = self.down4(x4) x6 = self.up1(x5, x4) x7 = self.up2(x6, x3) x8 = self.up3(x7, x2) x9 = self.up4(x8, x1) spikes.append(x9) x1 = x1.detach() x2 = x2.detach() x3 = x3.detach() x4 = x4.detach() x5 = x5.detach() x6 = x6.detach() x7 = x7.detach() x8 = x8.detach() x9 = x9.detach() accumulated_spks=torch.sum(torch.stack(spikes), dim=0) accumulated_spks=accumulated_spks.detach() output = self.outc(accumulated_spks) return output # very simple training loop just to mess around with def training(loss_f, model, Optimizer, dataloader): training_acc = 0 batch_no = 0 model.train() # set network to training mode for batch in dataloader: torch.cuda.empty_cache() print(f"iteration: {batch_no}") start = time.time() X = spikegen.rate(batch["X"], num_steps=num_steps).to(device) # ensure network and data are running on GPU if available Y = batch["Y"].to(device) pred = model(X) # compute network output loss = loss_f(pred, Y) # compute loss and backpropagate on it Optimizer.zero_grad() loss.backward() Optimizer.step() stop = time.time() batch_no += 1 if batch_no == 5: break; if batch_no % 50 == 0: print(f"Loss at batch no: {batch_no}: {loss}") PATH = os.path.join(root_dir, "model"+".pth") model = UNet(n_channels = 4, n_classes = 1).to(device) torch.save(model.state_dict(), PATH) loss_f = nn.BCEWithLogitsLoss() optimizer = torch.optim.Adam(model.parameters(), lr = 1e-3) training(loss_f, model, optimizer, train_dataloader)
问题定位与修复方案
1. LIF神经元隐藏状态未正确重置
snnTorch的Leaky神经元在init_hidden=True时会维护内部隐藏状态,但你的代码存在两个问题:
DoubleConv中调用UT.reset(self)无法递归重置nn.Sequential里的Leaky神经元;- UNet的forward中重复调用
UT.reset(self),导致状态重置不彻底。
修复:
删除DoubleConv中的UT.reset(self),在UNet的forward循环前添加递归重置:
def forward(self, x): spikes = [] # 递归重置所有子模块的LIF神经元状态 UT.reset(self, recurse=True) for step in range(num_steps): x1 = self.inc(x[step]) # 剩余代码不变
2. 不必要的detach操作引发内存堆积
你对x1到x9以及accumulated_spks的detach操作不仅切断了反向传播链路,还会导致张量引用无法被GC正常回收,加剧内存泄漏。
修复:
删除所有不必要的detach(),并在循环后主动删除中间张量引用:
def forward(self, x): spikes = [] UT.reset(self, recurse=True) for step in range(num_steps): x1 = self.inc(x[step]) x2 = self.down1(x1) x3 = self.down2(x2) x4 = self.down3(x3) x5 = self.down4(x4) x6 = self.up1(x5, x4) x7 = self.up2(x6, x3) x8 = self.up3(x7, x2) x9 = self.up4(x8, x1) spikes.append(x9) # 主动删除中间张量引用,帮助GC回收 del x1, x2, x3, x4, x5, x6, x7, x8, x9 accumulated_spks = torch.sum(torch.stack(spikes), dim=0) output = self.outc(accumulated_spks) return output
3. 滥用torch.cuda.empty_cache()
该函数仅释放未被PyTorch占用的显存块,无法释放正在使用的张量内存,频繁调用会增加显存分配开销,导致迭代变慢。
修复:
删除训练循环和forward中的torch.cuda.empty_cache(),仅在epoch结束等必要场景调用。
4. 额外优化建议
- 启用PyTorch自动混合精度训练,降低显存占用:
from torch.cuda.amp import GradScaler, autocast def training(loss_f, model, Optimizer, dataloader): scaler = GradScaler() training_acc = 0 batch_no = 0 model.train() for batch in dataloader: print(f"iteration: {batch_no}") start = time.time() X = spikegen.rate(batch["X"], num_steps=num_steps).to(device) Y = batch["Y"].to(device) Optimizer.zero_grad() with autocast(): pred = model(X) loss = loss_f(pred, Y) scaler.scale(loss).backward() scaler.step(Optimizer) scaler.update() stop = time.time() batch_no += 1 if batch_no == 5: break if batch_no % 50 == 0: print(f"Loss at batch no: {batch_no}: {loss}")
- 用
nn.ConvTranspose2d替代双线性插值,更适配SNN的离散脉冲特性,减少内存开销。
内容的提问来源于stack exchange,提问作者Kernel_Sahanderz
相关产品推荐
相关产品推荐

