☰
MMF 框架中的 M4C 模型:基于 Pointer-Augmented Multimodal Transformers 的 TextVQA 迭代答案预测实战指南
2026/10/12 4:20:22 网站建设 项目流程
  • 多模态
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 预训练

【免费下载链接】mmf

A modular framework for vision & language multimodal research from Facebook AI Research (FAIR)

项目地址:https://gitcode.com/gh_mirrors/mm/mmf
点击查看免费下载

本篇技术指南以 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_stvqaval accuracy - 40.55%;test accuracy - 40.46%Rosetta-en OCR;以 ST-VQA 为额外数据(官方最佳模型)
TextVQA(textvqa)textvqa/defaults.yamlm4c.textvqa.aloneval accuracy - 39.40%;test accuracy - 39.01%Rosetta-en OCR
TextVQA(textvqa)textvqa/ocr_ml.yamlm4c.textvqa.ocr_mlval accuracy - 37.06%Rosetta-ml OCR
ST-VQA(stvqa)stvqa/defaults.yamlm4c.stvqa.defaultsval ANLS - 0.472(accuracy - 38.05%);test ANLS - 0.462Rosetta-en OCR
OCR-VQA(ocrvqa)ocrvqa/defaults.yamlm4c.ocrvqa.defaultsval 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):

  1. 文本编码(_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不一致,则插入一个线性投影层。

  2. 物体编码(_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)。

  3. 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五个开关做消融实验(置零对应特征)。

  4. 多模态 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)。

  5. 输出层(_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_frcn0.1Faster R-CNN fc7 层的学习率缩放(预训练部分用小学习率微调)
lr_scale_text_bert0.1文本 BERT 的学习率缩放
lr_scale_mmt1.0多模态 Transformer 的学习率缩放(不缩放)
text_bert_init_from_bert_basetrue是否从bert-base-uncased初始化文本编码
text_bert.num_hidden_layers3文本 BERT 层数
obj.mmt_in_dim2048物体外观特征维度
obj.dropout_prob0.1物体编码 Dropout
ocr.mmt_in_dim3002OCR 特征维度(300 FastText + 604 PHOC + 2048 Faster R-CNN + 50 遗留)
ocr.dropout_prob0.1OCR 编码 Dropout
mmt.hidden_size768多模态 Transformer 隐层维度
mmt.num_hidden_layers4多模态 Transformer 层数
classifier.typelinear固定词表分类器类型
classifier.ocr_max_num50OCR 复制槽位上限(在总输出维度中扣除)
classifier.ocr_ptr_net.hidden_size/query_key_size768 / 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)

项目地址:https://gitcode.com/gh_mirrors/mm/mmf
点击查看免费下载

相关推荐

上一篇:openeuler/rockchip镜像构建终极指南:支持RK3399/RK3588的完整步骤
下一篇:rpmdepsearch开发者指南:如何贡献代码和扩展功能

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

立即咨询