基于 Hugging Face Transformers 在 MM-IMDb 上微调多模态双 Transformer(MMBT)分类模型实战指南

发布时间:2026/9/25 16:55:06
基于 Hugging Face Transformers 在 MM-IMDb 上微调多模态双 Transformer(MMBT)分类模型实战指南 推理引擎大模型【免费下载链接】FlexGenRunning large language models on a single GPU for throughput-oriented scenarios.项目地址https://gitcode.com/gh_mirrors/fl/FlexGen点击查看免费下载本文以当前仓库中 mm-imdb 示例目录 的官方 README 为核心骨架完整讲解如何利用run_mmimdb.py在 MM-IMDb 多模态数据集上训练与评估一个融合「电影海报图像 剧情文本」的 MMBTMultimodal Bitransformer多标签分类模型。读完本文你将掌握 MM-IMDb 数据集的构成、MMBT 的模态融合原理、完整可复制的训练命令与全部参数含义并能结合仓库源码理解从数据加载、图像编码到早停评估的整条调用链。一、MM-IMDb 数据集与 MMBT 模型简介MM-IMDb 是一个多模态Multimodal数据集包含约 26,000 部电影每部电影同时携带海报图像、剧情简介plots以及其他元数据metadata。它天然适合训练一个「看图 读文」的联合分类模型模型需要同时利用电影海报的视觉信息与剧情文本的语义信息为电影打上类型标签。在 run_mmimdb.py 对应的研究项目中采用的模型是 MMBTSupervised Multimodal Bitransformer。从 modeling_mmbt.py 的模型说明可以看到MMBT 是一种有监督的多模态双 Transformer 模型它把文本编码器如 BERT与另一个模态此处为图像的编码器输出的特征融合起来再送入同一个 Transformer 编码器进行联合建模从而在多模态分类基准上取得当时领先的效果。在 MM-IMDb 场景中两个编码器分别是文本编码器bert-base-uncased等预训练 BERT由--model_name_or_path指定图像编码器ResNet-152详见下文「图像编码器」小节。二、训练环境准备本示例位于仓库的第三方依赖目录下属于 Hugging Face Transformers 的研究型示例research_projects。运行脚本前需要确认 Python 环境已安装 PyTorch 与 Transformers 相关依赖安装仓库维护的 Transformers fork。根据 benchmark/third_party/README.md 的说明当前仓库维护的是 huggingface/transformers v4.24.0 的 fork可按如下方式安装cd FlexGen/benchmark/third_party/transformers pip3 install -e . pip3 install accelerate0.15.0脚本还依赖sklearn用于 F1 指标计算、torchvision用于 ResNet-152 图像编码器、Pillow用于图像读取以及 TensorBoardtorch.utils.tensorboard缺失时脚本会回退到tensorboardX见 run_mmimdb.py。注意本示例脚本是为研究实验设计的并不依赖 GPU 加速以外的特殊硬件CPU 亦可运行但训练速度会显著变慢。三、在 MM-IMDb 上训练与评估的完整命令README 给出了一个可直接复制的训练命令模板。请将其中的路径替换为你本机的真实路径python run_mmimdb.py \ --data_dir /path/to/mmimdb/dataset/ \ --model_type bert \ --model_name_or_path bert-base-uncased \ --output_dir /path/to/save/dir/ \ --do_train \ --do_eval \ --max_seq_len 512 \ --gradient_accumulation_steps 20 \ --num_image_embeds 3 \ --num_train_epochs 100 \ --patience 5对上述命令的解读与源码中的参数定义逐一对应命令行参数源码默认值作用说明--data_dir无必填MM-IMDb 数据目录内部应包含train.jsonl与dev.jsonl两个 JSONL 文件源码在 load_examples 中按evaluate标志拼接文件名读取--model_name_or_path无必填预训练模型路径或 huggingface.co 上的模型标识如bert-base-uncased--output_dir无必填模型预测结果与检查点checkpoint的输出目录--do_trainFalse是否执行训练--do_evalFalse是否在 dev 集上执行评估--max_seq_lenmax_seq_length128文本 token 化的最大序列长度注意 README 中写作max_seq_len而源码参数名实际为max_seq_length见 run_mmimdb.py使用时以源码参数名为准--gradient_accumulation_steps1梯度累积步数即累积多少步更新后再执行一次反向/更新--num_image_embeds1图像编码器输出的图像嵌入数量决定后续 AdaptiveAvgPool2d 的池化尺寸--num_train_epochs3.0训练总轮数README 示例中用 100 配合早停使用--patience5早停Early Stopping耐心值连续多少个 epoch 的 micro-F1 未创新高则提前终止训练只要保证--data_dir下存在train.jsonl/dev.jsonl并把--output_dir指向一个当前为空或不存在的可写目录上述命令即可完成「训练 评估」全流程。四、核心代码路径与关键调用链解析README 虽短但其背后是两条彼此独立、可分别复用的代码路径。理解它们有助于你调试或把该多模态方案迁移到自己的数据上。4.1 数据加载与多模态样本构造训练与评估共用同一套数据管道入口是 load_examples()def load_examples(args, tokenizer, evaluateFalse): path os.path.join(args.data_dir, dev.jsonl if evaluate else train.jsonl) transforms get_image_transforms() labels get_mmimdb_labels() dataset JsonlDataset(path, tokenizer, transforms, labels, args.max_seq_length - args.num_image_embeds - 2)这里有一个值得注意的细节传给数据集的文本最大长度是args.max_seq_length - args.num_image_embeds - 2。减法中的- 2是因为 JsonlDataset.getitem会把 tokenize 后句子的首 token通常是[CLS]和尾 token通常是[SEP]拆出来分别作为「图像起始 token」与「图像结束 token」而- num_image_embeds是为拼接在前的图像嵌入预留的序列长度。这样拼接后的总序列长度恰好不超过max_seq_length。具体的数据结构由 utils_mmimdb.py 定义JsonlDataset读取 JSONL 文件每行是一个电影样本字段至少包含text剧情文本、img相对data_dir的图像路径与label类型标签列表。__getitem__返回image_start_token、image_end_token、sentence、image、label五个字段。collate_fn把一批样本整理成定长张量返回顺序为(text_tensor, mask_tensor, img_tensor, img_start_token, img_end_token, tgt_tensor)——这与训练/评估循环中batch[0]~batch[5]的取用方式一一对应见 run_mmimdb.py 的 train。get_mmimdb_labels()返回 23 个电影类型标签Crime、Drama、Thriller、Action、Comedy、Romance、Documentary、Short、Mystery、History、Family、Adventure、Fantasy、Sci-Fi、Western、Horror、Sport、War、Music、Musical、Animation、Biography、Film-Noir标签以 one-hot 多标签形式编码。4.2 图像编码器ResNet-152 自适应平均池化ImageEncoder是图像模态的核心组件model torchvision.models.resnet152(pretrainedTrue) modules list(model.children())[:-2] # 去掉最后的 avgpool 与全连接层 self.model nn.Sequential(*modules) self.pool nn.AdaptiveAvgPool2d(POOLING_BREAKDOWN[args.num_image_embeds])其 forward 的维度变换注释为Bx3x224x224 - Bx2048x7x7 - Bx2048xN - BxNx2048即输入3×224×224的 RGB 海报图像224 来自 get_image_transforms 中的 Resize(256) CenterCrop(224) 预处理并做了针对该数据集的均值/方差归一化经过 ResNet-152 骨干输出2048×7×7特征图由AdaptiveAvgPool2d池化到N个位置N --num_image_embeds展平并转置为B×N×2048其中2048正是MMBTConfig中modal_hidden_size的默认值。POOLING_BREAKDOWN表见 utils_mmimdb.py规定了不同num_image_embeds对应的池化网格num_image_embeds池化尺寸1(1, 1)2(2, 1)3(3, 1)4(2, 2)5(5, 1)6(3, 2)7(7, 1)8(4, 2)9(3, 3)4.3 MMBT 模态融合文本与图像在 Embedding 层拼接模型的组装发生在 run_mmimdb.py 的 main() 中transformer_config AutoConfig.from_pretrained(...) tokenizer AutoTokenizer.from_pretrained(...) transformer AutoModel.from_pretrained(...) img_encoder ImageEncoder(args) config MMBTConfig(transformer_config, num_labelsnum_labels) model MMBTForClassification(config, transformer, img_encoder)其中MMBTConfig会把文本 Transformer 的全部配置属性拷贝过来并追加modal_hidden_size2048与num_labels见 configuration_mmbt.py。本任务中num_labels 23对应 23 个电影类型标签。从 modeling_mmbt.py 的 MMBTModel 可以看到融合方式ModalEmbeddings先把图像编码器的输出B×N×2048经一个线性层投影到 BERT 的hidden_size768再在序列最前面拼接start_token[CLS]的 word embedding、在末尾拼接end_token[SEP]的 word embedding并加上位置与 token type 嵌入文本侧按正常流程取 BERT 的 word 嵌入两者沿序列维torch.cat成一个完整的嵌入序列送入 BERT 的 Transformer encoder 做联合自注意力建模MMBTForClassification取池化输出经过 Dropout 与一个nn.Linear(hidden_size, num_labels)分类头得到 logits见 modeling_mmbt.py。也就是说MMBT 并没有在 Transformer 之后做简单的向量拼接而是让图像特征以「虚拟 token」的身份进入 BERT 的注意力层与文本 token 进行深度交互——这正是它被称为 Bitransformer 的原因。五、训练循环、损失函数与早停机制5.1 多标签损失与类别不均衡处理由于电影可以同时属于多个类型脚本没有使用交叉熵而是在 main() 中构造了带正样本权重pos_weight的二元交叉熵label_frequences train_dataset.get_label_frequencies() label_frequences [label_frequences[l] for l in labels] label_weights (torch.tensor(label_frequences) / len(train_dataset)) ** -1 criterion nn.BCEWithLogitsLoss(pos_weightlabel_weights)get_label_frequencies()统计每个类型在训练集中的出现次数见 utils_mmimdb.py罕见类别的pos_weight更大从而缓解类型分布不均带来的训练偏向。5.2 训练循环的关键步骤train()实现了完整的训练管线要点包括优化器与调度器使用AdamW且对bias与LayerNorm.weight不施加 weight decay配合get_linear_schedule_with_warmup线性预热与衰减默认learning_rate5e-5、weight_decay0.0、warmup_steps0梯度累积loss 除以gradient_accumulation_steps后再反向满足 README 示例中「小 batch 大累积」的训练策略混合精度通过--fp16启用 NVIDIA Apex AMPfp16_opt_level默认O1分布式与多卡单卡多 GPU 时自动套nn.DataParallel多机时通过--local_rank走DistributedDataParallelfind_unused_parametersTrue检查点保存每--save_steps默认 50步保存checkpoint-{global_step}内含pytorch_model.bin与training_args.binTensorBoard 日志主进程记录 loss 与学习率等标量按--logging_steps默认 50输出。5.3 基于 micro-F1 的早停每个 epoch 结束后脚本都会在 dev 集上评估一次并以 micro-F1 作为早停指标见 run_mmimdb.pyresults evaluate(args, model, tokenizer, criterion) if results[micro_f1] best_f1: best_f1 results[micro_f1] n_no_improve 0 else: n_no_improve 1 if n_no_improve args.patience: train_iterator.close() break这正是 README 示例中把--num_train_epochs设为 100、--patience设为 5 的原因模型最多训练 100 轮但一旦连续 5 轮 micro-F1 没有提升就提前停止兼顾效果与时间成本。5.4 评估指标evaluate()在 dev 集上以sigmoid(logits) 0.5作为多标签判定阈值并计算三个指标loss平均二元交叉熵损失macro_f1每个类型 F1 的算术平均averagemacromicro_f1按样本-标签对整体统计的 F1averagemicro是早停与最优模型选择的核心指标。结果会写入output_dir/{prefix}/eval_results.txt训练结束后若--do_eval开启脚本还会对output_dir下保存的模型权重做最终评估若加--eval_all_checkpoints则逐一评估所有 checkpoint见 run_mmimdb.py。六、其他常用训练参数速查除上述命令涉及的参数外脚本还支持一系列常用的微调参数默认值与含义均来自 run_mmimdb.py 的 argparse 定义参数默认值说明--config_name/--tokenizer_name与模型名不同的配置/分词器名称或路径--cache_dirNone预训练模型下载缓存目录--per_gpu_train_batch_size8每 GPU 训练 batch 大小--per_gpu_eval_batch_size8每 GPU 评估 batch 大小--learning_rate5e-5Adam 初始学习率--weight_decay0.0权重衰减系数--adam_epsilon1e-8Adam 优化器 epsilon--max_grad_norm1.0梯度裁剪范数上限--max_steps-1若大于 0则覆盖num_train_epochs限定总训练步数--warmup_steps0线性预热步数--logging_steps50每多少更新步记录一次日志--save_steps50每多少更新步保存一次 checkpoint--evaluate_during_trainingFalse训练过程中是否在每个日志步执行评估--eval_all_checkpointsFalse评估所有 checkpoint--no_cudaFalse强制不使用 CUDA--num_workers8DataLoader 数据加载线程数--overwrite_output_dirFalse允许覆盖非空的输出目录否则训练前会报错退出--overwrite_cacheFalse覆盖缓存的数据集--seed42随机种子脚本在训练前通过set_seed固定--fp16/--fp16_opt_levelFalse/O1是否启用 Apex 混合精度及其 AMP 优化级别--local_rank-1分布式训练的 local_rank-1 表示单进程--server_ip/--server_port远程调试ptvsd附加地址一个值得强调的坑README 示例中的--max_seq_len与源码 argparse 定义的--max_seq_length不一致。若直接照抄 README程序会因未知参数报错实际运行时应使用--max_seq_length 512。七、进阶把多模态方案迁移到自己的数据若希望复用这套代码处理自己的「图像 文本」多标签任务可以遵循以下最小改造路径当前仓库为只读请在本地另建副本修改准备 JSONL 数据仿照 MM-IMDb 格式每行包含text文本内容、img相对数据目录的图像路径、label标签名列表并拆分为train.jsonl与dev.jsonl替换标签表修改get_mmimdb_labels()返回你自己的标签列表适配图像统计量若图像内容差异大可重新统计均值/方差并修改get_image_transforms()中的 Normalize 参数调整池化--num_image_embeds控制图像特征 token 数若图像分辨率或语义粒度不同可通过POOLING_BREAKDOWN表格重新映射更换文本底座--model_name_or_path支持任何 AutoModel 兼容的预训练文本模型如bert-base-uncasedMMBT 会自动把其 hidden size 作为融合维度。八、总结MM-IMDb 示例是理解 MMBT 多模态建模范式的绝佳入口。其 README 虽然简短但背后由 run_mmimdb.py 与 utils_mmimdb.py 构成了一个完整、可复现的实验闭环JSONL 多模态数据加载 → ResNet-152 图像编码 → MMBT 嵌入级模态融合 → 带类别权重的多标签 BCE 训练 → 基于 micro-F1 的早停与评估。掌握它之后无论是复现 MM-IMDb 基准、实验不同num_image_embeds的图像特征粒度还是迁移到自定义多模态分类任务你都能快速上手。赞分享推理引擎大模型【免费下载链接】FlexGenRunning large language models on a single GPU for throughput-oriented scenarios.项目地址https://gitcode.com/gh_mirrors/fl/FlexGen点击查看免费下载相关推荐使用 Flower 与 Hugging Face Transformers 联邦微调大语言模型IMDB 情感分类快速入门指南使用 Flower 与 Hugging Face Transformers 联邦微调大语言模型IMDB 情感分类快速入门指南 本指南基于 Flower 官方人工智能联邦学习机器学习深度学习CANN/asc-devkitAscend C SIMD API存储非对齐数据接口asc_storeunalign_post_postupdate 产品支持情况 | 产品 | 是否支持 | | : | : :| | Ascend 950PR/人工智能深度学习算子库CANNAscendHugging Face课程Transformer模型调试实战指南Hugging Face课程Transformer模型调试实战指南 引言 在自然语言处理 NLP 项目中使用预训练Transformer模型进行微调和推理时文档教程人工智能NLP深度学习上一篇PluginEval 锚定评分标准全解judge 四维 Rubrics 的分级细则与源码实现下一篇RuView homecore-server 运维评审清单从 homecore metaharness 的 operate-server playbook 到服务器源码的逐项印证创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询