- 推理引擎
- 大模型
【免费下载链接】FlexGen
Running large language models on a single GPU for throughput-oriented scenarios.
本篇技术指南以 benchmark/third_party/transformers/examples/tensorflow/multiple-choice/README.md 为骨架,结合 run_swag.py 完整源码,系统讲解如何使用 🤗 Transformers 在 TensorFlow/Keras 上微调多选问答(Multiple-choice)模型。读完本文,你将掌握该脚本的三段式参数体系、SWAG 数据预处理与自定义 Data Collator 的实现原理、多 GPU/TPU 分布式训练的配置方法,并能将其直接迁移到自己的多选任务(如常识推理、阅读理解候选题)中。该示例位于本仓库benchmark/third_party/transformers目录下,作为第三方模型基线之一被 FlexGen 性能评测脚本引用对照,理解其训练与推理形态也有助于读懂后续的基准对比逻辑。
一、脚本概述:一个开箱即用的多选问答微调示例
multiple-choice目录下包含三个文件:核心训练脚本 run_swag.py、依赖清单 requirements.txt 以及本 README。该脚本演示的是**多选作答(multiple-choice answering)**范式:给定一段上下文和若干候选结尾/选项,模型需要预测正确的那一个。SWAG(Situations With Adversarial Generations)是这一任务最具代表性的公开数据集,脚本默认直接使用它。
对于常规使用场景,该脚本无需任何修改即可直接运行;源码中通过注释明确标出了需要针对自有项目进行调整的部分(如数据集列名、选项数量、预处理逻辑等)。从实现看,脚本依赖以下关键 Transformers 组件(见 run_swag.py):
TFAutoModelForMultipleChoice:多选任务专用的 TensorFlow 模型基类;TFTrainingArguments:涵盖训练超参、分布式策略、输出与 Hub 推送等全部训练参数;HfArgumentParser:将命令行/JSON 参数解析为类型安全的 dataclass;create_optimizer:构建带 warmup 与权重衰减的 AdamW 优化器及学习率调度;DefaultDataCollator与自定义DataCollatorForMultipleChoice:负责批次组装与动态填充。
脚本开头的check_min_version("4.24.0")强制要求 Transformers 版本不低于 4.24.0,低于此版本会直接报错,这是使用前需要确认的环境前提。
二、环境依赖与安装前提
requirements.txt 给出的依赖极其精简,只有三项:
sentencepiece != 0.1.92 protobuf tensorflow >= 2.3其中sentencepiece用于某些 BPE 类分词器(如 XLM、XLNet 等)的底层切词,排除 0.1.92 版本是为了规避该版本已知问题;protobuf是 TensorFlow 生态的公共依赖;tensorflow >= 2.3声明了脚本可运行的最低 TensorFlow 版本。此外,脚本运行时还需要datasets、transformers两个库(二者在 run_swag.py 顶部导入)。安装方式为常规 pip 安装上述依赖后直接运行脚本,无需编译任何自定义算子。
三、参数体系:三类 dataclass 的完整对照
脚本通过HfArgumentParser同时解析三类参数(见 run_swag.py):模型参数ModelArguments、数据参数DataTrainingArguments、训练参数TFTrainingArguments。支持两种传参方式:
- 命令行传参:
python run_swag.py --model_name_or_path ...; - JSON 文件传参:当只传入一个以
.json结尾的参数时,自动走parser.parse_json_file()分支,把 JSON 反序列化为上述三类参数。
3.1 模型参数(ModelArguments)
| 参数 | 默认值 | 含义 |
|---|---|---|
model_name_or_path | 必填 | 预训练模型路径或 Hugging Face Hub 模型标识(如distilbert-base-cased) |
config_name | None | 预训练配置名称/路径,与模型名不同时使用 |
tokenizer_name | None | 分词器名称/路径,与模型名不同时使用 |
cache_dir | None | 下载的预训练模型缓存目录 |
use_fast_tokenizer | True | 是否使用 tokenizers 库支撑的快速分词器 |
model_revision | "main" | 模型版本(分支名、标签名或 commit id) |
use_auth_token | False | 是否使用huggingface-cli login生成的令牌(访问私有模型时需要) |
参数定义见 run_swag.py。其中config_name与tokenizer_name的设计允许用户对同一个模型权重搭配不同的配置与分词器(例如更换词表或调整结构)。
3.2 数据参数(DataTrainingArguments)
| 参数 | 默认值 | 含义 |
|---|---|---|
train_file | None | 本地训练数据文件(仅支持 csv 或 json,脚本会断言扩展名) |
validation_file | None | 本地评估数据文件(仅支持 csv 或 json) |
overwrite_cache | False | 是否覆盖预处理缓存 |
preprocessing_num_workers | None | 预处理并行进程数 |
max_seq_length | None | 分词后的最大序列长度;超长截断、超短填充 |
pad_to_max_length | False | 是否把全部样本填充到最大长度;False时按批次动态填充(GPU 上更高效,但对 TPU 很不友好) |
max_train_samples | None | 截断训练样本数(调试/快速验证用) |
max_eval_samples | None | 截断评估样本数(调试/快速验证用) |
参数定义见 run_swag.py。需要特别说明的是pad_to_max_length与max_seq_length的组合行为:
- 当
max_seq_length未指定时,脚本读取tokenizer.model_max_length,若该值大于 1024 则自动回退到 1024(日志中会给出警告,可通过--max_seq_length xxx覆盖); - 当显式传入的
max_seq_length超过tokenizer.model_max_length时,脚本取二者较小值并给出警告(见 run_swag.py)。
3.3 训练参数(TFTrainingArguments)
该参数集来自 Transformers 库的TFTrainingArguments,脚本运行示例中直接用到--output_dir、--do_eval、--do_train,其余常用项包括:
--output_dir:输出目录(模型与结果保存位置,必填);--do_train/--do_eval:是否执行训练/评估(评估不依赖训练,可单独运行);--per_device_train_batch_size/--per_device_eval_batch_size:单设备(单 GPU 或 TPU 核)批大小,总批大小会在训练时按副本数放大;--num_train_epochs:训练轮数;--learning_rate、--warmup_steps/--warmup_ratio:学习率与预热(二者任一大于 0 即生效,warmup_ratio按总步数比例计算);--adam_beta1/--adam_beta2/--adam_epsilon:Adam 超参;--weight_decay:权重衰减系数;--max_grad_norm:梯度全局裁剪范数;--xla:启用jit_compile即 XLA 编译;--seed:随机种子(在模型初始化前调用set_seed);--overwrite_output_dir:输出目录已存在且非空时,是否覆盖继续;--push_to_hub系列:--push_to_hub_model_id、--push_to_hub_organization、--push_to_hub_token,用于把模型推送到 Hub。
完整参数可通过python run_swag.py --help查看(源码注释指向src/transformers/training_args.py)。
四、数据加载与预处理全流程
4.1 数据集来源两种模式
脚本支持两种数据来源(见 run_swag.py):
- 默认模式:不传
train_file/validation_file时,自动从 Hub 下载 SWAG 数据集,即load_dataset("swag", "regular"); - 本地文件模式:传入
train_file/validation_file后,按扩展名(csv/json)加载本地数据。分布式场景下,load_dataset保证同一数据只由一个本地进程负责下载。
4.2 SWAG 数据结构与列名约定
SWAG 数据集的字段约定被硬编码在脚本中(见 run_swag.py):
ending_names = [f"ending{i}" for i in range(4)]:4 个候选结尾列ending0~ending3;context_name = "sent1":上下文(前提句);question_header_name = "sent2":附加问题/引导句,与每个候选结尾拼接成完整选项。
换用自有数据集时,这里就是要改的第一个位置:选项数量变了就调整ending_names的长度,字段名不同就替换三个变量名,并在预处理函数中保持同样的拼接语义。
4.3 预处理函数 preprocess_function
核心预处理逻辑(见 run_swag.py)分四步:
- 上下文复制:把
sent1复制 4 份,得到[[context] * 4 ...],与 4 个候选一一对应; - 选项拼接:对每个样本,把
sent2(header)与其 4 个ending分别拼接,构成第二个句子; - 扁平化:
list(chain(*...))把(N, 4)结构展平为4N个独立序列,交给tokenizer(first_sentences, second_sentences, truncation=True, max_length=max_seq_length)成批分词; - 还原形状:把分词结果按每 4 个一组切回
(N, 4)结构,产出每个样本一个input_ids/attention_mask的 4×seq_len 张量。
4.4 为什么需要自定义 Data Collator
多选任务的难点在于:每个训练样本实际包含 4 个输入序列(4 个选项),但标签只有一个整数下标。Transformers 自带的数据整理器无法直接处理这种"一标签对多序列"的结构,因此脚本在 run_swag.py 中实现了DataCollatorForMultipleChoice,其工作流程是:
- 从每个 feature 中弹出
label(兼容labels命名); - 把
(batch_size, num_choices)的特征扁平化为batch_size * num_choices个单序列; - 调用
tokenizer.pad()统一填充,支持三种填充策略(True/'longest'动态填充到批内最长、'max_length'按最大长度、False不填充),以及pad_to_multiple_of参数将序列对齐到某值的整数倍; - 用
tf.reshape(v, (batch_size, num_choices, -1))还原 3D 形状,使模型能以(batch, choices, tokens)的维度接收输入; - 重新附上
tf.int64类型的labels。
其中pad_to_multiple_of的实战价值在于:将序列长度对齐为 8 的倍数可满足 NVIDIA Volta 及更新架构(算力 ≥ 7.5)上 Tensor Core 的 GEMM 对齐要求,显著提升 GPU 利用率。
脚本对两种填充路径做了分流(见 run_swag.py):pad_to_max_length=True时使用DefaultDataCollator(固定长度,适合 TPU);否则使用上述自定义 collator(动态填充,GPU 更高效但"对 TPU 非常不友好")。
五、模型构建、优化器与训练执行
5.1 TFAutoModelForMultipleChoice 与支持的模型家族
脚本通过TFAutoModelForMultipleChoice.from_pretrained(...)按model_name_or_path自动加载对应架构的多选头模型(见 run_swag.py)。从 modeling_tf_auto.py 的映射表可以确认,该 Auto 类当前覆盖 17 个模型家族:ALBERT、BERT、CamemBERT、ConvBERT、DistilBERT、ELECTRA、FlauBERT、Funnel、Longformer、MobileBERT、MPNet、RemBERT、RoBERTa、RoFormer、XLM、XLM-RoBERTa、XLNet。也就是说,把--model_name_or_path换成这些家族中任意一个 Hub 标识或本地路径,脚本即可直接微调。
5.2 分布式策略与批大小核算
在进入训练前,脚本把所有模型构建放在training_args.strategy.scope()上下文内。默认策略是MirroredStrategy:只要机器上有多个 GPU 就会被自动利用,这是 README 中"多 GPU 有效使用"承诺的实现基础。--tpu参数则用于指定 TPU 资源名,将策略切换为 TPU 策略。
总批大小按副本数放大(见 run_swag.py):
num_replicas = training_args.strategy.num_replicas_in_sync total_train_batch_size = training_args.per_device_train_batch_size * num_replicas total_eval_batch_size = training_args.per_device_eval_batch_size * num_replicas即配置中的per_device_*是单设备批大小,实际送入model.fit的批大小会自动乘以设备数。
5.3 优化器与学习率调度
训练步数与优化器均由create_optimizer统一创建(见 run_swag.py):
num_train_steps = (len(train_dataset) // total_train_batch_size) * int(training_args.num_train_epochs)num_warmup_steps的优先级为:warmup_steps > 0时直接采用;否则若warmup_ratio > 0按总步数比例折算;两者都未配置则为 0。优化器为 AdamW,融合了weight_decay_rate与adam_global_clipnorm(梯度全局裁剪)。模型随后执行:
model.compile(optimizer=optimizer, metrics=["accuracy"], jit_compile=training_args.xla)jit_compile=training_args.xla意味着传入--xla即可启用 XLA 编译加速。
5.4 tf.data 管线:prepare_tf_dataset 的正确打开方式
训练与评估数据通过model.prepare_tf_dataset()封装(见 run_swag.py)。从底层实现 modeling_tf_utils.py 看,该方法会把datasets.Dataset包装为可直接送入 Kerasfit()/evaluate()的tf.data.Dataset,并自动完成三件事:根据模型call签名推断输入列名、按标签名识别label/labels列、剔除与模型无关的多余列。因此它是官方推荐的接入方式;若需完全掌控列名映射,才降级使用dataset.to_tf_dataset()手写细节。
脚本还在数据集上设置了tf.data.Options(),把experimental_distribute.auto_shard_policy显式置为AutoShardPolicy.OFF(见 run_swag.py),避免多副本自动分片带来的数据重复/丢失问题。评估数据集额外传了drop_remainder=True,保证每批形状完整。
训练本体即一次标准的 Keras 调用:
history = model.fit( tf_train_dataset, validation_data=validation_data, epochs=int(training_args.num_train_epochs), callbacks=callbacks, )callbacks中仅在--push_to_hub开启时包含PushToHubCallback(默认模型 ID 形如{model_name}-finetuned-multiplechoice)。
5.5 检查点自动恢复
脚本启动时会检查输出目录(见 run_swag.py):
- 若目录中同时存在
CONFIG_NAME(config.json)与TF2_WEIGHTS_NAME(TF 权重文件),则判定为已有检查点,自动从该目录恢复训练,并打印提示日志; - 若目录非空但缺少上述两个文件,直接抛出
ValueError,提示改用--overwrite_output_dir或更换--output_dir。
这一机制让中断后的续训变得安全且显式。
六、评估、结果落盘与模型保存
评估支持两种路径:
- 训练中评估:
do_train与do_eval同时开启时,model.fit的validation_data负责产出每个 epoch 的验证指标,最后取history.history各指标末位值作为eval_metrics(见 run_swag.py); - 独立评估:仅
do_eval时,单独构建评估数据集执行model.evaluate(tf_eval_dataset),得到{"val_loss": ..., "val_accuracy": ...}(见 run_swag.py)。
无论哪种路径,指标最终都以 JSON 形式写入{output_dir}/all_results.json(见 run_swag.py),方便后续程序化读取。未开启--push_to_hub时,训练/评估结束后调用model.save_pretrained(training_args.output_dir)在本地保存完整模型(见 run_swag.py)。
七、实战:从示例命令到自定义任务
README 给出的最小示例命令为:
python run_swag.py \ --model_name_or_path distilbert-base-cased \ --output_dir output \ --do_eval \ --do_train结合前文参数体系,一个面向实际训练、覆盖关键超参与 TPU 场景的完整命令可以是:
python run_swag.py \ --model_name_or_path bert-base-uncased \ --output_dir output_swag \ --do_train \ --do_eval \ --per_device_train_batch_size 32 \ --per_device_eval_batch_size 64 \ --num_train_epochs 3 \ --learning_rate 5e-5 \ --warmup_ratio 0.1 \ --weight_decay 0.01 \ --max_grad_norm 1.0 \ --max_seq_length 128 \ --seed 42- 多 GPU 环境无需额外配置,
MirroredStrategy自动生效;TPU 环境追加--tpu <TPU_RESOURCE_NAME>,同时建议开启--pad_to_max_length(固定长度填充更契合 TPU 的数据管线约束); - 本地自定义数据集(csv/json)只需追加
--train_file train.csv --validation_file valid.csv,并同步修改 run_swag.py 中的列名约定; - 调试阶段用
--max_train_samples 100 --max_eval_samples 100快速验证管线,正式训练前再移除。
关于内存的注意事项同样来自 README 的明确提示:脚本会将全部数据一次性载入内存。绝大多数多选数据集规模较小,这不成问题;但若数据集非常大,则必须改造脚本引入数据流式(streaming)加载——尤其是 TPU 场景,其对数据供给速率的要求更高,全量驻留内存的做法难以满足持续供数需求。
八、小结与仓库定位
本文完整拆解了run_swag.py从参数解析、数据预处理、自定义 collator、分布式训练到评估落盘的全链路:核心要点包括三类 dataclass 参数体系、SWAG 的 4 选项展平/还原预处理、DataCollatorForMultipleChoice的动态填充与 3D 形状还原、prepare_tf_dataset的自动列推断,以及MirroredStrategy/TPU 策略下的批大小核算。将该脚本换一个model_name_or_path、改一处列名约定,即可覆盖绝大多数基于预训练语言模型的多选问答微调场景。
需要提醒的是,本目录位于仓库的 benchmark/third_party/transformers 第三方基准模块中,其定位是作为独立的模型训练/推理基线,与本仓库 FlexGen 核心的吞吐优化引擎并无代码耦合。若你在研究 FlexGen 的基准对比,可进一步阅读 benchmark/README.md 与 benchmark/hf_ds/README.md,理解第三方模型基线与 FlexGen 评测套件的组织关系。
- 推理引擎
- 大模型
【免费下载链接】FlexGen
Running large language models on a single GPU for throughput-oriented scenarios.
相关推荐
FlexGen 仓库中的 TensorFlow 问答微调实战:基于 Transformers `run_qa.py` 训练与评估 SQuAD
FlexGen 仓库中的 TensorFlow 问答微调实战:基于 Transformers run_qa.py 训练与评估 SQuAD 导读:本仓库(Flex
推理引擎大模型OpCore Simplify:自动生成 OpenCore EFI 的黑苹果引导配置指南
OpCore Simplify:自动生成 OpenCore EFI 的黑苹果引导配置指南 手动跟过 OpenCore(macOS 的开源引导加载器)安装文档的人
开发工具CLI用 RoBERTa 微调 Commonsense QA:基于 Fairseq 的多选问答实战指南
用 RoBERTa 微调 Commonsense QA:基于 Fairseq 的多选问答实战指南 本文围绕 infoxlm/fairseq/examples/r
人工智能大模型预训练深度学习NLP计算机视觉多模态语音音频微调
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考