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

PyTorch自动编码器调用.to(device)时内存CPU耗尽,求助解决

自动编码器初始化内存溢出问题排查与解决

问题现象

运行自动编码器代码时,在执行autoencoder = Autoencoder().to(device)步骤出现RAM/CPU被完全占满、程序卡顿,甚至在Google Colab中直接因资源不足崩溃。相关代码及运行输出如下:

原自动编码器代码

import torch
import torch.nn as nn

class Autoencoder(nn.Module):
    def __init__(self):
        super(Autoencoder, self).__init__()
        
        self.flatten = nn.Flatten()
        
        self.enc_conv0 = nn.Sequential(
            nn.Conv2d(in_channels=3, out_channels=64, kernel_size=3, padding=(1, 1)),
            nn.ReLU(),
            nn.BatchNorm2d(64),

            nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3, padding=(1, 1)),
            nn.ReLU(),
            nn.BatchNorm2d(128)
        )
        
        self.enc_conv1 = nn.Sequential(
            nn.Conv2d(in_channels=128, out_channels=256, kernel_size=3, padding=(1, 1)),
            nn.ReLU(),
            nn.BatchNorm2d(256),

            nn.Conv2d(in_channels=256, out_channels=512, kernel_size=3, padding=(1, 1)),
            nn.ReLU(),
            nn.BatchNorm2d(512)
        )
        
        self.enc_fc = nn.Sequential(
            nn.Linear(in_features=512*64*64, out_features=4096),
            nn.ReLU(),
            nn.BatchNorm1d(4096),
            
            nn.Linear(in_features=4096, out_features=2048),
            nn.ReLU(),
            nn.BatchNorm1d(2048),
            
            nn.Linear(in_features=2048, out_features=dim_code)
        )
        
        self.dec_fc = nn.Sequential(
            nn.Linear(in_features=dim_code, out_features=2048),
            nn.ReLU(),
            nn.BatchNorm1d(2048),
            
            nn.Linear(in_features=2048, out_features=4096),
            nn.ReLU(),
            nn.BatchNorm1d(4096),
            
            nn.Linear(in_features=4096, out_features=512*64*64),
            nn.ReLU(),
            nn.BatchNorm1d(512*64*64)
        )
        
        self.dec_conv0 = nn.Sequential(
            nn.ConvTranspose2d(in_channels=512, out_channels=256, kernel_size=(3,3), padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(256),
            
            nn.ConvTranspose2d(in_channels=256, out_channels=128, kernel_size=(3,3), padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(128),
        )
        
        self.dec_conv1 = nn.Sequential(
            nn.ConvTranspose2d(in_channels=128, out_channels=64, kernel_size=(3,3), padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(64),
            
            nn.ConvTranspose2d(in_channels=64, out_channels=3, kernel_size=(3,3), padding=1)
        )

    def forward(self, x):
        e0 = self.enc_conv0(x)
        e1 = self.enc_conv1(e0)
        latent_code = self.enc_fc(self.flatten(e1))
        
        d0 = self.dec_fc(latent_code)
        d1 = self.dec_conv0(d0.view(-1, 512, 64, 64))
        reconstruction = self.dec_conv1(d1)

        return reconstruction, latent_code

训练初始化代码

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(device)

criterion = nn.BCELoss()
print('crit')

autoencoder = Autoencoder().to(device)
print('deviced')

运行输出

cuda
'crit'

问题根源

核心问题是全连接层参数规模远超硬件内存上限:

  • 编码器第一个全连接层nn.Linear(512*64*64, 4096)的权重参数数量为512*64*64*4096 = 8589934592(约86亿),单个层就需要数十GB内存才能存储,直接超出Kaggle/Colab提供的GPU内存配额。
  • 解码器对应的反向全连接层存在同样的超大参数问题,导致模型初始化时内存被瞬间占满。
  • 原编码器未对卷积后的特征图做下采样,输入64x64图像经过卷积后尺寸仍保持64x64,进一步放大了全连接层的输入维度。

解决方案

方案1:添加下采样缩小特征图,降低全连接层维度

在编码器的卷积块后加入MaxPool2d下采样,减小特征图宽高;解码器用带stride的转置卷积恢复图像尺寸,大幅压缩全连接层的参数规模:

import torch
import torch.nn as nn

dim_code = 128  # 提前定义latent code维度

class Autoencoder(nn.Module):
    def __init__(self):
        super(Autoencoder, self).__init__()
        
        self.flatten = nn.Flatten()
        
        # 编码器:加入MaxPool下采样,特征图从64->32->16
        self.enc_conv0 = nn.Sequential(
            nn.Conv2d(in_channels=3, out_channels=64, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(64),
            nn.MaxPool2d(kernel_size=2, stride=2),  # 64x64 → 32x32
            
            nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(128)
        )
        
        self.enc_conv1 = nn.Sequential(
            nn.Conv2d(in_channels=128, out_channels=256, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(256),
            nn.MaxPool2d(kernel_size=2, stride=2),  # 32x32 →16x16
            
            nn.Conv2d(in_channels=256, out_channels=512, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(512)
        )
        
        # 全连接层输入维度压缩为512*16*16=131072,参数规模大幅降低
        self.enc_fc = nn.Sequential(
            nn.Linear(in_features=512*16*16, out_features=4096),
            nn.ReLU(),
            nn.BatchNorm1d(4096),
            
            nn.Linear(in_features=4096, out_features=2048),
            nn.ReLU(),
            nn.BatchNorm1d(2048),
            
            nn.Linear(in_features=2048, out_features=dim_code)
        )
        
        self.dec_fc = nn.Sequential(
            nn.Linear(in_features=dim_code, out_features=2048),
            nn.ReLU(),
            nn.BatchNorm1d(2048),
            
            nn.Linear(in_features=2048, out_features=4096),
            nn.ReLU(),
            nn.BatchNorm1d(4096),
            
            nn.Linear(in_features=4096, out_features=512*16*16),
            nn.ReLU(),
            nn.BatchNorm1d(512*16*16)
        )
        
        # 解码器:用转置卷积的stride恢复尺寸,16x16->32x32->64x64
        self.dec_conv0 = nn.Sequential(
            nn.ConvTranspose2d(in_channels=512, out_channels=256, kernel_size=3, padding=1, stride=2, output_padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(256),
            
            nn.ConvTranspose2d(in_channels=256, out_channels=128, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(128),
        )
        
        self.dec_conv1 = nn.Sequential(
            nn.ConvTranspose2d(in_channels=128, out_channels=64, kernel_size=3, padding=1, stride=2, output_padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(64),
            
            nn.ConvTranspose2d(in_channels=64, out_channels=3, kernel_size=3, padding=1)
        )

    def forward(self, x):
        e0 = self.enc_conv0(x)  # [B,128,32,32]
        e1 = self.enc_conv1(e0)  # [B,512,16,16]
        latent_code = self.enc_fc(self.flatten(e1))  # [B, dim_code]
        
        d0 = self.dec_fc(latent_code)  # [B,512*16*16]
        d1 = self.dec_conv0(d0.view(-1, 512, 16, 16))  # [B,128,32,32]
        reconstruction = self.dec_conv1(d1)  # [B,3,64,64]

        return reconstruction, latent_code

方案2:全卷积自动编码器(移除全连接层)

完全去掉全连接层,仅用卷积和转置卷积完成编码解码,参数规模更可控,也更适配图像任务:

import torch
import torch.nn as nn

dim_code = 128

class Autoencoder(nn.Module):
    def __init__(self):
        super(Autoencoder, self).__init__()
        
        # 编码器:卷积+下采样
        self.encoder = nn.Sequential(
            nn.Conv2d(3, 64, 3, padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(64),
            nn.MaxPool2d(2, 2),  # 64→32
            
            nn.Conv2d(64, 128, 3, padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(128),
            nn.MaxPool2d(2, 2),  #32→16
            
            nn.Conv2d(128, 256, 3, padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(256),
            nn.MaxPool2d(2, 2),  #16→8
            
            nn.Conv2d(256, dim_code, 3, padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(dim_code)
        )
        
        # 解码器:转置卷积+上采样
        self.decoder = nn.Sequential(
            nn.ConvTranspose2d(dim_code, 256, 3, padding=1, stride=2, output_padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(256),
            
            nn.ConvTranspose2d(256, 128, 3, padding=1, stride=2, output_padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(128),
            
            nn.ConvTranspose2d(128, 64, 3, padding=1, stride=2, output_padding=1),
            nn.ReLU(),
            nn.BatchNorm2d(64),
            
            nn.Conv2d(64, 3, 3, padding=1)
        )

    def forward(self, x):
        latent = self.encoder(x)  # [B, dim_code, 8,8]
        # 若需要向量形式的latent code,可添加全局平均池化
        latent_vec = torch.mean(latent, dim=(2,3))
        reconstruction = self.decoder(latent)
        return reconstruction, latent_vec

额外注意事项

  • 若使用BCELoss,需给模型输出添加sigmoid激活,或改用BCEWithLogitsLoss(直接接收未归一化的输出),否则损失计算会出现异常。
  • 必须提前定义dim_code变量,原代码中未对该变量赋值,会导致运行报错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 17:55:55