在3D生成任务里,显存占用比2D图像生成更容易失控。一个中等规模的NeRF或3D扩散模型,只是前向渲染就会产生大量中间特征,如果再叠加多视角一致性约束,显存很快见底。解决这个问题不能只靠降低batch size,因为过小的batch会让梯度估计不稳定。激活检查点与梯度累积是两种互补的显存优化手段,一个用重计算换激活存储,一个用小batch累加模拟大batch训练。本文结合3D生成模型的特点,说明两者的实现方式、适用边界和组合策略。

一、3D生成训练的显存压力从何而来
3D生成模型与2D生成模型最大的不同在于表示方式。NeRF系列通过对每个像素发射光线并在光线上采样大量空间点,再把这些点送入MLP网络预测颜色和密度;3D Gaussian Splatting则维护数十万甚至上百万个高斯原语,渲染时需要逐点投影和叠加;3D扩散模型直接在体素、点云或三平面隐空间上去噪,特征图天然带有三维尺寸。这些过程都会在计算图中保留数量庞大的中间张量,激活显存往往比参数本身高出数倍。
以NeRF训练为例,一个像素可能对应64到128个采样点,每个采样点要经过多层MLP,多层之间的激活都会被autograd保存下来。如果一次迭代渲染4096条光线,激活张量数量很容易达到几百MB甚至几GB,这还没有算上多视角一致性损失、深度监督等额外分支。对于3D扩散模型,一个batch可能包含8到16个体素网格,每个网格的隐藏层特征尺寸是64×64×64,同样会产生巨大激活。
除了激活,优化器状态也占用大量显存。Adam系列优化器需要为每个参数保存一阶动量和二阶动量,如果模型参数为500MB,优化器状态就接近1GB。降低batch size可以线性降低激活显存,但过小的batch会破坏批归一化统计量,也会让梯度下降方向噪声变大。3D生成训练通常只能使用很小的batch size,此时单纯降batch已经不够,需要更细粒度的显存管理技术。
二、激活检查点的机制与代码实现
激活检查点的核心思想是:前向传播时不保存中间激活,只记录分段边界上的输入;反向传播到达该段时,重新执行这一段的前向计算,得到中间激活后再继续反向。这样做牺牲了额外的前向计算时间,换取激活显存的大幅下降。PyTorch提供了现成的接口torch.utils.checkpoint.checkpoint,可以直接包裹需要检查点的函数或模块。
对于3D生成模型,并非所有层都值得做激活检查点。体渲染里的MLP解码器、扩散模型中的注意力块、三平面特征解码器等计算密集且激活较大的模块收益最高。多层感知机前向计算成本相对较低,重算的时间开销可以接受;而卷积层、归一化层等如果激活本身不大,强行检查点反而增加调度开销。实际使用建议先通过显存分析工具定位激活热点,再对热点模块做分段。
import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint
class NeRFDecoder(nn.Module):
def __init__(self, dim=256):
super().__init__()
self.fc1 = nn.Linear(dim, dim)
self.fc2 = nn.Linear(dim, dim)
self.fc3 = nn.Linear(dim, dim)
self.act = nn.ReLU()
def _forward(self, x):
h = self.act(self.fc1(x))
h = self.act(self.fc2(h))
return self.fc3(h)
def forward(self, x):
# 只对计算量较大的连续MLP做检查点
return checkpoint(self._forward, x, use_reentrant=False)
上面的代码中,use_reentrant=False是PyTorch较新版本推荐的模式。它避免了可重入检查点对autograd图的一些限制,能更好地处理复杂分支和多次调用,但要求被包裹的函数不包含原地修改输入等特殊操作。对于大多数3D生成中的标准模块,这个参数可以放心开启。
激活检查点带来的时间开销通常表现为训练速度下降,而不是显存出现波动。因为重计算只发生在反向阶段,前向仍然只计算一次。如果3D生成任务中数据加载、光线采样或渲染预处理已经占用了大量时间,检查点带来的额外计算可能被掩盖,实际训练速度下降幅度有限。这个特性使得它在3D生成里比在纯图像分类里更有吸引力。
三、梯度累积如何在不增加显存的情况下扩大batch规模
梯度累积解决的是另一个问题:当物理显存只允许很小的batch size时,如何获得大batch训练的稳定梯度。它的做法是连续处理多个小batch,每次反向传播后不更新参数,而是把梯度累加在param.grad里,达到预设步数后再执行一次optimizer.step()。从梯度期望上看,这等价于把这些小batch合成一个大batch,因为总梯度等于各小批量损失梯度之和。
实现时需要注意损失归一化。如果每个小batch的损失直接backward(),累加后的梯度相当于大batch总损失的梯度;想要保持每个样本的平均梯度与单步训练一致,需要把损失除以累积步数。下面是一个典型的训练循环:
accumulation_steps = 4
optimizer.zero_grad()
for step, batch in enumerate(dataloader):
inputs, targets = batch
outputs = model(inputs)
loss = criterion(outputs, targets)
# 关键:除以累积步数,保持更新幅度一致
loss = loss / accumulation_steps
loss.backward()
if (step + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
在3D生成任务里,梯度累积的优势很明显。NeRF每次迭代采样的光线本身就是随机子集,把多次采样的梯度累加起来等效于一次采样更多光线;3D扩散模型在不同时间步上的去噪样本也可以自然累积。由于3D生成网络大多使用LayerNorm或GroupNorm,不像2D分类网络那样依赖BatchNorm的batch统计量,因此小batch累积带来的统计偏差要小得多。不过如果某些模块确实使用了批归一化,建议在累积阶段使用滑动统计量,并适当调整动量系数。
梯度累积与激活检查点互不冲突。前者只改变优化器更新频率,不减少单步前向的激活存储;后者只降低单步激活占用,不改变有效batch大小。两者叠加后,可以在显存受限的情况下先通过检查点把单步峰值压下来,再通过梯度累积把有效batch规模提上去。
四、在3D生成项目中组合使用的策略与调参建议
实际项目里,盲目套用两个技术往往达不到预期。更可靠的做法是先用torch.cuda.memory_summary()或PyTorch Profiler观察一次训练迭代的显存分布,区分参数、优化器状态、梯度和激活分别占了多少。如果激活占比超过一半,优先考虑激活检查点;如果batch size已经小到1但还希望提升训练稳定性,再叠加梯度累积。
以典型的NeRF训练为例,假设单卡显存为24GB,当前batch size为4096条光线,一次迭代激活约6GB,总占用12GB。对MLP解码器做检查点后,激活可能降到3.5GB,总占用降到9GB左右。此时如果想让有效batch size达到16384条光线,可以设置accumulation_steps=4,显存不会增加,但梯度估计会更稳定。训练时间会增加约15%到25%,主要来自重计算。对于3D Gaussian Splatting训练,激活占用通常比NeRF低,但模型参数和优化器状态更大,这时检查点的收益相对有限,梯度累积反而是更主要的优化方向。
检查点分段的大小也会影响效果。分段太粗,重计算范围大,时间开销高;分段太细,每次边界保存和恢复的额外开销增加。一般建议以注意力块、解码器层或MLP连续三层为一个检查点单元。梯度累积步数不宜设置过大,过大会让优化器更新频率过低,训练曲线出现台阶式下降,也可能放大延迟更新带来的不稳定性。常见范围是2到8步,具体需要结合学习率和数据分布做小幅搜索。
另外,激活检查点和梯度累积都可以与混合精度训练搭配。混合精度能进一步减少参数、梯度和激活的显存占用,并且3D生成模型中的矩阵乘法在FP16下通常有明显加速。但需要注意检查点重计算时保持精度设置一致,避免前向和重计算使用不同精度导致梯度不一致。综合使用这些手段,单卡训练中等规模的3D生成模型是可行的。