news 2026/9/25 15:36:28

使用 [特殊字符] Transformers 在 TensorFlow 中微调摘要模型:run_summarization 脚本实战与源码解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
使用 [特殊字符] Transformers 在 TensorFlow 中微调摘要模型:run_summarization 脚本实战与源码解析
  • 推理引擎
  • 大模型

【免费下载链接】FlexGen

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

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

本篇技术指南以 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):

  1. check_min_version("4.24.0"):Transformers 版本低于 4.24.0 直接报错;
  2. require_version("datasets>=1.8.0", ...):datasets库低于 1.8.0 会提示安装;
  3. 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中,判定顺序如下:

  1. 若指定了--tpu_name(即 README 中的--tpu),则通过TPUClusterResolver连接集群、初始化 TPU 系统,使用tf.distribute.TPUStrategy;
  2. 无 TPU 且机器上没有 GPU(len(gpus) == 0)时,回退到OneDeviceStrategy(device="/cpu:0");
  3. 只有 1 块 GPU 时,使用OneDeviceStrategy(device="/gpu:0");
  4. 有多块 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_nameNone与模型名不同时,单独指定配置名/路径
--tokenizer_nameNone与模型名不同时,单独指定分词器名/路径
--cache_dirNone预训练模型下载缓存目录
--use_fast_tokenizerTrue是否使用基于tokenizers库的快速分词器
--model_revision"main"模型版本(分支名、tag 或 commit id)
--use_auth_tokenFalse访问私有模型时传入huggingface-cli login生成的 token

5.2 DataTrainingArguments:数据与序列长度控制

核心参数(对应 run_summarization.py):

参数默认值说明
--dataset_name/--dataset_config_nameNone从 Hub 加载的数据集名与子配置
--text_column/--summary_columnNone自定义源文本列与摘要列;缺省时从内置映射或数据列自动推断
--train_file/--validation_file/--test_fileNone本地 CSV/JSON 数据文件(脚本断言扩展名必须是csv或json)
--max_source_length1024源文本 tokenize 后最大长度,超长截断、不足填充
--max_target_length128目标摘要最大长度
--val_max_target_lengthNone验证集目标长度,缺省回退到max_target_length;同时覆盖model.generate的max_length
--pad_to_max_lengthFalse是否将所有样本填充到模型最大长度;False时按 batch 内最大长度动态填充(GPU 更高效,但不利于 TPU/XLA 形状缓存)
--max_train_samples/--max_eval_samples/--max_predict_samplesNone调试用截断样本数
--num_beamsNone评估/预测时 beam search 的束宽,传入model.generate
--ignore_pad_token_for_lossTrue计算 loss 时是否忽略标签中的 padding token(内部替换为 -100)
--source_prefixNone加在每个源文本前的提示前缀(T5 模型常用,如summarize:)
--preprocessing_num_workers/--overwrite_cacheNone/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):

  1. Hub 数据集:传--dataset_name(及可选--dataset_config_name),通过load_dataset自动下载;
  2. 本地文件:传--train_file/--validation_file/--test_file,按扩展名推断格式后用load_dataset(extension, data_files=...)加载。

对于常见摘要数据集,脚本内置了一份列名映射summarization_name_mapping(run_summarization.py),自动确定“源文本列”与“摘要列”:

数据集源文本列摘要列
cnn_dailymailarticlehighlights
xsumdocumentsummary
samsumdialoguesummary
multi_newsdocumentsummary
big_patentdescriptionabstract
amazon_reviews_multireview_bodyreview_title
orange_sumtextsummary
pn_summaryarticlesummary
pscextract_textsummary_text
thaisumbodysummary
xgluenews_bodynews_title
wiki_summaryarticlehighlights

不在映射表中的数据集则默认取第一个列为源文本、第二个列为摘要;若自定义数据集列名不同,务必用--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, labels

10.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 任务微调与评估的基线脚本直接复用。围绕它你可以做三类扩展:

  1. 换模型:--model_name_or_path t5-base并加--source_prefix 'summarize: ',或换 Pegasus、ProphetNet 等任意TFAutoModelForSeq2SeqLM支持的模型;
  2. 换数据:直接换--dataset_name(利用内置列名映射),或用--train_file/--validation_file加载自有 CSV/JSON 数据并配合--text_column/--summary_column;
  3. 适配硬件:多卡默认 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.

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

相关推荐

上一篇:Figma 界面怎么变中文?FigmaCN 汉化插件 10 分钟上手,从此告别翻译软件
下一篇:三步彻底移除 Windows Defender:windows-defender-remover 快速上手与避坑指南

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/25 15:27:37

把资深 BA 装进团队:BA Master 工程化实战手册(TaoToken 配置篇)

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/25 15:24:42

寒武纪PyTorch理事会席位背后:AI芯片软件栈适配与算子实现全解析

1. 从“同桌”这个词说起&#xff1a;一个信号背后的技术分量“寒武纪拿下PyTorch最高席位&#xff0c;与英伟达同桌”——这个标题我第一次看到的时候&#xff0c;正在调一个模型训练脚本&#xff0c;手边跑着的是一台装了消费级显卡的机器。说实话&#xff0c;第一反应不是兴…

作者头像 李华
网站建设 2026/9/25 15:19:04

OI Wiki 离线版怎么部署:3 条路线选 1 条就够

OI Wiki 离线版怎么部署&#xff1a;3 条路线选 1 条就够 【免费下载链接】OI-wiki :star2: Wiki of OI / ICPC for everyone. &#xff08;某大型游戏线上攻略&#xff0c;内含炫酷算术魔法&#xff09; 项目地址: https://gitcode.com/GitHub_Trending/oi/OI-wiki 机房…

作者头像 李华
网站建设 2026/9/25 15:10:34

多协议以太网温湿度变送器在楼宇自控中的选型与部署要点

作为在楼宇自控行业摸爬滚打十几年的老工程师&#xff0c;这几年感触最深的变化&#xff0c;就是温湿度变送器这类最基础的传感设备&#xff0c;正在从“能测就行”快速转向“能接入、能互通、能联动”。越来越多的智慧楼宇项目在采购温湿度变送器时&#xff0c;已经不满足于传…

作者头像 李华