恒美微站 Logo 恒美微站
  • 首页
  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心
  • 联系我们

基于 JAX/Flax 微调序列到序列摘要模型:Transformers `run_summarization_flax.py` 实战指南

  • 首页
  • 资讯中心
  • /
  • 基于 JAX/Flax 微调序列到序列摘要模型:Transformers `run_summarization_flax.py` 实战指南

相关资讯

WeiXinMPSDK 微信 Native 支付实战指南:从付款码 URL 签名、二维码生成到回调统一下单 2026/9/25 3:54:41
CTF压缩包爆破实战:从密码原理到hashcat与掩码攻击全解析 2026/9/25 3:54:41
抖音视频去水印批量下载实战:开源工具douyin-downloader配置与踩坑指南 2026/9/25 3:49:41

最新资讯

Math Dice 数学骰子:Basic Computer Games 中的加法可视化训练程序全解析
Atlas 300V 24G推理卡部署YOLO全流程与踩坑指南
DiceBear Pixel Art 风格预设(Preset)完全指南:10 套开箱即用的渲染选项与源码机制解析
学生宿舍管理系统:从建库到前端的完整数据流闭环实现
Atlas 300V 24G实战:从PyTorch到OM的YOLO推理部署全流程指南
Atlas 300V 24G推理加速卡部署实战:从YOLO模型转换到ACL调优

今日推荐

AI元人文:从工具使用到思维重构的深度探索
Python+CNN车牌识别实战:从数据预处理到模型训练与部署
Vim基础操作全攻略:保存退出、模式切换与高频命令实战

本周热门

BrewUI:给Homebrew套上图形界面,让macOS软件包管理更简单
BrewUI:让Homebrew包管理变得可视化与高效
公式与文本对齐全攻略:从Word到LaTeX的实用技巧

本月精选

自研推理加速器Redwood:两周内实现PyTorch模型高效部署的实战教程
V4L2摄像头采集实战:从camera_client.rar到出图全流程解析
从“谁发明了钢琴键”到知识问答智能体:RAG与记忆工程实践

基于 JAX/Flax 微调序列到序列摘要模型:Transformers `run_summarization_flax.py` 实战指南

发布时间:2026/9/25 3:54:41
基于 JAX/Flax 微调序列到序列摘要模型:Transformers `run_summarization_flax.py` 实战指南 推理引擎大模型【免费下载链接】FlexGenRunning large language models on a single GPU for throughput-oriented scenarios.项目地址https://gitcode.com/gh_mirrors/fl/FlexGen点击查看免费下载本文以 Transformers Flax 摘要微调示例为骨架完整讲解如何用run_summarization_flax.py在 GPU/TPU 上对 BART、T5、Pegasus 等序列到序列Seq2Seq模型进行摘要任务的端到端微调。你将掌握完整的训练/评估/预测命令行用法、三类参数模型、数据、训练的默认值与语义并从源码层面理解 JAX/Flax 的函数式训练循环、分布式pmap并行、ROUGE 指标计算与 Model Hub 推送机制可直接复用到自己的摘要乃至其他 Seq2Seq 任务中。该文档位于 benchmark/third_party/transformers/examples/flax/summarization/README.md配套的完整训练脚本为 run_summarization_flax.py。该目录属于当前仓库 benchmark/third_party 下随附的 Transformers 框架源码作为开源参考实现随仓库分发。为什么用 JAX/Flax 做摘要微调摘要Summarization是典型的序列到序列任务输入长文档、输出短摘要。JAX/Flax 技术栈在这一场景下的核心优势体现在JAX暴露 NumPy 风格的 API并具备强大的转换能力jit可以把纯函数 trace 后编译为 GPU/TPU 上高效、融合的加速代码grad求任意梯度、pmap多设备并行、remat梯度检查点、vmap自动向量化、pjit自动分片的模型并行且这些转换可以任意组合。Flax在 JAX 之上提供基于 dataclass 的模块抽象代码简洁显式其lifted转换如vmap、remat允许任意嵌套。函数式与不可变性JAX/Flax 模型不可变以纯函数方式更新天然适合pmap级别的简单高效模型并行。需要特别说明的是这一时期的 Flax 示例没有 Trainer 抽象所有训练循环都是显式写在脚本中的参见 examples/flax/README.md因此阅读本脚本的源码也是学习 JAX/Flax 训练范式的最佳入口之一。环境准备与依赖安装脚本运行依赖的最小集合记录在 requirements.txtdatasets 1.1.3 jax0.2.8 jaxlib0.1.59 flax0.3.5 optax0.0.8 evaluate0.2.0若还要运行仓库附带的示例测试test_flax_examples.py需要额外安装 _tests_requirements.txt 中列出的pytest、nltk、rouge-score、seqeval、tensorboard、conllu等包。安装 JAX 本身需按运行环境区分GPUJAX 的 pip 安装与 CUDA/CuDNN 版本强相关需根据本机 CUDA 版本选择对应的jaxlib安装方式。TPUJAX/Flax 官方示例均以在 Cloud TPU 上高效运行为设计目标多设备并行开箱即用。另外脚本在启动时会自动检测并下载 NLTK 的punkt分词数据用于 ROUGE 计算前的句子切分离线环境设置了TRANSFORMERS_OFFLINE下会直接抛出提示要求先联网完成下载。支持哪些模型架构脚本通过FlaxAutoModelForSeq2SeqLM自动加载模型。从源码映射表 modeling_flax_auto.pyFLAX_MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING_NAMES可以确认本脚本支持以下 Seq2Seq 架构model_type对应 Flax 模型类bartFlaxBartForConditionalGenerationblenderbotFlaxBlenderbotForConditionalGenerationblenderbot-smallFlaxBlenderbotSmallForConditionalGenerationencoder-decoderFlaxEncoderDecoderModellongt5FlaxLongT5ForConditionalGenerationmarianFlaxMarianMTModelmbartFlaxMBartForConditionalGenerationmt5FlaxMT5ForConditionalGenerationpegasusFlaxPegasusForConditionalGenerationt5FlaxT5ForConditionalGeneration也就是说README 命令示例中的 BART 只是其中之一你也可以直接替换为t5-small、google/pegasus-xsum等 checkpoint。模型加载流程为AutoConfig.from_pretrained→AutoTokenizer.from_pretrained→FlaxAutoModelForSeq2SeqLM.from_pretrained或from_config从头训练权重精度由--dtype控制float32/float16/bfloat16默认float32。训练命令与预期结果README 核心示例README 给出的完整训练命令如下可以直接复制运行python run_summarization_flax.py \ --output_dir ./bart-base-xsum \ --model_name_or_path facebook/bart-base \ --tokenizer_name facebook/bart-base \ --dataset_namexsum \ --do_train --do_eval --do_predict --predict_with_generate \ --num_train_epochs 6 \ --learning_rate 5e-5 --warmup_steps 0 \ --per_device_train_batch_size 64 \ --per_device_eval_batch_size 64 \ --overwrite_output_dir \ --max_source_length 512 --max_target_length 64 \ --push_to_hub按 README 记载该配置在 6 个 epoch 后约37 分钟完成训练得到验证集 loss1.7785、ROUGE217.01训练统计可在 TensorBoard.dev 上查看。需要说明的是这里使用的是generate的默认生成参数README 特别提醒若针对xsum数据集的特点如合适的 beam 数、最小生成长度等显式配置生成参数ROUGE 分数可以进一步提升。上述耗时与指标与当时的硬件环境、依赖版本相关在不同机器上复现时会有差异不应视为固定基准。三类命令行参数详解对应源码 dataclass脚本使用HfArgumentParser解析三类参数源码中分别定义为ModelArguments、DataTrainingArguments、TrainingArguments三个 dataclass。除命令行外还支持把唯一参数写成 JSON 配置文件路径脚本会自动调用parse_json_file读取。模型参数ModelArguments参数默认值说明--model_name_or_pathNone预训练 checkpoint 名称或路径不设置则从头训练--model_typeNone从头训练时指定模型类型上述 10 种之一--config_nameNone与模型名不同的 config 名称/路径--tokenizer_nameNone与模型名不同的分词器名称/路径--cache_dirNone预训练模型下载缓存目录--use_fast_tokenizerTrue是否使用 tokenizers 库的快速分词器--dtypefloat32权重初始化与训练的浮点格式float32/float16/bfloat16--use_auth_tokenFalse加载私有模型时使用huggingface-cli login生成的 token数据参数DataTrainingArguments参数默认值说明--dataset_nameNone使用 Datasets Hub 上的数据集名称--dataset_config_nameNone数据集的 configuration 名称如多语言数据集的语言子集--text_column/--summary_columnNone自定义数据集中正文/摘要列名不指定则自动推断--train_file/--validation_file/--test_fileNone本地训练/验证/测试文件仅支持json或csv--max_source_length1024源文本最大 token 数超长截断、不足填充--max_target_length128目标摘要最大 token 数--val_max_target_lengthNone验证/预测时的目标长度默认取max_target_length同时会覆盖model.generate的max_length--max_train_samples/--max_eval_samples/--max_predict_samplesNone调试用截断样本数以加速--preprocessing_num_workersNone预处理进程数--source_prefixNone加在每条源文本前的前缀T5 类模型常用--predict_with_generateFalse是否用generate计算生成式指标ROUGE/BLEU--num_beamsNone评估时 beam 数会传给model.generate默认取模型 config--overwrite_cacheFalse覆盖预处理缓存训练参数TrainingArguments参数默认值说明--output_dir必填模型预测与 checkpoint 输出目录--overwrite_output_dirFalse输出目录已存在且非空时是否覆盖也可用于从 checkpoint 目录续训--do_train/--do_eval/--do_predictFalse是否执行训练/评估/预测--per_device_train_batch_size8每个 GPU/TPU core/CPU 上的训练 batch--per_device_eval_batch_size8每个设备上的评估 batch--learning_rate5e-5AdamW 初始学习率--weight_decay0.0AdamW 权重衰减--adam_beta1/--adam_beta2/--adam_epsilon0.9 / 0.999 / 1e-8AdamW 优化器超参--label_smoothing_factor0.0标签平滑系数0 表示不启用--adafactorFalse是否用 Adafactor 替代 AdamW--num_train_epochs3.0总训练 epoch 数--warmup_steps0线性预热步数--logging_steps500每 N 步输出一次日志--save_steps500每 N 步保存 checkpoint--eval_stepsNone每 N 步做一次评估--seed42随机种子--push_to_hubFalse训练后是否上传模型到 Model Hub--hub_model_idNoneHub 仓库全名含用户名/组织名如yourname/bart-base-xsum--hub_tokenNone推送 Hub 用的 token--gradient_checkpointingFalse梯度检查点以更慢的反向传播换取显存节省注意总 batch size 是每设备 batch × 设备数源码中train_batch_size per_device_train_batch_size * jax.device_count()脚本会自动使用检测到的全部 GPU/TPU core分布式训练开箱即用。数据加载与预处理从 Hub 或本地 jsonlines/csv脚本支持两条数据路径Hub 数据集指定--dataset_name如xsum通过load_dataset自动下载。本地文件指定--train_file/--validation_file/--test_file扩展名必须是csv或json源码中有断言校验load_dataset按扩展名自动解析。列名自动推断源码维护了一张summarization_name_mapping字典run_summarization_flax.py为常见数据集预置了正文/摘要列名数据集正文列摘要列cnn_dailymailarticlehighlightsxsumdocumentsummarysamsumdialoguesummarybig_patentdescriptionabstractxgluenews_bodynews_titleorange_sum/pn_summary/psc/thaisum/wiki_summary/amazon_reviews_multi见源码映射见源码映射若未命中映射或使用自定义文件默认取数据集第一列为正文、第二列为摘要也可用--text_column/--summary_column显式指定。预处理细节源码preprocess_function由于 jitted 函数需要固定长度输入预处理对源/目标统一使用paddingmax_length补齐到max_source_length/max_target_length并截断。针对 Flax 模型不接受labels的特性脚本通过模型模块的shift_tokens_right函数把目标序列右移一位生成decoder_input_ids同时保留decoder_attention_mask用于在损失中屏蔽 pad token。训练循环内部实现JAX/Flax 微调原理拆解阅读 run_summarization_flax.py 的训练部分可以完整还原 JAX/Flax 微调的典型范式训练状态自定义TrainState继承 Flaxtrain_state.TrainState额外携带dropout_rngreplicate()将参数复制到所有设备并把 dropout 随机键按设备分片shard_prng_key。学习率调度create_learning_rate_fn用optax.linear_schedule构造线性预热 线性衰减再通过optax.join_schedules拼接总步数 数据集大小 // 总batch × epoch数。权重衰减掩码decay_mask_fn遍历参数树识别名称中含layernorm/layer_norm/ln的 LayerNorm 参数及所有bias对它们不施加权重衰减——这是 AdamW 的标准最佳实践。优化器optax.adamw配合上述学习率调度与衰减掩码。损失函数loss_fn实现了带标签平滑的交叉熵one-hot 软化标签confidence 1 - label_smoothing_factor并用decoder_attention_mask屏蔽 padding token最后对非 pad 位置求均值归一化。梯度更新train_step内用jax.value_and_grad(compute_loss, has_auxTrue)求梯度用jax.lax.psum跨设备规约损失与样本数再归一化后apply_gradients更新状态。多设备并行jax.pmap(..., batch)把 train/eval/generate 步编译成 SPMD 并行程序训练数据经shard分发到各设备评估与生成使用pad_shard_unpad自动补齐不完整 batch、卸载结果。这就是 JAX纯函数 任意组合转换特性的直接体现。数据加载data_loader用jax.random.permutation生成随机 batch 索引drop_lastTrue时跳过不完整的末尾 batch评估时则保留。评估与生成ROUGE 指标计算流程当指定--predict_with_generate时评估/预测循环除了计算 loss还会执行生成并计算 ROUGE生成参数gen_kwargs由val_max_target_length或模型 config 的max_length与num_beams或模型 config 的num_beams构成通过model.generate(batch[input_ids], attention_mask..., **gen_kwargs)批量解码。指标计算compute_metrics先用 tokenizer 解码预测与标签skip_special_tokensTrue再用nltk.sent_tokenize把每段文本按句切分换行ROUGE-LSum 要求句子间换行最后evaluate.load(rouge).compute(..., use_stemmerTrue)得到各 ROUGE 分数并乘以 100同时记录平均生成长度gen_len。输出预测阶段结束后主进程把指标写入output_dir/test_results.json键名如test_rouge1、test_rouge2、test_rougeL、test_rougeLsum。日志、Checkpoint 与 Model Hub 推送TensorBoard主进程jax.process_index() 0用 FlaxSummaryWriter写入训练/评估标量train_loss、eval_loss、各 ROUGE 值等write_metric负责落盘。Checkpoint每个 epoch 结束后主进程从state.params取出首设备副本调用model.save_pretrained(output_dir, paramsparams)与tokenizer.save_pretrained(output_dir)保存模型与分词器。Hub 推送--push_to_hub会基于output_dir目录名或--hub_model_id指定的全名创建/克隆远端仓库每个 epoch 以Saving weights and logs of epoch N为提交信息异步推送。使用前需要本地登录huggingface-cli login或通过--hub_token传入认证 token。仓库如何验证该示例可运行test_flax_examples.py 中的test_run_summarization用t5-small配合tests/fixtures/tests_samples/xsum/sample.json小样本做冒烟测试num_train_epochs3、warmup_steps8、learning_rate2e-4、per_device_train_batch_size2、per_device_eval_batch_size1并断言test_rouge1 10、test_rouge2 2、test_rougeL 7、test_rougeLsum 7。这既验证了脚本端到端可运行也给出了小数据快速验证的参数参考——先用极小样本跑通流程再切换到完整数据集。自定义数据集jsonlines/csv 快速上手如果你有自己的摘要数据可按以下方式组织jsonlines每行一个 JSON 对象包含正文与摘要两个字段如{document: ..., summary: ...}。csv第一列为正文、第二列为摘要或通过--text_column/--summary_column指定列名。启动示例python run_summarization_flax.py \ --model_name_or_path facebook/bart-base \ --train_file ./data/train.json \ --validation_file ./data/val.json \ --test_file ./data/test.json \ --text_column document --summary_column summary \ --do_train --do_eval --do_predict --predict_with_generate \ --num_train_epochs 3 \ --per_device_train_batch_size 16 \ --output_dir ./my-bart-summarizer \ --overwrite_output_dir运行注意事项小结output_dir已存在且非空、同时启用了--do_train且未加--overwrite_output_dir时会直接报错终止这是防止误覆盖已有 checkpoint 的保护机制。模型 config 必须正确设置decoder_start_token_id否则脚本会在启动时抛错。max_source_length/max_target_length决定显存占用与训练速度摘要任务常见配置为源 5121024、目标 64128README 示例即 512/64。若目标是把模型发布到 Hub建议显式指定--hub_model_id保证仓库名符合用户名/模型名规范。文中涉及的文件均为仓库只读参考实现运行与配置均在本地完成无需修改仓库内容。赞分享推理引擎大模型【免费下载链接】FlexGenRunning large language models on a single GPU for throughput-oriented scenarios.项目地址https://gitcode.com/gh_mirrors/fl/FlexGen点击查看免费下载相关推荐Transformers 中的 PEGASUS-X面向长文本摘要的序列到序列模型实战指南Transformers 中的 PEGASUS X面向长文本摘要的序列到序列模型实战指南 PEGASUS X 是 Google 在 2022 年发布、并已完整人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态Transformers 摘要生成实战指南基于 T5 微调 BillSum 法律文本摘要模型Transformers 摘要生成实战指南基于 T5 微调 BillSum 法律文本摘要模型 摘要生成Summarization是 Transfor人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态Transformers 中的越南语大规模序列到序列模型BARTpho 架构、分词原理与文本摘要实战Transformers 中的越南语大规模序列到序列模型BARTpho 架构、分词原理与文本摘要实战 BARTpho 是面向越南语Vietnamese的大人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态上一篇终极系统备份与恢复指南使用Rescuezilla保护你的数据安全下一篇5步掌握Steam API打造个性化游戏数据解决方案创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于恒美微站

恒美微站专注于为个体商户、工作室提供极简自助建站服务,让每个人都能轻松拥有专业网站。

快速链接

  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心

服务项目

  • 可视化建站
  • 拖拽编辑
  • 主题定制
  • SEO 优化
  • 网站托管

联系方式

  • 📍 地址:北京市朝阳区建国路 88 号
  • 📞 电话:400-888-8888
  • ✉️ 邮箱:info@hmyw.cn
  • 🕐 时间:周一至周日 9:00-18:00

© 2024 恒美微站 hmyw.cn 版权所有 | 京 ICP 备 12345678 号