- 推理引擎
- 大模型
【免费下载链接】FlexGen
Running large language models on a single GPU for throughput-oriented scenarios.
本篇技术指南以 FlexGen 仓库基准测试套件中收录的 TensorFlow 摘要微调示例(README.md 及其配套脚本 run_summarization.py)为主体,系统讲解如何在单机多卡或 TPU 环境下,用 BART、T5 等 Seq2Seq 模型对 CNN/DailyMail 等数据集做摘要任务的训练与评估。读完本文,你将掌握该脚本的完整命令行用法、全部核心参数语义、数据预处理与 ROUGE 评估的内部实现,以及 MirroredStrategy/TPU 分布式策略的底层选择逻辑。
一、示例概览:文档说了什么
关联文档是一份精炼的脚本说明,核心信息有三点:
- 用途:演示如何使用 🤗 Transformers 库训练一个“摘要生成”(summarization)模型;对于标准场景可以直接复用,脚本内也通过注释标注了需要按自己项目调整的位置。
- 分布式能力:脚本默认使用
MirroredStrategy,在多 GPU 可用时会自动生效;通过--tpu参数传入 TPU 资源名即可切换到 TPU 训练。 - 一条开箱即用的训练命令:以
facebook/bart-base为起点,在cnn_dailymail3.0.0 数据集上做 3 个 epoch 的微调并同时训练、评估。
这份文档虽短,但与之配套的 run_summarization.py 是一份 700+ 行的完整实现,覆盖参数解析、数据加载、预处理、模型加载、优化器构建、ROUGE 评估与模型导出全流程。下文将以此为主线逐层展开。
二、环境准备与依赖版本要求
运行脚本前需要安装依赖,版本约束见 requirements.txt:
datasets >= 1.4.0 tensorflow >= 2.3.0 evaluate >= 0.2.0此外,脚本自身还做了三道硬性校验(对应 run_summarization.py):
check_min_version("4.24.0"):Transformers 版本低于 4.24.0 直接报错;require_version("datasets>=1.8.0", ...):datasets库低于 1.8.0 会提示安装;nltk的tokenizers/punkt数据缺失时,会自动通过nltk.download("punkt")下载;若处于离线模式(TRANSFORMERS_OFFLINE环境变量)则会明确抛出异常,提示先联网下载一次。
评估阶段依赖 ROUGE 指标库(evaluate.load("rouge")),因此evaluate也是必装项。
三、开箱即用的训练命令(原文核心命令详解)
README 给出的示例命令如下:
python run_summarization.py \ --model_name_or_path facebook/bart-base \ --dataset_name cnn_dailymail \ --dataset_config "3.0.0" \ --output_dir /tmp/tst-summarization \ --per_device_train_batch_size 8 \ --per_device_eval_batch_size 16 \ --num_train_epochs 3 \ --do_train \ --do_eval各参数含义与取值要点:
| 参数 | 含义 | 说明 |
|---|---|---|
--model_name_or_path | 预训练模型标识或本地路径 | 可填 Hugging Face Hub 模型 ID(如facebook/bart-base、t5-base),或本地已下载的 checkpoint 目录 |
--dataset_name | 数据集名称 | 通过datasets库从 Hub 加载,如cnn_dailymail、xsum、samsum等 |
--dataset_config | 数据集子配置名 | CNN/DailyMail 使用"3.0.0"版本(源码字段名为dataset_config_name,命令行可缩写为dataset_config) |
--output_dir | 输出目录 | 保存 checkpoint、评估结果与最终模型;目录非空且无 checkpoint 时会阻止训练(见下文断点续训) |
--per_device_train_batch_size | 每设备训练批大小 | 单卡/单 TPU 核上的批大小,总批大小会乘以设备数(num_replicas_in_sync) |
--per_device_eval_batch_size | 每设备评估批大小 | 同上,评估专用 |
--num_train_epochs | 训练轮数 | 默认 3.0 |
--do_train | 开启训练 | 需要数据集中存在train分片 |
--do_eval | 开启评估 | 需要数据集中存在validation分片,评估指标为 ROUGE |
如果只想做一次独立的评估(不训练),可只传--do_eval并配合--model_name_or_path指向一个已微调好的模型;脚本会走独立的 XLA 编译生成评估路径。
四、分布式策略:MirroredStrategy 与 TPU 的底层逻辑
README 提到“默认使用 MirroredStrategy,多 GPU 自动生效;传--tpu即可用 TPU”。这套自动选型逻辑实现在 training_args_tf.py 的_setup_strategy中,判定顺序如下:
- 若指定了
--tpu_name(即 README 中的--tpu),则通过TPUClusterResolver连接集群、初始化 TPU 系统,使用tf.distribute.TPUStrategy; - 无 TPU 且机器上没有 GPU(
len(gpus) == 0)时,回退到OneDeviceStrategy(device="/cpu:0"); - 只有 1 块 GPU 时,使用
OneDeviceStrategy(device="/gpu:0"); - 有多块 GPU 时,才启用
tf.distribute.MirroredStrategy()做数据并行。
注意两点工程细节:
- 只想用部分 GPU:源码注释明确提示用
CUDA_VISIBLE_DEVICES=0这类环境变量来限定可见设备,否则脚本会使用全部 GPU; - 混合精度:开启
--fp16后,GPU 场景会设置全局策略mixed_float16,而一旦检测到 TPU 则改为mixed_bfloat16(training_args_tf.py),与 TPU 硬件特性匹配。
脚本中实际使用的是training_args.strategy.scope()上下文(run_summarization.py),模型加载、tf.data构建、编译与fit都在该策略作用域内执行,从而保证多副本同步。
五、参数体系全景:三类 Dataclass 逐项解析
脚本用HfArgumentParser((ModelArguments, DataTrainingArguments, TFTrainingArguments))解析命令行(run_summarization.py),同时支持两种传入方式:逐个命令行参数,或一个 JSON 文件路径(脚本会根据参数个数自动识别并以parse_json_file解析)。
5.1 ModelArguments:模型 / 配置 / 分词器
| 参数 | 默认值 | 说明 |
|---|---|---|
--model_name_or_path | 必填 | 预训练模型路径或 Hub ID |
--config_name | None | 与模型名不同时,单独指定配置名/路径 |
--tokenizer_name | None | 与模型名不同时,单独指定分词器名/路径 |
--cache_dir | None | 预训练模型下载缓存目录 |
--use_fast_tokenizer | True | 是否使用基于tokenizers库的快速分词器 |
--model_revision | "main" | 模型版本(分支名、tag 或 commit id) |
--use_auth_token | False | 访问私有模型时传入huggingface-cli login生成的 token |
5.2 DataTrainingArguments:数据与序列长度控制
核心参数(对应 run_summarization.py):
| 参数 | 默认值 | 说明 |
|---|---|---|
--dataset_name/--dataset_config_name | None | 从 Hub 加载的数据集名与子配置 |
--text_column/--summary_column | None | 自定义源文本列与摘要列;缺省时从内置映射或数据列自动推断 |
--train_file/--validation_file/--test_file | None | 本地 CSV/JSON 数据文件(脚本断言扩展名必须是csv或json) |
--max_source_length | 1024 | 源文本 tokenize 后最大长度,超长截断、不足填充 |
--max_target_length | 128 | 目标摘要最大长度 |
--val_max_target_length | None | 验证集目标长度,缺省回退到max_target_length;同时覆盖model.generate的max_length |
--pad_to_max_length | False | 是否将所有样本填充到模型最大长度;False时按 batch 内最大长度动态填充(GPU 更高效,但不利于 TPU/XLA 形状缓存) |
--max_train_samples/--max_eval_samples/--max_predict_samples | None | 调试用截断样本数 |
--num_beams | None | 评估/预测时 beam search 的束宽,传入model.generate |
--ignore_pad_token_for_loss | True | 计算 loss 时是否忽略标签中的 padding token(内部替换为 -100) |
--source_prefix | None | 加在每个源文本前的提示前缀(T5 模型常用,如summarize:) |
--preprocessing_num_workers/--overwrite_cache | None/False | 预处理进程数与是否覆盖缓存 |
__post_init__中有两条强约束:必须提供dataset_name或训练/验证文件之一;val_max_target_length未指定时回退为max_target_length。
5.3 TFTrainingArguments:训练循环与分布式配置
该类继承自通用TrainingArguments并在 training_args_tf.py 中补充 TF 特有字段,常用项:
- 训练控制:
--output_dir、--do_train、--do_eval、--num_train_epochs(默认 3.0)、--max_steps、--per_device_train_batch_size(默认 8)、--per_device_eval_batch_size(默认 8)、--gradient_accumulation_steps(默认 1); - 优化器:
--learning_rate(默认 5e-5)、--weight_decay、--adam_beta1(0.9)、--adam_beta2(0.999)、--adam_epsilon(1e-8)、--max_grad_norm(1.0)、--warmup_steps、--warmup_ratio; - 保存与日志:
--save_strategy、--save_steps(500)、--save_total_limit、--logging_strategy、--logging_steps; - 分布式与硬件:
--tpu_name、--tpu_zone、--gcp_project、--xla(是否启用 XLA 编译)、--fp16、--no_cuda、--seed(默认 42); - 模型分享:
--push_to_hub、--push_to_hub_model_id、--push_to_hub_organization、--push_to_hub_token。
TFTrainingArguments还提供strategy、n_replicas、train_batch_size、eval_batch_size等只读属性,脚本中正是用strategy.num_replicas_in_sync计算全局总批大小:total_train_batch_size = per_device_train_batch_size * num_replicas(run_summarization.py)。
六、数据集加载与内置列名映射
脚本支持两种数据来源(run_summarization.py):
- Hub 数据集:传
--dataset_name(及可选--dataset_config_name),通过load_dataset自动下载; - 本地文件:传
--train_file/--validation_file/--test_file,按扩展名推断格式后用load_dataset(extension, data_files=...)加载。
对于常见摘要数据集,脚本内置了一份列名映射summarization_name_mapping(run_summarization.py),自动确定“源文本列”与“摘要列”:
| 数据集 | 源文本列 | 摘要列 |
|---|---|---|
cnn_dailymail | article | highlights |
xsum | document | summary |
samsum | dialogue | summary |
multi_news | document | summary |
big_patent | description | abstract |
amazon_reviews_multi | review_body | review_title |
orange_sum | text | summary |
pn_summary | article | summary |
psc | extract_text | summary_text |
thaisum | body | summary |
xglue | news_body | news_title |
wiki_summary | article | highlights |
不在映射表中的数据集则默认取第一个列为源文本、第二个列为摘要;若自定义数据集列名不同,务必用--text_column/--summary_column显式指定,脚本会对不存在的列名直接抛错。
七、数据预处理:tokenize 与标签构建
核心逻辑在preprocess_function(run_summarization.py):
def preprocess_function(examples): inputs = examples[text_column] targets = examples[summary_column] inputs = [prefix + inp for inp in inputs] model_inputs = tokenizer(inputs, max_length=data_args.max_source_length, padding=padding, truncation=True) labels = tokenizer(text_target=targets, max_length=max_target_length, padding=padding, truncation=True) if padding == "max_length" and data_args.ignore_pad_token_for_loss: labels["input_ids"] = [ [(l if l != tokenizer.pad_token_id else -100) for l in label] for label in labels["input_ids"] ] model_inputs["labels"] = labels["input_ids"] return model_inputs要点如下:
- 前缀拼接:
source_prefix会拼到每个源文本前,这是 T5 系模型的约定用法。脚本专门做了 T5 特判(run_summarization.py):若使用t5-small/base/large/3b/11b却未传--source_prefix,会打警告提示应加--source_prefix 'summarize: '; - 目标侧 tokenize:通过
tokenizer(..., text_target=targets)对摘要文本做独立编码,max_target_length控制长度; - padding 策略:
pad_to_max_length为True时用max_length填充,False时动态填充; - -100 掩码:启用
ignore_pad_token_for_loss时,标签中所有pad_token_id被替换为 -100,从而在交叉熵 loss 中忽略 padding 位置。
预处理通过datasets的.map(batched=True, num_proc=..., remove_columns=column_names, load_from_cache_file=...)批量执行,训练集与验证集分别处理,并用max_train_samples/max_eval_samples支持调试截断。
八、模型加载与 TF Dataset 构建
8.1 加载模型与检查配置
在training_args.strategy.scope()内,通过TFAutoModelForSeq2SeqLM.from_pretrained(...)加载模型(run_summarization.py),随后model.resize_token_embeddings(len(tokenizer))对齐词表。
脚本还强制校验model.config.decoder_start_token_id is not None,否则直接报错——解码器起始 token 是自回归生成的前提(例如 BART 配置中decoder_start_token_id=2,见 configuration_bart.py)。
8.2 DataCollatorForSeq2Seq
数据整理使用DataCollatorForSeq2Seq(run_summarization.py):
label_pad_token_id = -100 if data_args.ignore_pad_token_for_loss else tokenizer.pad_token_id data_collator = DataCollatorForSeq2Seq( tokenizer, model=model, label_pad_token_id=label_pad_token_id, pad_to_multiple_of=128, # Reduce the number of unique shapes for XLA, especially for generation return_tensors="tf", )pad_to_multiple_of=128的注释点明了关键:把 batch 内所有序列长度对齐到 128 的倍数,可大幅减少 XLA 需要编译的输入形状数量,尤其对生成阶段(beam search + 自回归)的编译缓存非常友好。
8.3 prepare_tf_dataset 与自动分片
随后用model.prepare_tf_dataset()把 Hugging Face Dataset 包成tf.data.Dataset(run_summarization.py),并设置tf.data.experimental.AutoShardPolicy.OFF关闭自动分片,避免分布式训练中数据集的 shard 策略干扰:
dataset_options = tf.data.Options() dataset_options.experimental_distribute.auto_shard_policy = tf.data.experimental.AutoShardPolicy.OFF tf_train_dataset = model.prepare_tf_dataset( train_dataset, collate_fn=data_collator, batch_size=total_train_batch_size, shuffle=True, ).with_options(dataset_options)源码注释指出:prepare_tf_dataset能自动从模型输入名推断列名,比底层的to_tf_dataset()更省心,是 Keras 训练推荐方式。
九、优化器与学习率调度
训练步数与 warmup 的计算逻辑(run_summarization.py):
num_train_steps = int(len(tf_train_dataset) * training_args.num_train_epochs) if training_args.warmup_steps > 0: num_warmup_steps = training_args.warmup_steps elif training_args.warmup_ratio > 0: num_warmup_steps = int(num_train_steps * training_args.warmup_ratio) else: num_warmup_steps = 0 optimizer, lr_schedule = create_optimizer( init_lr=training_args.learning_rate, num_train_steps=num_train_steps, num_warmup_steps=num_warmup_steps, adam_beta1=training_args.adam_beta1, adam_beta2=training_args.adam_beta2, adam_epsilon=training_args.adam_epsilon, weight_decay_rate=training_args.weight_decay, adam_global_clipnorm=training_args.max_grad_norm, )warmup 的优先级是warmup_steps优先、其次warmup_ratio;create_optimizer会生成带线性 warmup + 衰减的 Adam 优化器与配套学习率调度器,max_grad_norm通过adam_global_clipnorm实现全局梯度裁剪。若只评估不训练(do_eval且无do_train),则optimizer = None。
十、ROUGE 评估与 KerasMetricCallback 实现
10.1 生成参数与文本后处理
评估阶段加载evaluate.load("rouge"),构造生成参数(run_summarization.py):
gen_kwargs = { "max_length": data_args.val_max_target_length, "num_beams": data_args.num_beams, "no_repeat_ngram_size": 0, # Not supported under XLA right now }no_repeat_ngram_size被强制设为 0,注释说明当前 XLA 下不支持该约束,而部分模型配置默认开启它。
ROUGE 的rougeLSum变体要求每个句子之间以换行分隔,因此postprocess_text(run_summarization.py)会先用nltk.sent_tokenize分句,再以"\n".join(...)重组预测与参考文本:
def postprocess_text(preds, labels): preds = [pred.strip() for pred in preds] labels = [label.strip() for label in labels] preds = ["\n".join(nltk.sent_tokenize(pred)) for pred in preds] labels = ["\n".join(nltk.sent_tokenize(label)) for label in labels] return preds, labels10.2 compute_metrics 与指标汇总
compute_metrics(run_summarization.py)负责解码与打分:预测 token 用tokenizer.batch_decode(predictions, skip_special_tokens=True)解码;标签中的 -100 先还原为pad_token_id再解码;最后metric.compute(..., use_stemmer=True)计算 ROUGE,并只取各指标的mid.fmeasure * 100保留两位小数。
10.3 KerasMetricCallback 的作用与参数
由于 ROUGE 需要字符串比较和生成循环,无法写成可被 TF 编译的普通 Keras 指标,脚本引入KerasMetricCallback(run_summarization.py):
metric_callback = KerasMetricCallback( metric_fn=compute_metrics, eval_dataset=tf_eval_dataset, predict_with_generate=True, use_xla_generation=True, generate_kwargs=gen_kwargs, )该回调的实现位于 keras_callbacks.py,其 docstring 明确指出:回调在每个 epoch 结束时,先在eval_dataset上执行预测/生成,再把结果以np.ndarray形式传给metric_fn。几个关键参数:
metric_fn:接收(predictions, labels),返回“指标名 → 数值”字典;predict_with_generate:是否用model.generate()产出结果;use_xla_generation:是否用 XLA 编译生成过程。源码注释称这可以带来“最高约 100 倍”的生成加速,但每种输入形状都需要一次新的 XLA 编译,因此建议配合pad_to_multiple_of或固定长度 padding 减少形状数量;generate_kwargs:透传给model.generate的关键字参数。
一个典型的metric_fn返回形如{'rouge1': 37.4199, 'rouge2': 13.9768, 'rougeL': 34.361, 'rougeLsum': 35.0781},与任何 Keras 指标一样记录进训练历史。
10.4 独立评估路径
若只评估不训练(do_eval且未do_train),脚本走独立评估分支(run_summarization.py):把生成函数包装为@tf.function(jit_compile=True)以获取 XLA 加速,逐 batch 生成、解码、metric.add_batch,最后输出mid.fmeasure * 100的指标字典。
十一、训练、断点续训与模型导出
11.1 编译与 fit
模型编译与训练(run_summarization.py):
model.compile(optimizer=optimizer, jit_compile=training_args.xla) history = model.fit(tf_train_dataset, epochs=int(training_args.num_train_epochs), callbacks=callbacks)jit_compile=training_args.xla把--xla透传给 Keras 编译。脚本会提示:启用 XLA 但未设--pad_to_max_length时,前期因需编译各种输入形状可能较慢,属正常现象。
11.2 断点检测
训练前用get_last_checkpoint(output_dir)检查输出目录(run_summarization.py):
- 若目录已存在且非空、且找不到 checkpoint,直接抛错,提示用
--overwrite_output_dir覆盖或换输出目录; - 若发现 checkpoint 且未传
--resume_from_checkpoint,则自动从最近 checkpoint 续训并打日志。
11.3 结果与模型导出
训练/评估结束后(run_summarization.py):
- 评估指标(ROUGE 各分项)写入
{output_dir}/all_results.json; - 未启用
--push_to_hub时,model.save_pretrained(output_dir)保存本地副本。
11.4 推送 Hub(可选)
启用--push_to_hub后,脚本追加PushToHubCallback(run_summarization.py),每个 epoch 保存一次并自动生成模型卡片;默认模型 ID 为{模型名}-finetuned-{数据集名},model_card_kwargs会带上finetuned_from、tasks: "summarization"以及数据集标签/配置信息,便于模型卡片在 Hub 上被正确展示与检索。
十二、在 FlexGen 仓库中的定位与扩展用法
本示例位于 FlexGen 仓库的基准测试目录benchmark/third_party/transformers/examples/tensorflow/summarization/,是仓库收录的 Hugging Face Transformers 参考实现之一,可作为摘要类 Seq2Seq 任务微调与评估的基线脚本直接复用。围绕它你可以做三类扩展:
- 换模型:
--model_name_or_path t5-base并加--source_prefix 'summarize: ',或换 Pegasus、ProphetNet 等任意TFAutoModelForSeq2SeqLM支持的模型; - 换数据:直接换
--dataset_name(利用内置列名映射),或用--train_file/--validation_file加载自有 CSV/JSON 数据并配合--text_column/--summary_column; - 适配硬件:多卡默认 MirroredStrategy;有 TPU 时传
--tpu <资源名>(配合--tpu_zone/--gcp_project);需要 XLA 加速则加--xla并建议同时开启--pad_to_max_length。
结语
这份看似简短的 README 背后,是一套从分布式策略自动选型、数据集列名映射、-100 掩码标签构建,到 ROUGE 生成评估、断点续训与模型导出的完整 TensorFlow 摘要微调流水线。本文以 README.md 的示例命令为起点,结合 run_summarization.py 及 training_args_tf.py、keras_callbacks.py 等源码逐层还原了其实现细节。读者既可照抄命令快速跑通基线,也能依据各节参数表与源码定位按需定制,让摘要模型微调真正做到开箱即用、按需可改。
- 推理引擎
- 大模型
【免费下载链接】FlexGen
Running large language models on a single GPU for throughput-oriented scenarios.
相关推荐
使用 Transformers 与 Seq2SeqTrainer 微调 T5 实现文本摘要
使用 Transformers 与 Seq2SeqTrainer 微调 T5 实现文本摘要 本指南基于 🤗 Transformers 官方文档中的 Summa
人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态Transformers 摘要生成实战指南:基于 T5 微调 BillSum 法律文本摘要模型
Transformers 摘要生成实战指南:基于 T5 微调 BillSum 法律文本摘要模型 摘要生成(Summarization)是 🤗 Transfor
人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态Transformers 脚本化训练实战:用 run_summarization.py 微调文本摘要模型(分布式、TPU、Accelerate 与自定义数据集全流程)
Transformers 脚本化训练实战:用 run_summarization.py 微调文本摘要模型(分布式、TPU、Accelerate 与自定义数据集全
人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考