尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

Transformers 摘要生成(Summarization)微调实战:run_summarization.py 与 run_summarization_no_trainer.py 完整流程解析

发布时间:2026/9/25 2:17:19

资讯中心
01
ARTICLE

Transformers 摘要生成(Summarization)微调实战:run_summarization.py 与 run_summarization_no_trainer.py 完整流程解析

Transformers 摘要生成(Summarization)微调实战:run_summarization.py 与 run_summarization_no_trainer.py 完整流程解析
推理引擎大模型【免费下载链接】FlexGenRunning large language models on a single GPU for throughput-oriented scenarios.项目地址https://gitcode.com/gh_mirrors/fl/FlexGen点击查看免费下载本文是围绕 HuggingFace Transformers 官方示例中 Summarization 任务的实战技术指南。它基于本仓库内 Summarization 示例目录 的说明文档展开系统讲解如何利用run_summarization.py与run_summarization_no_trainer.py两个脚本在 CNN/DailyMail、XSum 等标准数据集或自备 CSV/JSONLINES 数据上对 BART、Pegasus、T5 等序列到序列模型进行摘要生成任务的微调与评估。读完本文你将掌握两种训练范式Trainer高层封装与Accelerate裸训练循环的完整命令行用法、关键参数语义、数据格式约定以及 ROUGE 评估的实现原理。一、示例定位它在当前仓库中的角色在 FlexGen 仓库中benchmark/third_party/transformers是为基准测试维护的 HuggingFace Transformers v4.24.0 分支安装方式见 benchmark/third_party/README.md在目录下执行pip3 install -e .并安装accelerate0.15.0。本文聚焦的 Summarization 示例位于run_summarization.py基于Trainer的高层封装脚本共 732 行run_summarization_no_trainer.py基于Accelerate的裸训练循环脚本共 759 行requirements.txt运行所需依赖清单README.md官方使用说明本文主体骨架。两个脚本均通过check_min_version(4.24.0)与require_version(datasets1.8.0)做版本校验因此建议在 Transformers ≥ 4.24.0、datasets ≥ 1.8.0 的环境下运行。原 README 中已废弃的bertabs与旧版finetune_trainer.py相关内容不在本仓库示例范围内本文不展开。二、支持的模型架构run_summarization.py通过AutoModelForSeq2SeqLM自动加载模型官方支持以下条件生成conditional generation架构架构说明BartForConditionalGenerationBART广泛用于摘要与生成FSMTForConditionalGeneration仅用于翻译场景fairseq 机器翻译MBartForConditionalGeneration多语言 BART需要指定--lang与--forced_bos_tokenMarianMTModelMarian 机器翻译模型PegasusForConditionalGenerationGoogle Pegasus专为摘要设计T5ForConditionalGenerationT5 文本到文本统一框架需--source_prefixMT5ForConditionalGeneration多语言 T5在源码层面run_summarization.py脚本通过AutoModelForSeq2SeqLM.from_pretrained加载权重AutoConfig加载配置、AutoTokenizer加载分词器。若model_name_or_path包含.ckpt后缀会自动以from_tfTrue从 TensorFlow checkpoint 转换run_summarization.py。三、环境与依赖准备requirements.txt 列出完整依赖accelerate datasets 1.8.0 sentencepiece ! 0.1.92 protobuf rouge-score nltk py7zr torch 1.3 evaluate其中py7zr用于解压 CNN/DailyMail 等以 7z 压缩包分发的数据集rouge-score与nltk用于 ROUGE 指标计算sentencepiece与protobuf是 T5、mBART、Pegasus 等模型分词器所需accelerate供无 Trainer 脚本使用。此外nltk的punkt分词数据会在脚本首次运行时自动下载脚本内置了FileLock与离线模式处理见 run_summarization.py。若使用本仓库内的 Transformers 分支安装命令为见 benchmark/third_party/README.mdcd benchmark/third_party/transformers pip3 install -e . pip3 install accelerate0.15.0四、使用 Trainer 微调run_summarization.py4.1 最小可运行示例以 T5-small 在 CNN/DailyMail 3.0.0 配置上微调为例对应 README.md 中的官方命令脚本路径按本仓库调整为仓库根目录相对路径python benchmark/third_party/transformers/examples/pytorch/summarization/run_summarization.py \ --model_name_or_path t5-small \ --do_train \ --do_eval \ --dataset_name cnn_dailymail \ --dataset_config 3.0.0 \ --source_prefix summarize: \ --output_dir /tmp/tst-summarization \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --overwrite_output_dir \ --predict_with_generate命令逐项解读--model_name_or_path t5-small指定预训练模型HuggingFace Hub 模型 ID 或本地路径--do_train/--do_eval开启训练与验证阶段脚本要求至少传入其一否则提示 There is nothing to do 并退出见 run_summarization.py--dataset_name cnn_dailymail --dataset_config 3.0.0从 datasets Hub 加载 CNN/DailyMail 数据集及 3.0.0 版本配置--source_prefix summarize: T5 系模型必须使用的任务前缀--output_dir模型与日志输出目录--overwrite_output_dir允许覆盖已有输出目录不传时若目录非空且未检测到 checkpoint脚本会直接报错见 run_summarization.py--predict_with_generate评估与预测阶段使用model.generate做自回归解码若省略compute_metrics不会被挂载到 Trainer 上见 run_summarization.py。4.2 T5 的 source_prefix 约定只有 T5 系列模型t5-small、t5-base、t5-large、t5-3b、t5-11b必须附加--source_prefix summarize: 参数。源码在加载模型后专门做了检查若未提供source_prefix且模型为上述 T5 之一会输出警告见 run_summarization.py。前缀会在预处理阶段拼接到每条源文本前inputs [prefix inp for inp in inputs]run_summarization.py。T5 模型在预训练阶段使用带任务前缀的格式微调时保持一致的输入格式是取得良好效果的前提而 BART/Pegasus 等模型不需要此前缀。4.3 切换数据集CNN/DailyMail 与 XSum原 README 特别说明XSumExtreme Summarization是另一个常用的摘要数据集。只需将数据集参数替换python benchmark/third_party/transformers/examples/pytorch/summarization/run_summarization.py \ --model_name_or_path t5-small \ --do_train \ --do_eval \ --dataset_name xsum \ --source_prefix summarize: \ --output_dir /tmp/tst-summarization \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --overwrite_output_dir \ --predict_with_generate之所以换数据集时无需手动指定文本/摘要列名是因为脚本内置了summarization_name_mapping字典run_summarization.py自动映射各数据集对应的字段cnn_dailymail→(article, highlights)xsum→(document, summary)samsum→(dialogue, summary)big_patent→(description, abstract)multi_news→(document, summary)amazon_reviews_multi→(review_body, review_title)以及orange_sum、pn_summary、psc、thaisum、xglue、wiki_summary等映射逻辑在 run_summarization.py若用户未传--text_column/--summary_column优先使用映射中的列名否则退回数据集第一个/第二个字段。4.4 使用自己的数据文件CSV / JSONLINES原 README 明确指出摘要任务支持自定义 CSV 与 JSONLINES 两种格式。使用自备数据时将--dataset_name替换为--train_file、--validation_filepython benchmark/third_party/transformers/examples/pytorch/summarization/run_summarization.py \ --model_name_or_path t5-small \ --do_train \ --do_eval \ --train_file path_to_csv_or_jsonlines_file \ --validation_file path_to_csv_or_jsonlines_file \ --source_prefix summarize: \ --output_dir /tmp/tst-summarization \ --overwrite_output_dir \ --per_device_train_batch_size4 \ --per_device_eval_batch_size4 \ --predict_with_generateCSV 文件训练与验证文件应各含一列原文、一列摘要。若 CSV 只有两列如下示例第一列默认为text、第二列为summarytext,summary Im sitting here in a boring room. Its just another rainy Sunday afternoon. Im wasting my time I got nothing to do. Im hanging around Im waiting for you. But nothing ever happens. And I wonder,Im sitting in a room where Im waiting for something to happen I see trees so green, red roses too. I see them bloom for me and you. And I think to myself what a wonderful world. I see skies so blue and clouds so white. The bright blessed day, the dark sacred night. And I think to myself what a wonderful world.,Im a gardener and Im a big fan of flowers. Christmas time is here. Happiness and cheer. Fun for all that children call. Their favorite time of the year. Snowflakes in the air. Carols everywhere. Olden times and ancient rhymes. Of love and dreams to share,Its that time of year again.若 CSV 包含多列如id,date,text,summary需显式指定要使用的列--text_column text \ --summary_column summary \JSONLINES 文件第二种格式为每行一个 JSON 对象{text: Im sitting here in a boring room. Its just another rainy Sunday afternoon. Im wasting my time I got nothing to do. Im hanging around Im waiting for you. But nothing ever happens. And I wonder, summary: Im sitting in a room where Im waiting for something to happen} {text: I see trees so green, red roses too. I see them bloom for me and you. And I think to myself what a wonderful world. I see skies so blue and clouds so white. The bright blessed day, the dark sacred night. And I think to myself what a wonderful world., summary: Im a gardener and Im a big fan of flowers.} {text: Christmas time is here. Happiness and cheer. Fun for all that children call. Their favorite time of the year. Snowflakes in the air. Carols everywhere. Olden times and ancient rhymes. Of love and dreams to share, summary: Its that time of year again.}与 CSV 相同默认取第一个键为原文、第二个键为摘要因此键名可以任意示例用了text与summary也可用--text_column text --summary_column summary显式指定。文件格式的合法性在脚本参数校验阶段被强制检查train_file与validation_file的扩展名必须是csv或json否则断言失败run_summarization.py。加载时脚本依据扩展名调用load_dataset(csv|json, data_files...)run_summarization.py。4.5 预处理与数据整理源码解读预处理函数preprocess_functionrun_summarization.py的核心逻辑过滤掉原文或摘要为空的样本在每条原文前拼接source_prefix用tokenizer(inputs, max_lengthmax_source_length, paddingpadding, truncationTrue)对原文做截断/填充用tokenizer(text_targettargets, ...)即text_target关键字参数独立对摘要做 tokenize若采用定长填充--pad_to_max_length且ignore_pad_token_for_lossTrue将标签中的pad_token_id替换为-100使 padding 部分不参与损失计算。相关参数默认值定义于 run_summarization.py 的DataTrainingArguments参数默认值说明--max_source_length1024输入原文最大长度超长截断、不足填充--max_target_length128训练时摘要标签最大长度--val_max_target_length跟随max_target_length验证/预测时的目标长度同时覆盖model.generate的max_length--pad_to_max_lengthFalse是否定长填充False 时按 batch 内最大长度动态填充GPU 更高效TPU 上不推荐--ignore_pad_token_for_lossTrue是否在损失中忽略 padding 标签替换为 -100--num_beamsNone评估/预测时model.generate的 beam 数--source_prefix加在每条原文前的任务前缀--preprocessing_num_workersNone数据预处理进程数--overwrite_cacheFalse是否覆盖预处理缓存--max_train_samples/--max_eval_samples/--max_predict_samplesNone调试用截取样本子集加快实验--lang/--forced_bos_tokenNone多语言模型mBART 等所需--dataset_config_nameNone数据集配置名如 3.0.0模型侧参数ModelArguments见 run_summarization.py--config_name、--tokenizer_name、--cache_dir、--use_fast_tokenizer默认 True、--model_revision默认 main、--use_auth_token默认 False访问私有模型时配合huggingface-cli login使用、--resize_position_embeddings当max_source_length超过模型位置编码数时是否自动扩展见 run_summarization.py。对 mBART 等多语言模型脚本还会校验decoder_start_token_id若缺失则按--lang从分词器映射设置并要求--forced_bos_token指定首生成 token 为目标语言 tokenrun_summarization.py。4.6 数据收集器与 ROUGE 评估数据整理阶段使用DataCollatorForSeq2Seqrun_summarization.py标签 padding 默认用-100配合ignore_pad_token_for_lossFP16 训练时按 8 的倍数对齐pad_to_multiple_of8。评估指标为 ROUGE实现要点run_summarization.py通过evaluate.load(rouge)加载指标postprocess_text用nltk.sent_tokenize将预测与参考按句子分行——rougeLSum 要求每句后带换行符这是 ROUGE-L 变体计算的格式前提metric.compute(predictions..., references..., use_stemmerTrue)计算 rouge1/rouge2/rougeL/rougeLsum结果乘以 100 并保留 4 位小数额外统计gen_len生成序列平均长度。训练结束后脚本会在--output_dir下保存模型与分词器并生成generated_predictions.txt--do_predict且--predict_with_generate时见 run_summarization.py若不传--push_to_hub则调用trainer.create_model_card()生成模型卡片run_summarization.py。五、使用 Accelerate 微调run_summarization_no_trainer.py5.1 与 Trainer 版本的区别run_summarization_no_trainer.py同样支持上述全部架构与数据集核心区别在于它暴露了完整的裸训练循环方便快速实验和任意定制如直接修改优化器或 DataLoader 配置。它牺牲了一部分Trainer的内置选项但通过Accelerate库天然支持分布式训练、TPU 与混合精度。官方 README 建议先安装 Acceleratepip install accelerate对应本仓库 benchmark/third_party/README.md 中指定的accelerate0.15.0版本。5.2 直接运行python benchmark/third_party/transformers/examples/pytorch/summarization/run_summarization_no_trainer.py \ --model_name_or_path t5-small \ --dataset_name cnn_dailymail \ --dataset_config 3.0.0 \ --source_prefix summarize: \ --output_dir ~/tmp/tst-summarization5.3 通过 accelerate 启动器运行该脚本的优势在于支持多种运行环境先交互式生成配置accelerate config回答引导问题后可用accelerate test验证环境是否就绪然后启动训练accelerate launch benchmark/third_party/transformers/examples/pytorch/summarization/run_summarization_no_trainer.py \ --model_name_or_path t5-small \ --dataset_name cnn_dailymail \ --dataset_config 3.0.0 \ --source_prefix summarize: \ --output_dir ~/tmp/tst-summarization同一条命令即可适配以下全部环境原 README 明确列出纯 CPU 环境单 GPU 环境多 GPU 分布式训练单节点或多节点TPU 训练。5.4 裸训练循环的实现要点从源码看run_summarization_no_trainer.py该脚本的关键设计参数解析使用标准argparseparse_args见 run_summarization_no_trainer.py与 Trainer 版本共享大部分数据参数max_source_length1024、max_target_length128、text_column、summary_column、source_prefix等训练参数默认值为per_device_train_batch_size8、learning_rate5e-5、num_train_epochs3、gradient_accumulation_steps1、lr_scheduler_typelinear、num_warmup_steps0Accelerator 初始化Accelerator(gradient_accumulation_steps...)统一管理设备与梯度累积run_summarization_no_trainer.py优化器分组将参数按是否需要权重衰减分成两组——bias、LayerNorm.weight、layer_norm.weight不衰减其余参数应用--weight_decay优化器为AdamWrun_summarization_no_trainer.py学习率调度通过get_scheduler生成调度器支持linear、cosine、cosine_with_restarts、polynomial、constant、constant_with_warmup六种类型并按梯度累积步数换算 warmup 与总步数run_summarization_no_trainer.py断点续训--resume_from_checkpoint支持从step_{n}/epoch_{n}目录恢复若未指定路径则自动选取最近目录run_summarization_no_trainer.py评估每个 epoch 结束后以model.generate(max_lengthval_max_target_length, num_beams...)生成摘要经pad_across_processes、gather_for_metrics跨进程对齐后计算 ROUGErun_summarization_no_trainer.py最终结果写入output_dir/all_results.json包含eval_rouge1/rouge2/rougeL/rougeLsum四项指标run_summarization_no_trainer.py可选的实验追踪--with_tracking配合--report_to支持 tensorboard / wandb / comet_ml默认 all记录损失与指标。5.5 加速参数速查参数默认值说明--per_device_train_batch_size8每个设备训练 batch 大小--per_device_eval_batch_size8每个设备评估 batch 大小--learning_rate5e-5初始学习率--weight_decay0.0权重衰减系数--num_train_epochs3训练轮数被--max_train_steps覆盖时失效--max_train_stepsNone总训练步数优先级高于轮数--gradient_accumulation_steps1梯度累积步数--lr_scheduler_typelinear调度器类型六选一--num_warmup_steps0学习率预热步数--checkpointing_stepsNone每 N 步或每 epoch 保存一次状态--resume_from_checkpointNone断点续训目录--push_to_hub/--hub_model_id/--hub_token-训练中实时推送模型到 Hub--with_tracking/--report_toall实验指标追踪--use_slow_tokenizerFalse是否使用慢速分词器六、总结围绕 Summarization 任务本仓库中的两个脚本提供了两条互补的技术路线run_summarization.pyTrainer 路线开箱即用、选项丰富适合快速落地标准流程——加载数据Hub 数据集或 CSV/JSONLINES、预处理前缀、截断、-100 标签屏蔽、Seq2SeqTrainer训练、ROUGE 评估与模型卡片生成全部内置run_summarization_no_trainer.pyAccelerate 路线训练循环完全透明优化器分组、调度器、断点续训、多进程指标聚合均可在脚本内直接修改配合accelerate config / test / launch一条命令打通 CPU、单卡、多卡、TPU 全场景。两套脚本共享同一套数据约定text/summary列映射、CSV/JSONLINES 格式、source_prefix规则与评估逻辑ROUGE nltk.sent_tokenize分行、use_stemmerTrue、结果放大 100 倍读者可根据对训练过程控制粒度的需求任选其一并参照本文的参数速查表完成定制。若需深入了解 Transformer 内部实现可继续阅读本仓库中 tasks/summarization.mdx 等官方文档源文件。赞分享推理引擎大模型【免费下载链接】FlexGenRunning large language models on a single GPU for throughput-oriented scenarios.项目地址https://gitcode.com/gh_mirrors/fl/FlexGen点击查看免费下载相关推荐Transformers 文本摘要实战指南基于 Seq2Seq 架构用 T5 完成抽象式摘要微调全流程Transformers 文本摘要实战指南基于 Seq2Seq 架构用 T5 完成抽象式摘要微调全流程 本篇技术指南围绕 Transformers 仓库的摘要人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态Transformers 摘要生成实战指南基于 T5 微调 BillSum 法律文本摘要模型Transformers 摘要生成实战指南基于 T5 微调 BillSum 法律文本摘要模型 摘要生成Summarization是 Transfor人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态Transformers 脚本化训练实战用 run_summarization.py 微调文本摘要模型分布式、TPU、Accelerate 与自定义数据集全流程Transformers 脚本化训练实战用 run_summarization.py 微调文本摘要模型分布式、TPU、Accelerate 与自定义数据集全人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态上一篇探秘MiniOS一款轻量级、便携式的操作系统构建工具下一篇探索未来UI设计augmented-ui 项目推荐创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

◈

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

◐

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

▲

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。