InternVideo核心代码解析:从3D位置编码到Flash Attention的模型源码深读
【免费下载链接】InternVideo[ECCV2024] Video Foundation Models & Data for Multimodal Understanding项目地址: https://gitcode.com/OpenGVLab/InternVideo
本文带你深读 InternVideo 视频基础模型的核心源码,聚焦两大关键模块:3D 位置编码(让模型理解视频的时间+空间结构)与Flash Attention(让视频 Transformer 训练效率大幅提升)。即使你是刚接触 Transformer 的新手,也能跟着本文快速看懂 InternVideo2 视觉骨干的设计思路与工程技巧。
项目结构一览:核心代码在哪里?
InternVideo 仓库覆盖了 InternVideo1 → InternVideo3 的完整演进,而单模态视觉骨干的源码集中在 InternVideo2/single_modality/ 目录。与本文主题相关的三个关键文件:
| 文件 | 职责 |
|---|---|
| pos_embed.py | 1D / 2D / 3D 正弦余弦位置编码 |
| flash_attention_class.py | Flash Attention 封装模块 |
| internvideo2.py | InternVideo2 主模型(Attention、Block、PatchEmbed、forward 全流程) |
其他变体如 internvideo2_pretrain.py(预训练版)、internvideo2_ap.py(注意力池化版)、internvideo2_teacher.py(蒸馏教师版)共享同一套位置编码与注意力设计。
3D位置编码:给视频模型"时空感知"的关键
静态图像只需要 2D 位置信息,而视频还多了一个时间维度。InternVideo 的解法在 pos_embed.py#L9-L54 的get_3d_sincos_pos_embed函数中,思路非常优雅:
- 维度切分:把嵌入维度
embed_dim切成 4 份,其中 3/4 分给空间(2D),1/4 分给时间(1D); - 空间编码:调用
get_2d_sincos_pos_embed_from_grid分别对高度、宽度坐标做正弦余弦编码(pos_embed.py#L98-L110); - 时间编码:对帧索引做 1D 正弦余弦编码(pos_embed.py#L113-L131),底层公式就是经典的
sin/cos(pos × 10000^(-2i/d)); - 拼接对齐:时间编码沿空间位置重复、空间编码沿时间位置重复,最后按
[T, H, W]顺序拼接成完整的 3D 位置编码。
正弦余弦编码的最大优点是无需训练、天然支持外推,且解析式生成、零存储成本。
两种初始化模式:联合 vs 可分离
在 internvideo2.py#L440-L464 的init_pos_embed方法中,InternVideo2 支持两种位置编码方案:
- 联合模式(默认):用 3D sincos 编码一次性初始化一个
pos_embed参数,之后可学习; - 可分离模式(
sep_pos_embed=True):空间编码pos_embed_spatial与时间编码pos_embed_temporal各自独立学习,forward 时再做相加组合(internvideo2.py#L510-L524)。
可分离方案参数更少、分辨率泛化更强,而联合方案表达更灵活——两种模式在代码中通过一个布尔开关无缝切换,是很好的工程参考。
Flash Attention:视频 Transformer 提速的核心
视频帧数多、Token 数量巨大,标准注意力 O(N²) 的显存开销是训练视频模型的最大瓶颈。InternVideo 的解法在 flash_attention_class.py#L10-L71:一个轻量级FlashAttention封装类。
封装类的设计要点
- 输入格式:接受 QKV 打包张量
(B, S, 3, H, D),直接对接底层flash_attn_varlen_qkvpacked_func算子; - 变长序列支持:当存在
key_padding_mask时,先用unpad_input把有效 Token 压紧成连续序列,算完注意力再pad_input还原,避免为无效填充位浪费算力; - 硬件约束:强制要求
float16/bfloat16且运行在 CUDA 上(flash_attention_class.py#L36-L37),这正是 Flash Attention 高效的前提。
双路径切换:_naive_attn 与 _flash_attn
Attention 类 同时实现了两条前向路径:
| 路径 | 说明 |
|---|---|
_naive_attn | 标准 PyTorch 注意力,作为 CPU / 调试时的回退方案 |
_flash_attn | 走 Flash Attention 内核,配合 QK-Normalization 与融合 RMSNorm |
forward里一行代码完成切换(internvideo2.py#L217-L219),而模型构造函数通过断言保证use_flash_attn、use_fused_rmsnorm、use_fused_mlp三个开关必须一致(internvideo2.py#L370-L371)——因为融合算子之间在内存布局上相互依赖。
配套的Block类(internvideo2.py#L249-L299)还集成了 LayerScale、DropPath(随机深度)与梯度检查点,是大规模视频训练省显存的常用组合拳。
串起来看:InternVideo2 前向全流程
以 forward 方法 为主线,数据流清晰可见:
- 3D 切块:PatchEmbed 用一个
Conv3d(时间维步长tubelet_size、空间维步长patch_size)把视频一次性切成时空 Token,输出网格尺寸为(T, H, W); - 拼接 CLS Token:在序列头部插入可学习的分类 Token,随后加上位置编码;
- N 层 Transformer Block:RMSNorm → 注意力 → LayerScale → MLP,逐块堆叠(默认深度 40 层、1408 维嵌入);
- 注意力池化投影:
AttentionPoolingBlock用交叉注意力把数百个视频 Token 压缩成一个向量,投影到 768 维的 CLIP 空间(internvideo2.py#L109-L116),实现与文本塔的语义对齐; - 分类头输出:LayerNorm + Linear 得到最终类别分数。
新手学习路线建议 🧭
如果你想继续深入 InternVideo 源码,建议按以下顺序阅读:
- 先跑通 run_pretraining.py 与 run_finetuning.py,建立"数据 → 模型 → 损失"的全局认知;
- 精读 pos_embed.py 与 flash_attention_class.py,掌握本文两大核心模块;
- 对照 scripts/finetuning/ 下的训练脚本,理解各模型规模的超参差异;
- 关注 MODEL_ZOO.md 与 INSTALL.md,补齐权重下载与部署知识。
总结:InternVideo 的源码把"视频时空建模"与"大模型训练效率"两个难题,分别用可组合的 3D sincos 位置编码和双路径 Flash Attention 封装给出了解答。理解了这两个模块,你就掌握了读懂大多数视频 Transformer 骨干的核心钥匙 🔑。
【免费下载链接】InternVideo[ECCV2024] Video Foundation Models & Data for Multimodal Understanding项目地址: https://gitcode.com/OpenGVLab/InternVideo
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考