- 推理引擎
- 大模型
【免费下载链接】FlexGen
Running large language models on a single GPU for throughput-oriented scenarios.
本文以当前仓库中 mm-imdb 示例目录 的官方 README 为核心骨架,完整讲解如何利用run_mmimdb.py在 MM-IMDb 多模态数据集上训练与评估一个融合「电影海报图像 + 剧情文本」的 MMBT(Multimodal Bitransformer)多标签分类模型。读完本文,你将掌握 MM-IMDb 数据集的构成、MMBT 的模态融合原理、完整可复制的训练命令与全部参数含义,并能结合仓库源码理解从数据加载、图像编码到早停评估的整条调用链。
一、MM-IMDb 数据集与 MMBT 模型简介
MM-IMDb 是一个多模态(Multimodal)数据集,包含约 26,000 部电影,每部电影同时携带海报图像、剧情简介(plots)以及其他元数据(metadata)。它天然适合训练一个「看图 + 读文」的联合分类模型:模型需要同时利用电影海报的视觉信息与剧情文本的语义信息,为电影打上类型标签。
在 run_mmimdb.py 对应的研究项目中,采用的模型是 MMBT(Supervised 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 accelerate==0.15.0- 脚本还依赖
sklearn(用于 F1 指标计算)、torchvision(用于 ResNet-152 图像编码器)、Pillow(用于图像读取),以及 TensorBoard(torch.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_train | False | 是否执行训练 |
--do_eval | False | 是否在 dev 集上执行评估 |
--max_seq_len | max_seq_length=128 | 文本 token 化的最大序列长度;注意 README 中写作max_seq_len,而源码参数名实际为max_seq_length(见 run_mmimdb.py),使用时以源码参数名为准 |
--gradient_accumulation_steps | 1 | 梯度累积步数,即累积多少步更新后再执行一次反向/更新 |
--num_image_embeds | 1 | 图像编码器输出的图像嵌入数量,决定后续 AdaptiveAvgPool2d 的池化尺寸 |
--num_train_epochs | 3.0 | 训练总轮数(README 示例中用 100 配合早停使用) |
--patience | 5 | 早停(Early Stopping)耐心值:连续多少个 epoch 的 micro-F1 未创新高则提前终止训练 |
只要保证--data_dir下存在train.jsonl/dev.jsonl,并把--output_dir指向一个(当前为空或不存在)的可写目录,上述命令即可完成「训练 + 评估」全流程。
四、核心代码路径与关键调用链解析
README 虽短,但其背后是两条彼此独立、可分别复用的代码路径。理解它们有助于你调试或把该多模态方案迁移到自己的数据上。
4.1 数据加载与多模态样本构造
训练与评估共用同一套数据管道,入口是 load_examples():
def load_examples(args, tokenizer, evaluate=False): 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(pretrained=True) 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_labels=num_labels) model = MMBTForClassification(config, transformer, img_encoder)其中MMBTConfig会把文本 Transformer 的全部配置属性拷贝过来,并追加modal_hidden_size=2048与num_labels(见 configuration_mmbt.py)。本任务中num_labels = 23,对应 23 个电影类型标签。
从 modeling_mmbt.py 的 MMBTModel 可以看到融合方式:
ModalEmbeddings先把图像编码器的输出B×N×2048经一个线性层投影到 BERT 的hidden_size(768),再在序列最前面拼接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_weight=label_weights)get_label_frequencies()统计每个类型在训练集中的出现次数(见 utils_mmimdb.py),罕见类别的pos_weight更大,从而缓解类型分布不均带来的训练偏向。
5.2 训练循环的关键步骤
train()实现了完整的训练管线,要点包括:
- 优化器与调度器:使用
AdamW,且对bias与LayerNorm.weight不施加 weight decay;配合get_linear_schedule_with_warmup线性预热与衰减;默认learning_rate=5e-5、weight_decay=0.0、warmup_steps=0; - 梯度累积:loss 除以
gradient_accumulation_steps后再反向,满足 README 示例中「小 batch + 大累积」的训练策略; - 混合精度:通过
--fp16启用 NVIDIA Apex AMP(fp16_opt_level默认O1); - 分布式与多卡:单卡多 GPU 时自动套
nn.DataParallel,多机时通过--local_rank走DistributedDataParallel(find_unused_parameters=True); - 检查点保存:每
--save_steps(默认 50)步保存checkpoint-{global_step},内含pytorch_model.bin与training_args.bin; - TensorBoard 日志:主进程记录 loss 与学习率等标量,按
--logging_steps(默认 50)输出。
5.3 基于 micro-F1 的早停
每个 epoch 结束后,脚本都会在 dev 集上评估一次,并以 micro-F1 作为早停指标(见 run_mmimdb.py):
results = 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 的算术平均(average="macro");micro_f1:按样本-标签对整体统计的 F1(average="micro"),是早停与最优模型选择的核心指标。
结果会写入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_dir | None | 预训练模型下载缓存目录 |
--per_gpu_train_batch_size | 8 | 每 GPU 训练 batch 大小 |
--per_gpu_eval_batch_size | 8 | 每 GPU 评估 batch 大小 |
--learning_rate | 5e-5 | Adam 初始学习率 |
--weight_decay | 0.0 | 权重衰减系数 |
--adam_epsilon | 1e-8 | Adam 优化器 epsilon |
--max_grad_norm | 1.0 | 梯度裁剪范数上限 |
--max_steps | -1 | 若大于 0,则覆盖num_train_epochs,限定总训练步数 |
--warmup_steps | 0 | 线性预热步数 |
--logging_steps | 50 | 每多少更新步记录一次日志 |
--save_steps | 50 | 每多少更新步保存一次 checkpoint |
--evaluate_during_training | False | 训练过程中是否在每个日志步执行评估 |
--eval_all_checkpoints | False | 评估所有 checkpoint |
--no_cuda | False | 强制不使用 CUDA |
--num_workers | 8 | DataLoader 数据加载线程数 |
--overwrite_output_dir | False | 允许覆盖非空的输出目录(否则训练前会报错退出) |
--overwrite_cache | False | 覆盖缓存的数据集 |
--seed | 42 | 随机种子(脚本在训练前通过set_seed固定) |
--fp16/--fp16_opt_level | False/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-uncased),MMBT 会自动把其 hidden size 作为融合维度。
八、总结
MM-IMDb 示例是理解 MMBT 多模态建模范式的绝佳入口。其 README 虽然简短,但背后由 run_mmimdb.py 与 utils_mmimdb.py 构成了一个完整、可复现的实验闭环:JSONL 多模态数据加载 → ResNet-152 图像编码 → MMBT 嵌入级模态融合 → 带类别权重的多标签 BCE 训练 → 基于 micro-F1 的早停与评估。掌握它之后,无论是复现 MM-IMDb 基准、实验不同num_image_embeds的图像特征粒度,还是迁移到自定义多模态分类任务,你都能快速上手。
- 推理引擎
- 大模型
【免费下载链接】FlexGen
Running large language models on a single GPU for throughput-oriented scenarios.
相关推荐
使用 Flower 与 Hugging Face Transformers 联邦微调大语言模型:IMDB 情感分类快速入门指南
使用 Flower 与 Hugging Face Transformers 联邦微调大语言模型:IMDB 情感分类快速入门指南 本指南基于 Flower 官方
人工智能联邦学习机器学习深度学习CANN/asc-devkit:Ascend C SIMD API存储非对齐数据接口
asc_storeunalign_post_postupdate 产品支持情况 | 产品 | 是否支持 | | : | : :| | Ascend 950PR/
人工智能深度学习算子库CANNAscendHugging Face课程:Transformer模型调试实战指南
Hugging Face课程:Transformer模型调试实战指南 引言 在自然语言处理 NLP 项目中,使用预训练Transformer模型进行微调和推理时
文档教程人工智能NLP深度学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考