← 返回 PaperDaily 大模型与智能体

DenseNet内存优化:264层单机可训,ImageNet冲到20.26%

DenseNet真正的短板不是算力,而是显存。本文把“存不下”这个工程老大难拆成共享内存和按需重算两步,硬是把264层模型塞进单机训练里。

DenseNet内存优化:264层单机可训,ImageNet冲到20.26%
🐉 龙哥读论文知识星球来了!
公众号每日8篇拆解不够看?星球无上限更AI领域论文、资讯、招聘、招博、开源代码,一站式干货,每日2分钟刷完即赚!
👇扫码加入「龙哥读论文」知识星球,前沿干货、实用资源一站式拿捏~ xingqiu_header

龙哥导读:
DenseNet真正的短板不是算力,而是显存。本文把“存不下”这个工程老大难拆成共享内存和按需重算两步,硬是把264层模型塞进单机训练里。


原论文信息如下:
论文标题:
Memory-Efficient Implementation of DenseNets
发表日期:
2017年07月
发表单位:
Cornell University, Facebook AI Research, Fudan University
原文链接:
https://arxiv.org/pdf/1707.06990.pdf

内存瓶颈:极深DenseNet的训练难题

DenseNet 的尴尬很典型:模型本身很省参数,训练时却可能很吃显存。这就像一台车发动机很省油,结果后备箱塞满了行李,跑得再优雅也得先想办法把东西装下。
这篇技术报告讨论的不是“DenseNet 好不好”,而是一个更接地气的问题:为什么 DenseNet 明明算得不多,GPU 却先顶不住了?答案不在算法本身,而在实现方式。原论文指出,很多朴素实现会把中间特征图、归一化结果、拼接结果以及它们的梯度都老老实实存着,结果显存占用随着网络深度迅速膨胀,最后把“继续堆深度”这条路直接堵死。
图1:DenseNet架构的高层示意图
图1:DenseNet 架构的高层示意图。每一层都直接连接到此前所有层,核心思想就是“特征复用”,而不是层层重复造轮子。
DenseNet 的连接方式很特别:第 层输入不是上一层的输出,而是前面所有层输出的拼接。这样做的好处很直接——早期特征能被后面每一层反复利用,参数也不用像普通网络那样一层层“重新发明”。但坏处也很现实:拼接和归一化这些中间步骤,在朴素实现里会制造大量临时张量,GPU 显存就开始吃紧。

揭秘根源:预激活与连续拼接的内存陷阱

这篇文章把问题拆得很清楚:DenseNet 的显存膨胀,主要不是来自“最终输出”,而是来自两个容易被忽略的细节——预激活批归一化连续拼接。这两个操作单独看都很正常,放到“每层都连所有层”的 DenseNet 里,就开始放大成内存杀手。
图2:预激活与后激活DenseNet的对比,以及连续与非连续卷积操作的速度对比
图2:左图对比了预激活与后激活 DenseNet 架构。预激活能明显降低错误率,但会产生数量可观的中间输出;右图对比了连续与非连续卷积操作的速度,连续内存更快,却往往意味着要复制更多特征。
先看预激活。这里的“预激活”指的是把批归一化(Batch Normalization, BN)和 ReLU 放在卷积前面。BN 的全称是 Batch Normalization,中文一般叫批归一化。原论文引用的是 He 等人在残差网络里的预激活设计。这个设计对精度有帮助,但代价是每层都可能产生一份“经过归一化的旧特征副本”。层数一多,这些副本就像外卖盒子一样越堆越多,显存自然不够用了。
再看连续拼接。卷积通常喜欢连续存储的张量,很多框架也默认这样做。问题在于 DenseNet 的输入来自不同层,原本分散在内存里,想喂给卷积就得先拼成连续块。于是又多了一轮复制。为了让卷积舒服一点,显存先累到喘不过气,这事非常工程味,也非常真实。
从计算图角度看,朴素实现的问题更直白:每层都在生成新的拼接结果、新的 BN 结果、新的梯度缓存,前向和反向都在“存存存”。DenseNet 的输出特征本来只需要线性增长,但这些中间张量却会让显存占用接近二次增长。换句话说,不是 DenseNet 天生吃显存,而是实现方式太老实

巧解内存:共享内存与按需重计算策略

这篇报告最值钱的地方,不是提出了什么全新网络结构,而是把一个很朴素的工程判断做到了极致:既然拼接和归一化都很便宜,那就别傻乎乎地永久保存它们。需要时再重算,往往比一直占着显存更划算。
图3:DenseNet层前向传播的朴素实现与高效实现对比
图3:DenseNet 层前向传播的朴素实现与高效实现对比。实心框表示实际分配的张量,半透明框表示指针。高效实现把拼接、批归一化等中间结果放进共享临时缓冲区,而不是每层都重新申请内存。
本方法的核心是两个共享缓冲区。第一个缓冲区给特征拼接用,第二个缓冲区给批归一化输出用。前向传播时,各层把临时结果写进这两个共享区域;反向传播时,再按需把这些结果重算出来。这样做的逻辑很像“借用同一块案板切菜”,切完就清空,不要每切一刀都买一块新案板。
这里还有一个关键前提:重算的对象必须足够便宜。DenseNet 正好满足这一点。拼接只是复制张量,BN 也只是缩放和平移,远比卷积便宜。原论文给出的判断很明确:这些操作大概只占一次前向-反向过程的很小一部分时间,因此用少量时间换大量显存,非常划算。
除了前向激活,反向传播里的梯度张量也会吃内存。本文进一步把梯度存储也改成共享方式,避免反向阶段继续“堆垃圾”。最后的效果很干脆:真正需要长期保留的,只剩卷积特征图和参数,中间的拼接、归一化和梯度缓存都可以共享或者重算。
    前向传播:
    1. 将上一层输出收集到共享拼接缓冲区
    2. 将拼接结果写入共享BN缓冲区
    3. 执行卷积,得到当前层输出
    4. 释放临时中间结果的独立存储
    
    反向传播:
    1. 按需重算拼接结果
    2. 按需重算BN结果
    3. 用重算后的中间值完成梯度计算
    4. 梯度也尽量写入共享存储

    线性存储:超越250层的DenseNet诞生

    方法好不好,不能只看概念,得看它到底把模型深度推到了哪一步。原始 DenseNet 实现里,层数很快就被显存卡死;而共享内存和按需重算之后,情况就变了:深度上限被显著抬高,原本“想都不敢想”的大模型终于能跑起来。
    图4:GPU内存消耗随网络深度变化的曲线
    图4:GPU 内存消耗随网络深度变化的曲线。随着层数增加,朴素实现的显存占用迅速上升;高效实现则显著压低了中间特征图的存储成本,使更深的 DenseNet 成为可能。
    原论文给出的数字很有说服力。以 bottleneck DenseNet-BC、每层增加 12 个特征为例,朴素实现的 160 层模型已经非常吃显存;而采用全部共享策略后,在同样预算下可以训练更深的网络。更关键的是,这种节省不是“牺牲模型能力换来的”,而是把那些本来就不该长期保存的中间量清理掉了。
    这里有个容易被忽略的细节:参数量本身仍然会随深度增长,DenseNet 的参数并不是线性的。但这和中间特征图的二次膨胀不是一回事。报告强调,真正把显存拖垮的主要是 feature maps,而不是参数本身。把这个锅甩清楚之后,工程优化的方向也就明确了。
    这类工作最让人服气的地方在于:它没有试图“发明一种更会算的网络”,而是把已有网络的内存管理打磨到能真正训练超深模型。对很多研究来说,能不能跑起来,往往比论文里写得多漂亮更重要。

    效率与深度兼得:15%时间换数倍深度

    工程优化最怕一种情况:显存省了,速度崩了。那就不是优化,是搬家。本论文的结果比较稳,因为它没有掉进这个坑。
    图5:计算时间对比
    图5:计算时间对比。共享梯度存储几乎不带来额外时间成本,而共享 BN 和拼接存储会因为反向阶段的重算带来约 15% 到 20% 的时间开销。
    这组结果很关键,因为它给出了一个很现实的交换比:大幅降低显存,代价只是小幅增加训练时间。对训练大模型的人来说,这通常是可以接受的。毕竟显存不够的时候,不是慢一点的问题,而是根本训不动的问题。
    从方法设计上看,这个折中也很合理。卷积是大头,拼接和 BN 是小头;把小头重算一遍,换来大头能继续堆深度,性价比很高。换成别的任务未必都成立,但在 DenseNet 这里,作者的判断是对的,而且实验也把这点坐实了。
    更值得注意的是,PyTorch 的自动梯度机制本身就已经比早期 LuaTorch 更省内存,这说明这篇工作不仅是在“发明新技巧”,也在提醒大家:框架层面的实现细节,真的能决定一个模型能不能训练。这话听起来朴素,工程里却非常值钱。

    ImageNet新纪录:264层DenseNet的惊艳表现

    如果说前面的内容还偏工程优化,那这一部分就是结果兑现。作者把新的实现策略真正用到了 ImageNet 上,而且不是“跑通就算赢”,而是把网络深度和精度都往前推了一截。
    图6:ImageNet上的Top-1分类错误率对比
    图6:ImageNet 上的 Top-1 分类错误率对比。星号标记的模型如果没有高效实现,根本训练不出来。新的 264 层 DenseNet 在单裁剪测试下取得了 20.26% 的 Top-1 错误率。
    这里有两个值得记住的点。第一,作者不仅训练了更深的 DenseNet,还尝试了不同增长率和学习率策略,包括标准训练流程和余弦退火学习率。第二,结果说明 DenseNet 的性能仍然会随着层数增加而继续改善,并没有在 161 层附近“撞墙”。这对当时的结论很重要:深度并没有失效,失效的是显存管理
    264 层 DenseNet 的意义不只是“更深”,而是证明了一个判断:只要把训练时的内存瓶颈解开,DenseNet 这类强特征复用架构还有继续变强的空间。对于研究者来说,这类结果比单纯刷一个小幅指标更有价值,因为它直接告诉大家,某些模型不是不行,只是以前装不下
    但也不能把这篇工作神化。它更像一把把门打开的钥匙,而不是重新造了一栋房子。它解决的是 DenseNet 训练中的工程瓶颈,带来的收益非常实在,但它并没有改变 DenseNet 的基本结构逻辑,也没有证明所有网络都能用同样方式无痛扩深。这个边界要看清楚,才算读得明白。

    龙迷三问

    下面是龙哥对于大家可能的一些问题的解答:

    这篇论文到底解决了什么问题?它解决的是 DenseNet 训练时的显存瓶颈。模型本身很高效,但朴素实现会把大量中间特征图和梯度都存下来,导致显存随着深度快速增长。本文通过共享内存和按需重计算,把中间特征图的存储从接近二次增长压到线性级别。

    BN、concat 这些词为什么总被提到?BN 是 Batch Normalization,中文叫批归一化;concat 是 concatenation,中文是拼接。DenseNet 里这些操作会频繁产生临时张量,朴素实现会把它们一直留在显存里。本文的关键就是把这些临时量改成共享缓冲区,反向传播时再重算。

    这类方法为什么只多花一点时间,却能省很多显存?因为被重算的部分本来就很便宜:拼接主要是复制,BN 主要是缩放和平移,远比卷积轻。拿这些小开销换掉大量中间存储,整体就很划算,所以作者观察到的时间增幅只有约 15% 到 20%,但深度和可训练规模却明显上去了。

    如果你还有哪些想要了解的,欢迎在评论区留言或者讨论~

    龙哥点评

    论文创新性分数:★★★☆☆

    创新点不在模型结构本身,而在训练实现的系统优化;思路不花哨,但抓住了 DenseNet 的真实痛点,属于“工程上很值钱”的那类工作。

    实验合理度:★★★★☆

    对比了朴素实现、共享梯度、完整共享策略,逻辑清楚;同时给出内存和时间两条曲线,能把“省显存是否值得”说透。

    学术研究价值:★★★★☆

    它证明了实现细节可以显著改变可训练深度,这对后续的模型压缩、重计算、显存优化都有启发,研究价值很实在。

    稳定性:★★★★☆

    方法本质是内存管理优化,不改网络语义,稳定性比“魔改结构”更好;但依赖框架实现和算子特性,跨平台仍要做验证。

    适应性以及泛化能力:★★★☆☆

    对 DenseNet 这类有大量可重算中间量的结构很适合,但并不是所有网络都同样适用;如果中间步骤本身很贵,重算策略就未必划算。

    硬件需求及成本:★★★★☆

    训练成本明显下降,单卡/单机可训练更深模型;代价是增加少量计算时间,但总体对硬件门槛的缓解很明显。

    复现难度:★★★☆☆

    思路不复杂,但要真正做到高效,得理解框架的内存分配、共享缓冲和反向重算;不是抄几行代码就能完全复刻。

    产品化成熟度:★★★☆☆

    对训练侧非常有用,尤其是显存紧张的大模型训练;但它更像底层优化方案,不是直接面向业务端的“开箱即用”产品功能。

    可能的问题:方法很务实,但论文更像技术报告而不是完整系统论文;对于不同框架和更复杂算子,重算收益是否始终稳定,还需要更广泛验证。


    主要参考文献

    [1] Pleiss G, Chen D, Huang G, et al. Memory-Efficient Implementation of DenseNets. arXiv:1707.06990, 2017.
    [2] Huang G, Liu Z, van der Maaten L, Weinberger K Q. Densely Connected Convolutional Networks. CVPR, 2017.
    [3] He K, Zhang X, Ren S, Sun J. Identity Mappings in Deep Residual Networks. ECCV, 2016.
    [4] Chen T, Xu B, Zhang C, Guestrin C. Training Deep Nets with Sublinear Memory Cost. arXiv:1604.06174, 2016.

    end
    欢迎加入龙哥读论文粉丝群,扫描下方二维码或者添加龙哥助手微信号加群:kangjinlonghelper。一定要备注:研究方向+地点+学校/公司+昵称(如 图像处理+上海+清华+龙哥),根据格式备注,可更快被通过且邀请进群。
    显存不够、论文看不懂、代码复现卡住?进群一起拆,少走弯路。
    wechat_helper dianzan
    转发文章 微博 X LinkedIn Facebook
    龙哥读论文 · PaperDaily

    本文基于龙哥读论文 PaperDaily 数据库整理,结合论文原文与工程视角进行解读。