ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

InternVideo核心代码解析:从3D位置编码到Flash Attention的模型源码深读

InternVideo核心代码解析:从3D位置编码到Flash Attention的模型源码深读 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.py1D / 2D / 3D 正弦余弦位置编码flash_attention_class.pyFlash Attention 封装模块internvideo2.pyInternVideo2 主模型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 分给空间2D1/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_embedTrue空间编码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_attnAttention 类 同时实现了两条前向路径路径说明_naive_attn标准 PyTorch 注意力作为 CPU / 调试时的回退方案_flash_attn走 Flash Attention 内核配合 QK-Normalization 与融合 RMSNormforward里一行代码完成切换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 BlockRMSNorm → 注意力 → 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),仅供参考
返回列表