做过视频生成或者视频理解模型训练的人,多半遇到过这样一种尴尬:数据集准备了半个月,模型结构也调通了,结果一段几秒钟的高分辨率视频喂进去,显卡直接报CUDA out of memory。更扎心的是,明明换了张显存翻倍的卡,也只是把能处理的视频长度从16帧提到了24帧,离目标还差得远。问题的根源在于视频数据的显存开销不是线性增长的,中间特征图的体积会在网络层层传递中不断膨胀。这篇文章就来拆解视频任务显存不足的成因,并重点介绍分块处理和梯度检查点这两条最实用的解决路线。

先搞清楚:视频显存到底被什么吃掉了
要解决显存不足,得先知道显存花在哪里。以一个典型的视频生成网络为例,输入张量的形状通常是(B, C, T, H, W),即批次、通道、帧数、高、宽。假设输入是16帧的512x512视频,经过第一层卷积升到256通道,特征图体积就已经是原始RGB输入的近百倍。再往后走,注意力层会为每一对时空位置计算权重矩阵,显存开销还会进一步放大。
具体来说,显存消耗主要来自三块:第一是模型参数和优化器状态,这部分随模型规模固定增长,视频任务和图像任务差别不大;第二是激活值缓存,这是视频任务的重灾区,反向传播时需要保存每一层的中间输出,通道数乘以帧数再乘以空间分辨率,叠加起来非常可观;第三是临时缓冲区,比如卷积的im2col展开、注意力中间矩阵,这些在计算峰值时同样会挤占显存。
理解了这个结构,优化思路也就清晰了:参数量动不了,那就从激活值和临时缓冲区下手。分块处理对付的是单次前向的输入体积,梯度检查点对付的是激活值的全程缓存,两者方向不同,甚至可以叠加使用。
方案一:分块处理,把大视频切成小块喂进去
分块处理的核心思想很朴素:既然整段视频一次算不动,那就切开算。常见的切法有两种,一是按时间维切,把一段64帧的视频分成4个16帧的片段,逐片段过网络;二是按空间维切,把高分辨率帧切成重叠的小块,处理完再拼接。视频生成领域很多工作本质上都是时间维分块,比如先训练图像生成,再扩展到短视频片段,最后通过滑动窗口的方式生成任意长度的视频。
在PyTorch中实现时间维分块并不复杂,关键在于处理块与块之间的信息传递。如果各块完全独立,生成的视频会出现明显的帧间不连贯。一种常见做法是让每个分块额外接收前一个分块末尾的若干帧作为条件:
import torch
def chunked_forward(model, video, chunk_size=16, overlap=4):
"""
video: 形状 (B, C, T, H, W)
chunk_size: 每块帧数
overlap: 与前一块重叠的帧数,用于保证连贯性
"""
B, C, T, H, W = video.shape
outputs = []
prev_tail = None
for start in range(0, T, chunk_size):
end = min(start + chunk_size, T)
chunk = video[:, :, start:end]
if prev_tail is not None:
# 拼接上一块尾部帧作为上下文条件
chunk = torch.cat([prev_tail, chunk], dim=2)
out = model(chunk)
# 保留末尾几帧作为下一块的上下文
prev_tail = out[:, :, -overlap:]
outputs.append(out[:, :, overlap:])
return torch.cat(outputs, dim=2)
这段代码里有几个细节值得注意。overlap参数控制上下文重叠帧数,取值太小会导致块边界处画面跳变,太大又浪费计算量,实践中通常取4到8帧。另外,分块推理和分块训练的难度差异很大:推理时逐块前向即可,训练时如果各块独立计算损失,梯度无法跨块传播,时序一致性会打折,需要设计合适的条件传递机制。
空间维分块的思路类似,主要用于超高分辨率场景。切块时要留出足够的重叠区域避免拼接缝隙,最后对重叠区域做加权融合。它的代价是计算量增加,重叠部分会被重复计算,所以空间分块一般只在分辨率确实跑不动时才启用。
方案二:梯度检查点,用重计算换显存
梯度检查点是另一条思路。标准反向传播要求把所有中间激活值保留到反向阶段,而梯度检查点的做法是:前向时只保存少数几个检查点位置的输出,中间的激活值直接丢弃;反向传播走到某个检查点时,从上一个检查点重新做一段前向计算,把丢掉的激活值现场算回来。本质上是拿约30%的额外计算时间,换取激活值显存的大幅下降。
PyTorch对这个技术的支持已经非常完善,核心入口是torch.utils.checkpoint.checkpoint。对视频模型来说,最实用的方式是按层包裹,比如对Transformer的每个Block启用检查点:
import torch
import torch.utils.checkpoint as cp
class VideoTransformerBlock(torch.nn.Module):
def __init__(self, dim, heads):
super().__init__()
self.attn = torch.nn.MultiheadAttention(dim, heads, batch_first=True)
self.ffn = torch.nn.Sequential(
torch.nn.Linear(dim, dim * 4),
torch.nn.GELU(),
torch.nn.Linear(dim * 4, dim),
)
self.norm1 = torch.nn.LayerNorm(dim)
self.norm2 = torch.nn.LayerNorm(dim)
def forward(self, x, use_checkpoint=False):
if use_checkpoint and self.training:
# 训练时启用梯度检查点,推理时直接前向
return cp.checkpoint(self._forward_impl, x, use_reentrant=False)
return self._forward_impl(x)
def _forward_impl(self, x):
h = self.norm1(x)
attn_out, _ = self.attn(h, h, h)
x = x + attn_out
return x + self.ffn(self.norm2(x))
使用时有个非常容易踩的坑需要提醒:use_reentrant=False这个参数在较新版本的PyTorch中强烈建议显式指定。旧的reentrant模式要求被包裹的函数内部所有张量都参与梯度计算,一旦某层有不需要梯度的分支就会报错或静默出错;非重入模式则更宽容,对控制流的支持也更好。另一个坑是Dropout,重计算时随机数状态如果不一致,前后两次前向的结果会有细微差别,非重入模式已经处理了随机数种子的问题,这也是推荐它的原因之一。
实际测下来,对一个12层的视频Transformer启用逐层检查点,激活值显存通常能降到原来的三分之一左右,训练速度下降大约25%到35%。如果觉得全部层都开太慢,还可以采用隔层开检查点的折中方案,显存和速度可以按需调配。另外,Hugging Face的 accelerate库提供了gradient_checkpointing_enable接口,一行代码就能给支持的模型启用,省去手动包裹的工作。
两种方案怎么选,以及配套优化手段
分块处理和梯度检查点并不互斥,选择时可以按问题类型来判断。如果显存瓶颈出在单次前向的输入体积上,比如视频分辨率高、帧数多到连一层都过不去,那分块处理是唯一解,因为不管怎么省激活值,输入本身就把显存占满了。如果输入能过网络但激活值累积撑爆显存,那梯度检查点是首选,改动小、效果直接。多数实际项目是两者叠加:先分块控制单次输入规模,再在块内用梯度检查点压激活值。
除了这两大主力手段,还有几个配套优化值得一起做。混合精度训练用FP16或BF16存储激活值,显存直接减半,配合梯度缩放几乎不影响精度,是性价比最高的基础操作。梯度累积则在不增加显存的前提下变相扩大批次,把一个大批次的梯度分多次累计后再更新。如果显存还是紧张,可以把不活跃的参数或优化器状态临时转移到CPU,也就是CPU Offload,代价是PCIe数据搬运的时间开销。
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for step, (video, target) in enumerate(dataloader):
optimizer.zero_grad()
# 混合精度 + 梯度累积,累积4个小批次等效1个大批次
with autocast(dtype=torch.bfloat16):
loss = compute_loss(model, video, target) / 4
loss.backward()
if (step + 1) % 4 == 0:
# BF16下可以不用缩放,FP16则用scaler.step(optimizer)
optimizer.step()
optimizer.zero_grad()
最后给一个排查建议:动手优化前,先用torch.cuda.memory_summary()或者NVIDIA的Nsight Systems看一眼显存分布,确认瓶颈到底在参数、激活值还是临时缓冲区。盲目的优化往往事倍功半,比如显存明明是被注意力中间矩阵吃掉的,你去启用梯度检查点效果就有限,这时候更需要的是FlashAttention这类注意力优化。定位准确,再组合使用分块、检查点、混合精度这些手段,大部分视频任务都能在有限的显卡上跑起来。