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

ProGAN添加Self Attention模块后大输入VRAM占用过高问题咨询

问题:ProGAN添加Self Attention后256x256尺寸显存溢出优化

问题背景

在Progressive GAN(ProGAN)的最后一层(to RGB之前)添加Self Attention模块后,当模型输出尺寸达到256x256时进程被终止。测试发现两个Attention实现均占用约16G VRAM,无法在3090显卡上运行完整模型,但32x32尺寸下训练正常。

测试代码

import torch
import torch.nn as nn

class SelfAttention(nn.Module):
    def __init__(self, channels):
        super(SelfAttention, self).__init__()
        self.channels = channels
        num_heads = 4
        self.mha = nn.MultiheadAttention(channels, num_heads, batch_first=True)
        self.ln = nn.LayerNorm([channels])
        self.ff_self = nn.Sequential(
            nn.LayerNorm([channels]),
            nn.Linear(channels, channels),
            nn.GELU(),
            nn.Linear(channels, channels)
        )

    def forward(self, x):
        size = x.shape[3]
        print("SIZE", size)
        print("CHANNELS", self.channels)
        x = x.view(-1, self.channels, size * size).swapaxes(1, 2)
        print()
        x_ln = self.ln(x)
        attention_value, _ = self.mha(x_ln, x_ln, x_ln)
        attention_value = attention_value + x
        attention_value = self.ff_self(attention_value) + attention_value
        return attention_value.swapaxes(2, 1).view(-1, self.channels, size, size)


class SelfAttention2(nn.Module):
    def __init__(self, channels):
        super(SelfAttention2, self).__init__()
        self.query = nn.Conv2d(channels, channels // 8, kernel_size=1)
        self.key = nn.Conv2d(channels, channels // 8, kernel_size=1)
        self.value = nn.Conv2d(channels, channels, kernel_size=1)
        self.softmax = nn.Softmax(dim=-1)

    def forward(self, x):
        N, C, H, W = x.size()
        query = self.query(x).view(N, -1, W*H).permute(0, 2, 1) # (N, C, H*W)
        key = self.key(x).view(N, -1, W*H) # (N, C, H*W)
        energy = torch.bmm(query, key) # (N, H*W, H*W)
        attention = self.softmax(energy)
        value = self.value(x).view(N, -1, W*H) # (N, C, H*W)
        out = torch.bmm(value, attention.permute(0, 2, 1))
        out = out.view(N, C, H, W)
        return out


if __name__ == '__main__':
    
    x = torch.randn((1, 64, 256, 256))
    
    sa1 = SelfAttention(64)
    sa1(x)

    sa2 = SelfAttention2(64)
    sa2(x)

优化建议

  • 缩小注意力计算的空间维度
    无需在256x256的全空间计算注意力,先通过stride=2的卷积将特征图下采样至128x128或64x64,计算注意力后再通过转置卷积上采样回原尺寸。此举可将注意力矩阵从(256256)×(256256)缩小至(128128)×(128128),显存占用直接降至原来的1/4。

  • 优化通道维度的注意力拆分
    SelfAttention2已采用通道拆分(channels//8),可进一步降低比例至channels//16;同时用分组卷积替代普通卷积减少参数总量。SelfAttention1中的MultiheadAttention可调整head数,比如增加head数同时降低单head通道数(保持总通道数不变),减少单head的计算显存开销。

  • 替换为FlashAttention实现
    PyTorch 2.0+提供的torch.nn.functional.scaled_dot_product_attention(FlashAttention)比原生MultiheadAttention更节省显存,计算速度更快。将SelfAttention1中的mha替换为该函数,可大幅降低显存占用。

  • 调整Attention的接入位置
    不必局限于在最后一层(to RGB前)添加Attention,可选择在128x128尺寸的过渡层插入,此时特征图尺寸更小,显存压力更低,同样能提升模型细节生成能力。

  • 启用混合精度训练
    通过torch.cuda.amp.autocast()启用混合精度训练,将大部分计算从float32转为float16,显存占用直接减半,且3090显卡支持FP16加速,几乎不影响训练精度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 17:35:30