MIL-NCE分布式训练全解析:从PyTorch实现到HowTo100M高效训练
2026/9/14 20:31:54 网站建设 项目流程

简介:面向大规模视频与多模态表征学习场景的 PyTorch 分布式训练示例代码包,聚焦 MIL-NCE 预训练方法在 HowTo100M 数据集上的实现。压缩包共 22 个文件,以 13 个 Python 脚本为主,涵盖模型定义、数据加载、损失计算、分布式训练入口与下游任务评估;同时附有 CSV 数据索引、环境说明、README 与开源许可证,整体体积 22.02MB,结构轻量便于研读。已有 184 人学习,适合想通过实际代码掌握 PyTorch GPU 分布式训练(如 DistributedDataParallel)、多进程数据采样和大规模视频特征提取的开发者。通过运行和改造其中脚本,可以直观理解多实例学习与噪声对比估计的组合方式,也能复用其数据加载与评估流程,快速迁移到个人视频理解或跨模态检索项目中。代码目录按数据、脚本、模型、工具分层,注释清晰,可作为入门分布式训练的最小可运行样例。

1. MIL-NCE 靠负样本吃饭,但 HowTo100M 的单卡 batch 喂不饱它

MIL-NCE(Multiple Instance Noise-Contrastive Estimation)是 DeepMind 在 HowTo100M 上做视频-文本联合嵌入训练时提出的一类对比学习损失。它的正样本不是一条字幕,而是一“包”相邻字幕:因为百万级教学视频的 ASR 字幕与画面没有逐句对齐,训练时只能容忍“这段时间里至少有一句沾边”。真正决定效果的是负样本,而负样本全部来自当前 batch 内的其他句子。global batch 越大,负样本越丰富,但视频特征与句子特征算相似度时生成的 logits 矩阵也会迅速撑爆单卡显存。所以标题里的 GPU 分布式在 MIL-NCE 场景下不是可选项,而是凑出有效负样本的前提。下面沿着损失函数在 PyTorch 里的矩阵写法、HowTo100M 数据怎么切成 rank、torchrun 怎么启动、出问题时按什么顺序排查这条线走完,最后一章给一个能直接把训练周期砍半的缓存技巧。

2. MIL-NCE 损失函数在 PyTorch 里的矩阵实现:bag 收集与 chunked 负样本

2.1 为什么正样本是一包句子,而不是一句

标准 InfoNCE 的每一对正样本是明确对齐的,比如图像增强对、视频片段与对应描述。但 HowTo100M 的 ASR 输出带有识别错误和时间戳漂移,同一段画面经常对应前后多个句子,人工细粒度对齐又不现实。MIL-NCE 的解法是给每个视频片段收集一个时间窗内的若干条字幕作为候选正样本,只要其中任意一句与该片段匹配就算正样例;负样本仍然是 batch 内其他视频、其他时间窗的句子。这就是 Multiple Instance Learning 叠加 NCE 的核心:正样本是一个包,而不是一条线。

这个设计带来的直接收益是模型不再被单句错误标签带偏。用 logsumexp 对包内候选句做“软最大”聚合,梯度会分配给与当前片段最相似的几句,既保留了对齐信息,又不像 max 那样掐死其他候选句的梯度。代价是每一步都要针对一个 [B, N] 的相似度矩阵计算,B 是当前 batch 的视频数,N 是这一步见到的全部句子数,通常 N 远大于 B。显存预算要先算清这笔账,再看 batch 怎么分配。

2.2 一个可直接搬进训练脚本的 mil_nce_logits 实现

下面这段是常用的实现方式,形状注释已经写在代码里。拿到任何 MIL-NCE 相关代码包后,先拿这个函数去对 loss 曲线,确认前后实现语义一致,再往下调数据管线:

import torch import torch.nn.functional as F def mil_nce_logits(vid_emb, cap_emb, bag_idx, temperature=0.07): # vid_emb: [B, D] 视频嵌入,未归一化 # cap_emb: [N, D] 本 step 内全部句子的嵌入 # bag_idx: [B, A] 视频 i 的 A 个候选正文本在 cap_emb 中的行号 vid_emb = F.normalize(vid_emb, dim=-1) cap_emb = F.normalize(cap_emb, dim=-1) B = vid_emb.size(0) N = cap_emb.size(0) logits = vid_emb @ cap_emb.t() # [B, N] logits = logits / temperature pos = logits.gather(1, bag_idx) # [B, A] log_pos = torch.logsumexp(pos, dim=1) # 正样本包的聚合 log_denom = torch.full((B,), -float('inf'), device=vid_emb.device) chunk = 1024 # 控制瞬时显存峰值 for i in range(0, N, chunk): part = vid_emb @ cap_emb[i:i+chunk].t() / temperature log_denom = torch.logaddexp( log_denom, part.logsumexp(dim=1)) return log_denom - log_pos # [B],外部自行 mean()

实现里有三个关键点。第一,两个嵌入都做 L2 归一化,点积结果就是余弦相似度,再除以温度放大差异;温度越小,正负样本之间的分数差越敏感。第二,gather把每个视频对应的 A 条句子相似度取出来,logsumexp完成包内聚合,它保留了对多个候选句的梯度,数值上比逐项exp再相加更稳。第三,分母必须包含全部句子,正样本也在其中,这样最终形式才是“正样本包占全部句子的概率比例取负对数”,等价于一个带噪声标签的多分类 softmax。

负样本部分用chunk循环是给大 batch 预留的。全量 [B, N] 矩阵在 B=256、N=4096 时还不到显存瓶颈,但反向传播时矩阵乘的中间梯度会明显放大占用;逐块累加logaddexp后,瞬时峰值被限制在 [B, 1024],B 再涨也不怕。如果你的显存余量充足,也可以直接用全量 logits 一次算完,结果一致。

2.3 bag 宽度、温度和 per-GPU batch 的相互制约

调参时最容易陷入的误区是只盯视频塔的 batch,忽略句子塔的规模。我一般先看三个数字:温度、bag 宽度 A、per-GPU batch。它们之间的牵制关系如下:

参数常用区间调大后的效果主要风险
temperature0.05 ~ 0.1对难负样本更敏感,收敛更尖锐低于 0.03 后梯度容易爆炸,fp16 下尤其明显
bag 宽度 A8 ~ 16 句正样本召回率提升过大时包里混入无关句子,标签噪声回升
per-GPU batch16 ~ 32 视频直接增加句子侧负样本数量句子嵌入矩阵与视频特征争抢显存
global batch128 ~ 512负样本多样性显著提升学习率需要同步放大,warmup 拉长

HowTo100M 的 ASR 句子都比较短,句子塔的内存压力小于视频特征。真正要留意的是 bag 索引的构建:如果同一视频相邻时间窗的句子被同时选进正样本包和负样本库,会出现“负样本撞车”。常见做法是在构建 bag 时直接把与当前时间窗有重叠的句子从负样本中剔除,实现上只需要在生成 bag_idx 时维护一个 mask。这个细节对 loss 绝对值影响不大,但会直接影响 R@10 这类检索指标曲线。

3. HowTo100M 的分布式数据管线与 PyTorch torchrun 启动

3.1 不要把 mp4 直接送进 DataLoader

HowTo100M 的发布形态是 CSV,每行包含 YouTube 视频 ID、起止时间和 ASR 字幕文本,视频本体需要另想办法获取,总量累计超过 15 万小时。如果直接把 mp4 丢进训练 DataLoader,每个 epoch 都要重复解码、重复抽帧,CPU 很快跑满,GPU 利用率掉到个位数。分布式场景下,8 张卡等 1 个视频解码 worker 的情况非常常见。

业界的通常做法是把“视频解码”和“训练”彻底拆开。先用一个预训练视频编码器把所有视频按固定步长抽帧编码成特征序列,存成 .npy 或 WebDataset;训练阶段只读特征,不再碰原始视频。MIL-NCE 原始工作里的视频塔选择是 S3D-G,也可以换成 TSM 或 Tiny VideoNet,关键是离线抽特征时保持采样率一致,否则训练时视频长度 T 对不上。这一步的预处理命令大致长这样:

python preprocess_videos.py --csv /data/howto100m/train.csv \ --video_root /data/howto/videos \ --feat_out /data/howto/feat_2s \ --encoder s3d --sample_rate 0.5 --batch 8 --workers 4

其中--sample_rate 0.5表示每 2 秒保存一个特征向量,输出形状是 [T, 512];--workers 4是给解码库用的并发数。视频文件损坏率在 HowTo100M 里不低,脚本里要记录跳过路径而不是直接中断,否则跑一整晚发现卡在第 3 万条坏视频上。

3.2 最小可运行的 DDP 训练骨架

训练侧的 DataLoader 只读 .npy 特征和预解析好的句子 id。用 PyTorch 自带的分布式数据并行时,最容易被忽略的是DistributedSampler:每个进程不能自己写随机采样,否则多卡之间数据重叠,负样本跨卡 concat 后会出现重复,等于负样本数量虚增。下面是一个骨架:

import os import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader, DistributedSampler def train(): rank = int(os.environ["LOCAL_RANK"]) world_size = int(os.environ["WORLD_SIZE"]) dist.init_process_group(backend="nccl") torch.cuda.set_device(rank) ds = HowToFeatureDataset(feat_dir="/data/howto/feat_2s", meta="/data/howto/train.csv") sampler = DistributedSampler(ds, num_replicas=world_size, rank=rank, shuffle=True) loader = DataLoader(ds, batch_size=16, sampler=sampler, num_workers=2, pin_memory=True, drop_last=True) model = JointEmbedding(vocab_size=30000, d_model=512).cuda(rank) model = DDP(model, device_ids=[rank]) opt = torch.optim.AdamW(model.parameters(), lr=1e-3) for epoch in range(20): sampler.set_epoch(epoch) for step, batch in enumerate(loader): vid = batch["vid"].cuda(rank) cap = batch["cap"].cuda(rank) bag = batch["bag_idx"].cuda(rank) loss = mil_nce_logits(vid.mean(1), cap, bag, temperature=0.07).mean() opt.zero_grad() loss.backward() opt.step()

set_epoch(epoch)必须在每个 epoch 开始时被调用,否则 DistributedSampler 内部随机种子不变,所有 epoch 的数据切片顺序完全相同。drop_last=True保证每个 rank 迭代步数一致,否则梯度同步在最后一个不完整 batch 上会挂住。使用 torchrun 启动时不需要手动传 rank,它会把 RANK、LOCAL_RANK、WORLD_SIZE 写进环境变量:

torchrun --nproc_per_node=4 --nnodes=1 --master_port=29500 \ train_mil_nce_ddp.py

多节点时,每台机器执行同一条命令,额外传--nnodes=2 --master_addr=主节点IP --master_port=29500,rank 由 torchrun 自动分配。第一次跑建议先--nnodes=1 --nproc_per_node=2验证脚本本身没有单卡依赖,再上多机。

3.3 全局 batch、梯度累积与学习率怎么对齐

MIL-NCE 的负样本收益来自全局 batch,但显存限制在 per-GPU batch,所以实际工程里几乎必然用到梯度累积。参数对照关系可以参考:

每卡 batch卡数累积步数全局 batch建议学习率
8441284e-4
16421284e-4
16845121e-3
32825121e-3

学习率按线性缩放是常见做法,基准是全局 batch 256 对应 8e-4,全局 batch 翻倍时学习率乘sqrt(2)或直接翻倍,具体要看 loss 曲线是否震荡。梯度累积实现时注意:累积多个 backward 后只调用一次optimizer.step(),DDP 的梯度同步发生在每次 backward 结束时,所以累积不会破坏梯度同步。如果模型里带 BN 层,建议在包装 DDP 之前调用torch.nn.SyncBatchNorm.convert_sync_batchnorm(model),让 BN 的均值方差跨卡同步,否则全局 batch 变大了 BN 统计量却还是单卡视角,负样本粒度和归一化粒度不一致。

4. MIL-NCE 分布式训练排错:NaN、梯度不同步、视频 worker 卡死

4.1 loss 在头几百步变 NaN,优先怀疑温度而不是权重衰减

MIL-NCE 训练里 NaN 很少直接出现在 loss 第一行,更多是跑到一半梯度爆炸,权重变成 NaN,后续 forward 全部跟着 NaN。检查顺序建议固定下来:先关 AMP 用 fp32 跑 200 步,确认是否还崩;第二步看温度系数,低于 0.05 的尝试提到 0.07;第三步检查 bag_idx 是否出现 0 之外的越界索引,gather越界在 CUDA 上通常不报错而是返回垃圾值。

温度排在第二位是因为它位于 logsumexp 的指数内部。温度缩小一倍,logits 就被放大一倍,半精度下exp(40)附近就开始触达 bf16/fp16 的表示上限。即使 forward 没溢出,反向传播的梯度也会因为指数放大而爆炸。fp32 下指数本身不会 NaN,但梯度一旦把 Adam 的状态撑爆,下一步 forward 必然 NaN。因此使用混合精度时,梯度裁剪和 GradScaler 要成对出现。可以在 backward 之后挂一个梯度范数检查:

def check_grad_norm(model, threshold=20.0): total = 0.0 for p in model.parameters(): if p.grad is not None: total += p.grad.detach().float().norm().item() ** 2 total **= 0.5 if total > threshold: print(f"grad norm {total:.2f} exceeds {threshold}")

这一步能第一时间看到梯度爆炸趋势,比等到 loss 变 NaN 再翻日志省时间。实际调参中,如果温度已经是 0.07 且 fp32 稳定,只是 fp16 下偶尔溢出,优先考虑增大 GradScaler 的init_scale,而不是继续降温度。

4.2 各 rank 的梯度范数应该基本一致,这是 DDP 健康的自检信号

DDP 在每次 backward 结束后会自动对梯度做 allreduce 平均,正常情况下各 rank 的平均梯度范数应当一致。如果发现 loss 曲线在每张卡上分叉,或者不同 rank 的指标差异很大,可以用下面的代码直接检查:

def debug_grad_sync(model, rank): g = model.module.video_fc.weight.grad.detach().clone() g_list = [torch.zeros_like(g) for _ in range(dist.get_world_size())] dist.all_gather(g_list, g) norms = [g_i.norm().item() for g_i in g_list] if rank == 0: print("grad norms:", norms) if max(norms) - min(norms) > 1e-4: print("data sharding mismatch")

判断时注意一个边界:如果开了梯度累积,每步累积的梯度是在各自 rank 上独立累加的,浮点运算顺序不同会导致范数有微小差异,阈值放宽到 1e-2 再判断。真正常见的“不同步”原因有两个:一是模型部分参数没有包进 DDP(比如某个 buffer 更新方式写错),二是 Dataset 内部自己做了随机采样而没有走 DistributedSampler。前者会让每张卡学到不同状态,后者会让同一条视频被不同 rank 重复看到,负样本统计失真。

4.3 视频 worker 静默崩溃,先查 fork 冲突再查解码容错

HowTo100M 源视频质量参差,有的只有 360p,有的是 4K 高码率,解码库与 ffmpeg 的兼容性问题在 DataLoader 多进程里会被放大。最常见的一个坑是 PyAV/Decord 与 fork 启动的 worker 冲突,表现是第一个 epoch 跑不完整体挂起,nvidia-smi显示全部 GPU 利用率为 0 但 CPU 有进程占满。解决办法是显式指定 spawn 启动方式,或者给 DataLoader 加timeout=120,让超时报错而不是永久阻塞。

另一个坑是某个 worker 解码异常直接退出,主进程收不到任何异常文本,表现为 dataloader 迭代卡住。建议在预处理阶段加一个轻量校验脚本,用 ffprobe 读取每个文件时长和编码格式,把损坏文件记录到黑名单;训练端再配合 WebDataset 做容错,遇到坏 shard 自动跳过而不卡死。排查这类问题时,先把num_workers降到 0 跑 50 步,如果问题消失,基本可以确认是解码进程与训练主进程的资源争抢。

5. 把 HowTo100M 的 bag 索引固化成 .npz,是复现 MIL-NCE 最值的提速操作

5.1 缓存文件至少包含三样东西

每个 epoch 都重新解析字幕 CSV、重新构造 bag 索引,是训练脚本里最隐蔽的 CPU 开销。HowTo100M 的句子数量接近千万级,每次启动都做一遍时间窗匹配和 tokenize,既没有随机性收益,还拖慢数据加载。我一般会在预处理阶段把训练样本的 bag 索引、句子 id、视频特征路径一次性写好,保存为 .npz,训练时直接 load 进内存:

import numpy as np # 伪代码:跑一次后得到以下数组 np.savez("train_bag.npz", vid_paths=np.array(vid_paths, dtype=object), cap_ids=np.array(cap_ids, dtype=np.int64), bag_idx=np.array(bag_idx, dtype=np.int64), cap_lens=np.array(cap_lens, dtype=np.int64))

训练端拿到 bag_idx 后直接torch.from_numpy(bag_idx)转张量,完全跳过 CSV 解析和字符串处理。缓存固定的 bag 索引不会降低数据多样性,因为负样本多样性来自 batch 内句子组合,而不是 bag 每次重新生成。对复现实验的好处更直接:所有 rank 读同一个文件,数据顺序完全一致,实验差异只来自随机种子,不会因为某次重新解析导致列表顺序变化。

5.2 用 step time 验证缓存是否生效

验证缓存有没有真正解决瓶颈,不需要看完整训练曲线,记录两个指标足够:dataloader单次迭代耗时,和 GPUtorch.cuda.synchronize后的 step time。用torch.profiler或者简单的 time 戳都能统计。如果 step time 从 0.8 秒降到 0.3 秒,说明 CPU 端解析已经不是瓶颈;如果 step time 没降但 GPU 利用率提升了,说明之前是数据加载饥饿,现在可以继续调大num_workers或开 WebDataset 分片。跑通缓存后,把 batch=256、epoch=20 的 step time 压在 0.5 秒以内,20k 步左右就能看到验证集 R@10 曲线稳定抬头。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询