☰
Transformers TensorFlow 多选问答微调实战:基于 SWAG 的 run_swag.py 脚本深度解析
2026/9/25 5:41:22 网站建设 项目流程
  • 推理引擎
  • 大模型

【免费下载链接】FlexGen

Running large language models on a single GPU for throughput-oriented scenarios.

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

本篇技术指南以 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_nameNone预训练配置名称/路径,与模型名不同时使用
tokenizer_nameNone分词器名称/路径,与模型名不同时使用
cache_dirNone下载的预训练模型缓存目录
use_fast_tokenizerTrue是否使用 tokenizers 库支撑的快速分词器
model_revision"main"模型版本(分支名、标签名或 commit id)
use_auth_tokenFalse是否使用huggingface-cli login生成的令牌(访问私有模型时需要)

参数定义见 run_swag.py。其中config_name与tokenizer_name的设计允许用户对同一个模型权重搭配不同的配置与分词器(例如更换词表或调整结构)。

3.2 数据参数(DataTrainingArguments)

参数默认值含义
train_fileNone本地训练数据文件(仅支持 csv 或 json,脚本会断言扩展名)
validation_fileNone本地评估数据文件(仅支持 csv 或 json)
overwrite_cacheFalse是否覆盖预处理缓存
preprocessing_num_workersNone预处理并行进程数
max_seq_lengthNone分词后的最大序列长度;超长截断、超短填充
pad_to_max_lengthFalse是否把全部样本填充到最大长度;False时按批次动态填充(GPU 上更高效,但对 TPU 很不友好)
max_train_samplesNone截断训练样本数(调试/快速验证用)
max_eval_samplesNone截断评估样本数(调试/快速验证用)

参数定义见 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):

  1. 默认模式:不传train_file/validation_file时,自动从 Hub 下载 SWAG 数据集,即load_dataset("swag", "regular");
  2. 本地文件模式:传入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)分四步:

  1. 上下文复制:把sent1复制 4 份,得到[[context] * 4 ...],与 4 个候选一一对应;
  2. 选项拼接:对每个样本,把sent2(header)与其 4 个ending分别拼接,构成第二个句子;
  3. 扁平化:list(chain(*...))把(N, 4)结构展平为4N个独立序列,交给tokenizer(first_sentences, second_sentences, truncation=True, max_length=max_seq_length)成批分词;
  4. 还原形状:把分词结果按每 4 个一组切回(N, 4)结构,产出每个样本一个input_ids/attention_mask的 4×seq_len 张量。

4.4 为什么需要自定义 Data Collator

多选任务的难点在于:每个训练样本实际包含 4 个输入序列(4 个选项),但标签只有一个整数下标。Transformers 自带的数据整理器无法直接处理这种"一标签对多序列"的结构,因此脚本在 run_swag.py 中实现了DataCollatorForMultipleChoice,其工作流程是:

  1. 从每个 feature 中弹出label(兼容labels命名);
  2. 把(batch_size, num_choices)的特征扁平化为batch_size * num_choices个单序列;
  3. 调用tokenizer.pad()统一填充,支持三种填充策略(True/'longest'动态填充到批内最长、'max_length'按最大长度、False不填充),以及pad_to_multiple_of参数将序列对齐到某值的整数倍;
  4. 用tf.reshape(v, (batch_size, num_choices, -1))还原 3D 形状,使模型能以(batch, choices, tokens)的维度接收输入;
  5. 重新附上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。

这一机制让中断后的续训变得安全且显式。

六、评估、结果落盘与模型保存

评估支持两种路径:

  1. 训练中评估:do_train与do_eval同时开启时,model.fit的validation_data负责产出每个 epoch 的验证指标,最后取history.history各指标末位值作为eval_metrics(见 run_swag.py);
  2. 独立评估:仅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.

项目地址:https://gitcode.com/gh_mirrors/fl/FlexGen
点击查看免费下载
上一篇:Binwalk终极升级指南:7步完成版本迁移与兼容性处理
下一篇:iOS骨架屏终极指南:SkeletonView集成CocoaPods与SPM完整教程

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

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

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

立即咨询