视频生成模型的训练成本极高,一次完整的训练动辄跑上几天甚至几周。在这种长周期任务里,最让人崩溃的事情莫过于半夜收到报警,打开日志发现loss已经变成了NaN,而且从出现NaN的那个step开始,后面的几十万个样本全部白跑。梯度爆炸和NaN问题在视频生成领域尤其常见,因为视频任务的时间维度让中间激活值的规模远超图像模型,且扩散模型的噪声预测目标本身数值范围就很宽。本文围绕检测、捕获、恢复三个环节,给出一套可以直接落地的工程方案。

一、为什么视频生成模型特别容易出现梯度爆炸
视频生成模型通常在扩散模型框架下工作,网络需要在多个时间步上预测噪声。时间步t接近0时,输入信噪比很高,模型的激活值容易偏大;时间步接近T时,输入几乎是纯噪声,梯度方差极大。这种跨时间步的数值分布差异,使得单一的一组归一化参数很难同时适配所有情况,稍有不慎就会在某些batch上产生异常大的梯度。
另一个因素是视频数据的时空注意力模块。3D注意力或者时空分离注意力的softmax在序列长度很长时,中间结果的数值精度要求更高。如果使用了AMP混合精度训练,FP16的表示范围只有FP32的几万分之一,一次上溢或者下溢就足以产生inf,随后inf参与运算就变成NaN,再经过反向传播污染所有梯度。
p还有一个隐蔽的来源是VAE编码器。很多视频生成模型先经过VAE把视频压缩到潜空间再训练,如果VAE的编码输出没有做数值截断,个别极端样本会编码出数值非常大的潜变量,直接把loss算成inf。排查时建议先对数据管线做一次全量扫描,确认进入模型的张量都在合理范围内。二、梯度爆炸的检测:在崩溃之前捕捉信号
NaN一旦出现就晚了,更有效的做法是监控梯度范数的变化趋势。正常的训练过程中,总梯度范数会波动但基本维持在一个数量级内;如果连续多个step的梯度范数呈指数级上升,往往就是爆炸的前兆。实现上可以在反向传播之后、优化器更新之前插入检查逻辑。
import torch
def check_gradients(model, max_norm=1e4):
total_norm = 0.0
for p in model.parameters():
if p.grad is not None:
param_norm = p.grad.data.norm(2)
total_norm += param_norm.item() ** 2
total_norm = total_norm ** 0.5
if not torch.isfinite(torch.tensor(total_norm)):
print("警告:梯度已包含NaN或Inf")
return False, total_norm
if total_norm > max_norm:
print(f"警告:梯度范数异常 {total_norm:.2e}")
return True, total_norm上面的函数返回梯度的总范数,建议每一步都记录下来并写入日志。实践中有价值的指标有两个:一是当前范数值本身,二是范数与滑动平均值的比值。当某一步的范数突然变成滑动平均的十倍以上,即使还没到NaN,也应该触发预警。可以把这个比值作为动态阈值,比固定阈值更适应不同模型的数值尺度。
对于使用梯度累积的场景要注意,检测应该放在所有累积步的反向传播完成之后进行,否则每次micro-batch的局部梯度范数会低估真实值。另外如果使用DeepSpeed或FSDP等分布式框架,梯度分片在不同设备上,需要用框架提供的接口聚合全局范数,单机单卡的写法在多卡环境下会漏检。
三、NaN值捕获:让训练循环具备自我感知能力
检测只解决发现问题,捕获则要精确定位NaN产生的位置。最直接的手段是在loss计算后立即判断,一旦发现非有限值就跳过这一步的参数更新,并记录导致问题的样本索引,方便后续排查数据问题。
loss = compute_loss(model, video_batch, timesteps)
if not torch.isfinite(loss):
skipped_batches.append(batch_indices)
logger.warning(f"loss非有限值,跳过step {global_step}")
optimizer.zero_grad(set_to_none=True)
continue
loss.backward()
finite, grad_norm = check_gradients(model)
if not finite:
optimizer.zero_grad(set_to_none=True)
continue
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()注意zero_grad(set_to_none=True)这个细节,把梯度直接置为None而不是清零,可以避免已经污染的梯度在下一轮被意外使用。跳过更新虽然会轻微拖慢收敛,但比起NaN扩散到全部权重,这点代价完全值得。
AMP场景下还需要处理GradScaler的行为。当发生溢出时,GradScaler会自动跳过该次更新并缩放scale因子,这一点很多人不了解,以为是自己的代码出了bug。可以通过scaler.get_scale()观察scale值的变化,如果scale频繁下降,说明模型经常处于溢出边缘,此时应该考虑调小初始scale或者对loss做数值平滑处理。
scaler = torch.cuda.amp.GradScaler(init_scale=2**14) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) # 溢出时会自动跳过 scaler.update()
四、崩溃恢复机制:从checkpoint自动续训
即使有完善的检测,也不能保证百分之百不崩溃,硬件的偶发错误、坏数据、学习率调度异常都可能引发问题。一套可靠的恢复机制需要包含三个部分:周期性保存带优化器状态的checkpoint、崩溃后自动回退到最近的有效checkpoint、以及可选的单步调试模式。
import os, torch, traceback
def save_checkpoint(model, optimizer, step, path="ckpt"):
os.makedirs(path, exist_ok=True)
state = {
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"step": step,
}
tmp = f"{path}/ckpt_{step}.tmp"
torch.save(state, tmp)
os.replace(tmp, f"{path}/ckpt_{step}.pt") # 原子写入防止保存中断导致损坏
def train_with_recovery(train_fn, model, optimizer, start_step, max_retries=3):
retries = 0
step = start_step
while retries < max_retries:
try:
step = train_fn(model, optimizer, step)
break
except (RuntimeError, FloatingPointError) as e:
traceback.print_exc()
retries += 1
ckpt = load_latest_checkpoint(path="ckpt")
model.load_state_dict(ckpt["model"])
optimizer.load_state_dict(ckpt["optimizer"])
step = ckpt["step"]
# 回退后降低学习率,减少再次崩溃概率
for g in optimizer.param_groups:
g["lr"] *= 0.5
return step原子写入这一点非常关键。保存checkpoint时如果进程被kill,会留下一个写了一半的损坏文件,恢复时加载失败等于雪上加霜。先写临时文件再用os.replace原子替换,可以保证磁盘上任何时刻都存在完整可用的checkpoint。同时建议保留最近三到五个checkpoint并循环覆盖,而不是只留最新一份。
恢复时降低学习率是一个实用技巧。梯度爆炸往往和当前学习率下模型处于不稳定区域有关,原学习率直接续训大概率会在同一位置再次崩溃。每次恢复将学习率乘以0.5,几次之后如果仍然崩溃,说明问题可能出在数据或者模型结构上,需要转入人工排查阶段,把出问题的batch固定下来做最小复现。
五、预防优于治疗:训练前的稳定性检查清单
最后总结一份上线训练前的检查清单。第一,对数据管线做数值审计,确认输入张量经过标准化后的最小最大值在合理区间,特别是VAE编码后的潜变量。第二,检查所有自定义loss中是否存在log、div、sqrt等可能产生inf的运算,必要时加上epsilon平滑。第三,attention计算中优先使用PyTorch内置的scaled_dot_product_attention,它内部对大序列做了数值稳定的处理。第四,warmup阶段使用线性预热而不是固定学习率,前几千步的梯度噪声最大,小学习率能显著降低早期崩溃概率。
把检测、捕获、恢复三层机制叠加起来,配合事前的稳定性检查,视频生成模型的长周期训练基本可以达到无人值守的可靠程度。建议从项目初期就把这些逻辑写进训练框架,而不是等第一次崩溃之后再补,那时损失的算力和时间已经无法挽回。