- 多模态
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 预训练
【免费下载链接】mmf
A modular framework for vision & language multimodal research from Facebook AI Research (FAIR)
本篇技术指南以 MMF(Facebook AI Research 开源的视觉与语言多模态研究框架)仓库中 M4C 项目文档 为骨架,系统讲解 M4C(Iterative Answer Prediction with Pointer-Augmented Multimodal Transformers)模型在 TextVQA、ST-VQA、OCR-VQA 三个数据集上的数据准备、预训练模型使用、训练与评估命令,并结合 mmf/models/m4c.py 源码剖析其"固定词表分类 + OCR 指针复制"的迭代解码原理。读完本文,你将掌握 M4C 在 MMF 中的完整使用流程(从安装到 EvalAI 预测文件生成),并理解其底层多模态 Transformer 与指针网络的关键实现。
一、M4C 模型是什么
M4C 出自论文Iterative Answer Prediction with Pointer-Augmented Multimodal Transformers for TextVQA(R. Hu, A. Singh, T. Darrell, M. Rohrbach,发表于 CVPR 2020),论文 BibTeX 引用如下:
@inproceedings{hu2020iterative, title={Iterative Answer Prediction with Pointer-Augmented Multimodal Transformers for TextVQA}, author={Hu, Ronghang and Singh, Amanpreet and Darrell, Trevor and Rohrbach, Marcus}, booktitle={Proceedings of the IEEE Conference on Computer Vision and Pattern Recognition}, year={2020} }在 MMF 仓库中,M4C 通过 mmf/models/m4c.py 中的@registry.register_model("m4c")注册为名为m4c的模型,模型类为M4C(BaseModel),其默认配置路径为configs/models/m4c/defaults.yaml。与早期只做"固定词表分类"的 VQA 模型(如 LoRRA)不同,M4C 的核心思路是:
- 用**多模态 Transformer(MMT)**同时建模问题文本、图像物体区域(Faster R-CNN 特征)和 OCR 文本 token;
- 用**指针网络(OCR Pointer Network)**在迭代解码的每一步动态地"复制"某个 OCR token 作为答案的一部分,从而可以回答出词表之外的、图片中实际存在的文字(例如招牌、标签、路牌上的文字);
- 输出分数由固定答案词表分类分数与动态 OCR 复制分数拼接而成(见 mmf/models/m4c.py 的
_forward_output)。
二、安装与依赖
M4C 随 MMF 一起安装。按照 MMF 的 安装指南 安装 MMF 即可,安装过程会一并解决 M4C 的所有依赖:
transformers:用于文本 BERT 编码与多模态 Transformer 的底层实现(TextBert、MMT均继承自BertPreTrainedModel,见 mmf/models/m4c.py);editdistance:用于 ST-VQA 的 ANLS 指标计算(见 mmf/utils/m4c_evaluators.py 中的STVQAANLSEvaluator);- PHOC 特征的 Python 接口:
mmf/utils/phoc/下的 C 扩展会在安装时编译,用于生成 OCR token 的 PHOC(Pyramidal Histogram of Characters)特征。
三、数据说明(TextVQA / ST-VQA / OCR-VQA)
本仓库支持 M4C 在三个数据集上的训练与评估:TextVQA、ST-VQA和OCR-VQA。运行命令时,数据集与相应依赖会通过 MMF 的 zoo 机制自动下载(对应dataset_config.<dataset>.zoo_requirements配置,见下文配置文件小节)。
3.1 ST-VQA 的图片质量问题
官方发现:下载的 ST-VQA 数据中有约1/3的图片来自 COCO-Text,且这些图片不知何故被缩放为 256×256,导致图像质量下降、宽高比失真。因此官方在发布物体与 OCR 特征时,用 COCO-Text 中的原始版本图片替换了这些被缩放的图片,再输入到物体检测与 OCR 系统中提取特征。
3.2 imdb 格式说明
官方发布的 imdb 中包含:
- OCR 识别结果(OCR tokens);
- 归一化边界框(范围在
[0,1]内):- 每个检测物体在
obj_normalized_boxes键下; - 每个 OCR token 在
ocr_normalized_boxes键下。
- 每个检测物体在
另外,ST-VQA 与 OCR-VQA 的 imdb 中答案被平铺(duplicated)为每题 10 个答案,以与 TextVQA imdb 的格式保持一致(TextVQA 每题天然有 10 个标注答案,评测时也按 10 个答案计算 soft score,见 mmf/utils/m4c_evaluators.py)。
3.3 TextVQA 的 OCR 版本
TextVQA 下载的文件同时包含两类 imdb:
- Rosetta-en OCR(性能更好,本文预训练模型表中的默认选择);
- Rosetta-ml OCR(与此前 LoRRA 模型使用的 OCR 结果一致)。
请下载与 OCR 版本对应的 OCR 特征文件:textvqa/ocr_en(Rosetta-en)或textvqa/ocr_ml(Rosetta-ml)。
3.4 特征提取脚本
- 物体 Faster R-CNN 特征由 tools/scripts/features/extract_features_vmb.py 提取(VMB = vqa-maskrcnn-benchmark);
- OCR Faster R-CNN 特征由 projects/m4c/scripts/extract_ocr_frcn_feature.py 提取。该脚本依赖
vqa-maskrcnn-benchmark(可从 ronghanghu 的 fork 安装),接收--detection_cfg、--detection_model、--imdb_file、--image_dir、--save_dir参数,对每个 imdb 条目按ocr_normalized_boxes还原出像素坐标的 OCR 框,送入检测模型得到每个 OCR 框的 2048 维 fc6 特征,并保存_info.npy(含 OCR 框与 token)。
四、预训练模型
官方在三个数据集上发布了以下预训练 M4C 模型。配置文件的基准目录是projects/m4c/configs/:
| 数据集 | 配置文件(位于projects/m4c/configs) | 预训练模型 Key | 指标 | 备注 |
|---|---|---|---|---|
TextVQA(textvqa) | textvqa/joint_with_stvqa.yaml(文档中写作join_with_stvqa.yaml,仓库实际文件名为joint_with_stvqa.yaml) | m4c.textvqa.with_stvqa | val accuracy - 40.55%;test accuracy - 40.46% | Rosetta-en OCR;以 ST-VQA 为额外数据(官方最佳模型) |
TextVQA(textvqa) | textvqa/defaults.yaml | m4c.textvqa.alone | val accuracy - 39.40%;test accuracy - 39.01% | Rosetta-en OCR |
TextVQA(textvqa) | textvqa/ocr_ml.yaml | m4c.textvqa.ocr_ml | val accuracy - 37.06% | Rosetta-ml OCR |
ST-VQA(stvqa) | stvqa/defaults.yaml | m4c.stvqa.defaults | val ANLS - 0.472(accuracy - 38.05%);test ANLS - 0.462 | Rosetta-en OCR |
OCR-VQA(ocrvqa) | ocrvqa/defaults.yaml | m4c.ocrvqa.defaults | val accuracy - 63.52%;test accuracy - 63.87% | Rosetta-en OCR |
这些模型在 MMF 的模型 zoo 中都有对应记录(见 mmf/configs/zoo/models.yaml),每个资源项都带版本号与 SHA-256 hashcode,例如m4c.textvqa.with_stvqa的版本为1.0_2020_06_30。注意这些预训练模型都依赖detectron.vmb_weights(VMB Faster R-CNN 权重),加载时会通过zoo_requirements自动补齐。
五、训练与评估
训练、评估流程与 MMF 通用流程一致(参见 快速开始 的 Training 一节)。mmf_run与mmf_predict是安装后由 setup.py 注册的两个命令行入口。
5.1 在 TextVQA 训练集上训练
mmf_run dataset=textvqa \ model=m4c \ config=projects/m4c/configs/textvqa/defaults.yaml \ env.save_dir=./save/m4c- 将
dataset=textvqa换成stvqa/ocrvqa、将config换成表 1 中对应的配置文件,即可切换到其他数据集与配置; env.save_dir可改成你偏好的任意保存路径。
5.2 用预训练模型在本地验证集上评估
以评估m4c.textvqa.with_stvqa为例:
mmf_run dataset=textvqa \ model=m4c \ config=projects/m4c/configs/textvqa/defaults.yaml \ env.save_dir=./save/m4c \ run_type=val \ checkpoint.resume_zoo=m4c.textvqa.with_stvqa同样可以按需替换dataset、config与checkpoint.resume_zoo。
注意:要评估你自己训练出的 checkpoint,应改用
checkpoint.resume=True且checkpoint.resume_best=True,而不是checkpoint.resume_zoo=...。更细粒度的 checkpoint 加载/恢复机制可参考 checkpointing 教程。
5.3 为 TextVQA 测试集生成 EvalAI 预测文件
mmf_predict dataset=textvqa \ model=m4c \ config=projects/m4c/configs/textvqa/defaults.yaml \ env.save_dir=./save/m4c \ run_type=test \ checkpoint.resume_zoo=m4c.textvqa.with_stvqa- 要在val 集上生成预测,把
run_type=test换成run_type=val; - 要对自己训练的 checkpoint 生成预测,同样把
checkpoint.resume_zoo=...换成checkpoint.resume=True且checkpoint.resume_best=True; - 要为表 1 中其他 TextVQA 预训练模型生成预测,替换
config与checkpoint.resume_zoo即可。
5.4 ST-VQA 联合训练配置
官方最佳 TextVQA 模型(m4c.textvqa.with_stvqa)使用 projects/m4c/configs/textvqa/joint_with_stvqa.yaml 配置。它通过includes: - ./defaults.yaml继承 TextVQA 默认配置,然后:
- 在
zoo_requirements中追加stvqa.defaults与stvqa.ocr_en; - 训练特征同时包含 TextVQA 与 ST-VQA 两套(
textvqa/defaults/features/open_images/detectron.lmdb,textvqa/ocr_en/features/ocr_en_frcn_features.lmdb与stvqa/defaults/features/detectron.lmdb,stvqa/ocr_en/features/ocr_en_frcn_features.lmdb); - 训练标注同样拼接
imdb_train_ocr_en.npy与imdb_subtrain.npy。
textvqa/ocr_ml.yaml则把特征与标注整体切换到 Rosetta-ml 版本(textvqa/ocr_ml/features/ocr_ml_frcn_features.lmdb与imdb_*_ocr_ml.npy)。
六、源码架构:M4C 是如何工作的
M4C 的build()方法将模型拆成五个组件依次构建(mmf/models/m4c.py):
文本编码(
_build_txt_encoding):TextBert(3 层 BERT,默认num_hidden_layers: 3),可由text_bert_init_from_bert_base: true从bert-base-uncased初始化,并挂入finetune_modules使用更小的学习率(lr_scale_text_bert: 0.1)。若其输出维度(768)与 MMT 的hidden_size不一致,则插入一个线性投影层。物体编码(
_build_obj_encoding):物体外观用finetune_faster_rcnn_fpn_fc7图像编码器把 2048 维 Faster R-CNN fc6 特征映射为 fc7;物体位置用 4 维归一化 bbox 坐标。两者经线性层 + LayerNorm 后相加,再经 Dropout 得到obj_mmt_in(mmf/models/m4c.py)。OCR 编码(
_build_ocr_encoding):每个 OCR token 的特征由四段拼接而成(mmf/models/m4c.py):- 300 维FastText词向量(
context_feature_0); - 604 维PHOC特征(
context_feature_1,由 mmf/datasets/processors/processors.py 的PhocProcessor调用mmf/utils/phoc生成); - 2048 维OCR Faster R-CNN fc7外观特征(
image_feature_1); - 50 维 OCR order 向量(LoRRA 遗留,置零,代码注释明确建议 TODO 移除)。
配置文件里的
ocr.mmt_in_dim: 3002正是300 + 604 + 2048 + 50。另外可通过remove_ocr_fasttext / remove_ocr_phoc / remove_ocr_frcn / remove_ocr_semantics / remove_ocr_bbox五个开关做消融实验(置零对应特征)。- 300 维FastText词向量(
多模态 Transformer(
_build_mmt):MMT(4 层 BERT encoder)。其forward将txt_emb、obj_emb、ocr_emb、dec_emb(上一步预测的嵌入)拼接成长序列,并使用类似prefix LM的注意力掩码:编码区元素彼此可互相 attend,解码区元素只能 causal 地 attend 自身及其之前的解码步(mmf/models/m4c.py)。输出层(
_build_output):OcrPtrNet:查询来自 MMT 解码输出mmt_dec_output,键来自mmt_ocr_output,用点积缩放(除以sqrt(query_key_size))加掩码后得到动态 OCR 分数(mmf/models/m4c.py);- 固定词表分类器
ClassifierLayer:输出维度为num_choices - classifier.ocr_max_num,即从"固定词表 + 最多 50 个 OCR 槽位"的总空间中扣除 OCR 复制维度; - 最终
scores = cat([fixed_scores, dynamic_ocr_scores], dim=-1)(mmf/models/m4c.py)。
6.1 迭代解码与教师强制
_forward_mmt_and_output(mmf/models/m4c.py):
- 训练时:从
train_prev_inds取上一步预测索引(由M4CAnswerProcessor在线采样一条答案解码序列,见 mmf/datasets/processors/processors.py),即教师强制(teacher-forcing); - 推理时:先以
BOS_IDX填充第 0 步,然后贪心解码——重复"前向 MMT → 计算分数 →argmax选出词表或 OCR 中得分最高者 → 写回prev_inds"这一循环,直到解码步数用尽(max_copy_steps: 12)。
6.2 上一步预测的嵌入
PrevPredEmbeddings(mmf/models/m4c.py)把"固定词表嵌入表 + OCR 嵌入表"拼接后,用_batch_gather按prev_inds取出上一步预测的嵌入,再加上位置嵌入与类型嵌入(token_type_ids = prev_inds.ge(ans_num),即 0 表示词表、1 表示 OCR),最后 LayerNorm + Dropout 得到解码嵌入。这里固定的最大解码长度为 100、类型数为 5。
6.3 损失函数
M4C 默认使用m4c_decoding_bce_with_mask损失(注册于 mmf/modules/losses.py):对scores与targets逐元素计算 BCE(binary_cross_entropy_with_logits),再乘以train_loss_mask(只对有效解码步施加损失),最后除以 mask 和作为归一化。该 mask 与prev_inds一样由M4CAnswerProcessor在数据管线中生成。
七、核心配置文件逐段解读
7.1 模型配置 mmf/configs/models/m4c/defaults.yaml
| 配置键 | 默认值 | 说明 |
|---|---|---|
lr_scale_frcn | 0.1 | Faster R-CNN fc7 层的学习率缩放(预训练部分用小学习率微调) |
lr_scale_text_bert | 0.1 | 文本 BERT 的学习率缩放 |
lr_scale_mmt | 1.0 | 多模态 Transformer 的学习率缩放(不缩放) |
text_bert_init_from_bert_base | true | 是否从bert-base-uncased初始化文本编码 |
text_bert.num_hidden_layers | 3 | 文本 BERT 层数 |
obj.mmt_in_dim | 2048 | 物体外观特征维度 |
obj.dropout_prob | 0.1 | 物体编码 Dropout |
ocr.mmt_in_dim | 3002 | OCR 特征维度(300 FastText + 604 PHOC + 2048 Faster R-CNN + 50 遗留) |
ocr.dropout_prob | 0.1 | OCR 编码 Dropout |
mmt.hidden_size | 768 | 多模态 Transformer 隐层维度 |
mmt.num_hidden_layers | 4 | 多模态 Transformer 层数 |
classifier.type | linear | 固定词表分类器类型 |
classifier.ocr_max_num | 50 | OCR 复制槽位上限(在总输出维度中扣除) |
classifier.ocr_ptr_net.hidden_size/query_key_size | 768 / 768 | 指针网络的查询与键维度 |
model_data_dir | ${env.data_dir} | Faster R-CNN fc7 权重目录 |
这些lr_scale_*会在get_optimizer_parameters(mmf/models/m4c.py)中生效:把挂入finetune_modules的模块参数按base_lr * lr_scale单独分组,其余参数使用默认学习率。
7.2 数据集配置(以 projects/m4c/configs/textvqa/defaults.yaml 为例)
关键点:
zoo_requirements:textvqa.defaults与textvqa.ocr_en,运行时自动下载数据与特征;features:train/val/test 均为textvqa/defaults/features/open_images/detectron.lmdb,textvqa/ocr_en/features/ocr_en_frcn_features.lmdb(逗号分隔物体特征与 OCR 特征两个 lmdb);- processors 管线(与 mmf/datasets/processors/processors.py 中的注册处理器一一对应):
bert_tokenizer:问题分词,max_seq_length: 20;m4c_answer:答案迭代解码目标构造,max_length: 50、max_copy_steps: 12、num_answers: 10,词表为fixed_answer_vocab_textvqa_5k.txt(5k 固定答案);M4CAnswerProcessor内部保证PAD_IDX == 0、BOS_IDX/EOS_IDX/UNK_IDX均有效(mmf/datasets/processors/processors.py);copy:OCR token 索引复制,max_length: 100;phoc:PHOC 特征,max_length: 50;fasttext:OCR FastText 词向量,model_file: wiki.en.bin(首次使用会从缓存目录下载);ocr_token_processor(simple_word)与bbox(归一化 bbox,max_length: 50);
- 开关:
return_features_info: true、use_ocr: true、use_ocr_info: true、use_order_vectors: true; - 优化器:Adam,
lr: 1e-4,eps: 1e-8,weight_decay: 0; - 训练计划:
max_updates: 24000、batch_size: 128、num_workers: 4,梯度裁剪max_grad_l2_norm: 0.25(clip_norm_mode: all),学习率在lr_steps: [14000, 19000]处以lr_ratio: 0.1衰减,并启用 warmup(warmup_factor: 0.2,warmup_iterations: 1000); - 评估指标:TextVQA 用
textvqa_accuracy,ST-VQA 用stvqa_accuracy+stvqa_anls,OCR-VQA 用ocrvqa_accuracy;early_stop.criteria分别指向对应指标。ST-VQA 与 OCR-VQA 的配置结构完全相同,仅词表(fixed_answer_vocab_stvqa_5k.txt/fixed_answer_vocab_ocrvqa_82.txt)、特征路径与max_updates(OCR-VQA 为 48000,lr_steps: [28000, 38000])不同。
八、评估指标实现
- TextVQA accuracy(mmf/utils/m4c_evaluators.py):按 EvalAI 的 soft score 计算——对 10 个标注答案去重后,每个唯一答案的分数为
min(1, 匹配数/3)的平均,预测答案命中该分数即得 acc; - ST-VQA accuracy(mmf/utils/m4c_evaluators.py):预测答案与任一 GT 完全匹配得 1 分,否则 0;
- ST-VQA ANLS(mmf/utils/m4c_evaluators.py):ANLS =
1 - edit_distance / max(len(s1), len(s2)),低于 0.5 的相似度计为 0,最终取对每个 GT 的最大值并求平均——这也是 M4C 安装依赖中包含editdistance的原因。
九、常见问题与建议
- 配置文件名差异:文档中的
textvqa/join_with_stvqa.yaml在仓库中的实际文件名为textvqa/joint_with_stvqa.yaml,使用时以仓库实际文件名为准; - OCR 版本必须匹配:用
ocr_ml.yaml训练/评估就必须使用textvqa.ocr_ml预训练模型与 Rosetta-ml 特征,混用 Rosetta-en/ml 会导致指标不一致; - 显存与速度:MMT 输入序列为"问题(≤20)+ 物体(≤100)+ OCR(≤50)+ 解码步(≤12)",batch_size 128、num_workers 4 为官方默认训练配置,实际可依据硬件调整;
- 消融实验:可通过
model_config.m4c.ocr.remove_ocr_*五个开关分别关闭 FastText、PHOC、OCR 外观、OCR 语义与 OCR bbox 特征,验证各模态对最终精度的贡献。
通过本文的配置表与源码解读,你可以直接从mmf_run dataset=textvqa model=m4c起步完成 M4C 的复现与实验,也可以进一步阅读 mmf/models/m4c.py 与各数据集配置深入定制模型。
- 多模态
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 预训练
【免费下载链接】mmf
A modular framework for vision & language multimodal research from Facebook AI Research (FAIR)
相关推荐
MMF 中的 M4C 模型:面向 TextVQA 的指针增强多模态 Transformer 实战指南
MMF 中的 M4C 模型:面向 TextVQA 的指针增强多模态 Transformer 实战指南 本篇技术指南围绕 MMF 框架内实现的 M4C(Itera
多模态人工智能深度学习NLP计算机视觉预训练MMF 实战:用 M4C 模型参加 TextVQA Challenge 的完整训练、评估与提交指南
MMF 实战:用 M4C 模型参加 TextVQA Challenge 的完整训练、评估与提交指南 TextVQA Challenge 是一项要求模型「阅读」图
多模态人工智能深度学习NLP计算机视觉预训练MMF 快速上手指南:使用 M4C 模型在 TextVQA 数据集上完成训练与推理
MMF 快速上手指南:使用 M4C 模型在 TextVQA 数据集上完成训练与推理 本指南以 MMF(Facebook AI Research 开源的视觉与语言
多模态人工智能深度学习NLP计算机视觉预训练
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考