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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 06:04:57