python transformers笔记(TrainingArguments类)

TrainingArguments类

TrainingArguments是Hugging Face Transformers库中用于集中管理超参数和配置的核心类。它定义了模型训练、评估、保存和日志记录的所有关键参数,并通过Trainer类实现自动化训练流程。

1、核心功能

(1)统一管理训练配置:替代手动定义学习率、批次大小等分散的超参数

(2)自动化流程配置:控制训练、评估、保存、日志记录等行为的触发策略

(3)支持混合精度训练(FP16/BP16)、多GPU/TPU分布式训练、梯度积累等

2、关键参数

(1)基础训练配置

参数名 作用 示例 默认值
output_dir 模型和日志的保存路径 "./results" "tmp_trainer"
num_train_epochs 训练总轮次 3 3.0
per_device_train_batch_size 每个GPU/CPU的批次大小 16 8
per_device_eval_batch_size 评估时的批次大小 32 8
learning_rate 初始学习率 2e-5 5e-5
weight_decay 权重衰减(L2正则化)系数 0.01 0.0

(2)训练策略控制

参数名 作用 示例 默认值
gradient_accumulation_steps 梯度累积步数(模拟更大批次) 4 1
evaluation_strategy 评估触发策略:"epoch"/"steps"/"no" "epoch" "no"
save_strategy 模型保存策略(需要与evaluation_strategy一致) "epoch" "steps"
logging_strategy 日志记录策略 "steps" "steps"
logging_steps 每多少步记录一次日志(当logging_strategy="steps"时生效) 100 500
save_steps 每多少步保存一次模型 100 500
eval_steps 每多少步评估一次模型 100 None

(3)硬件与性能优化

参数名 作用 示例 默认值
fp16 是否启用FP16混合精度训练(NVIDIA GPU) True False
bf16 是否启用BF16混合精度训练(AMD/Intel GPU/TPU) False False
optim 优化器类型(如"adamw_torch"、"adafactor") "adamw_torch" "adamw_hf"
dataloader_num_workers 数据加载的线程数 4 0
gradient_checkpointing 是否启用梯度检查点(节省显存,但降低速度) False False

(4)模型保存或加载

参数名 作用 示例 默认值
save_total_limit 最大保存的检查点数量(旧的会被删除) 3 None
load_best_model_at_end 训练结束后是否加载最佳模型(需启用评估) True False
metric_for_best_model 用于选择最佳模型的指标(如"eval_loss"或自定义指标) "eval_accuracy" None
great_is_better 指标是否越大越好(需与metric_for_best_model配合) None None

(5)日志与监控

参数名 作用 示例 默认值
report_to 日志上报工具(如"tensorboard"、"wandb") "tensorboard"
logging_dir TensorBoard日志记录 "./logs"
disable_tqdm 是否禁用进度条 False
log_level 日志级别("debug"、"warning"、"info"等) "debug" "passive"

3、使用示例

python 复制代码
from transformers import TrainingArguments

training_args = TrainingArguments(
    output_dir="./results",
    per_device_train_batch_size=8,
    per_device_eval_batch_size=16,
    num_train_epochs=3,
    evaluation_strategy="epoch",
    save_strategy="epoch",
    learning_rate=5e-5,
    weight_decay=0.01,
    fp16=True,
    load_best_model_at_end=True,
    metric_for_best_model="eval_loss",
    logging_dir="./logs",
    report_to=["tensorboard"],
)

4、高级用法

(1)自定义优化器和调度器

(2)动态参数调整

(3)超参数搜索

5、配置文件支持

python 复制代码
# 保存配置
training_args.save_to_json("./training_args.json")

# 加载配置
new_args = TrainingArguments.from_json_file("./training_args.json")

6、注意事项

(1)策略一致性

save_strategy和evaluation_strategy必须相同(除非设置为"no")

(2)混合精度选择

NVIDIA GPU用fp16,AMD/Intel GPU或TPU用bf16

(3)内存优化

遇到OOM错误时,减小per_device_train_batch_size或增加gradient_accumulation_steps。

相关推荐
外收内放17 分钟前
Python基础语法练习题(59-原始版本与优化版本)
开发语言·python
ASS-ASH20 分钟前
从力学与机械原理角度深度剖析:尊界V800刹车踏板支架断裂事件的技术逻辑
人工智能·python·数据可视化·机械工程·汽车安全·物理仿真·制动系统
xhy_070721 分钟前
AI 读了 .env,密钥会原样发给大模型吗?WES Code 的密钥打码(附实测)
人工智能·python·数据挖掘·flask·ai编程·wes code
泡海椒28 分钟前
JQuick-Excel 实战:用 STYLE 配置单元格与区域样式
开发语言·python·excel
威联通网络存储34 分钟前
TS-h1283XU-RP 在医药批记录与检验场景的部署
网络·python
Patrick在香港43 分钟前
时间戳凭空早了 8 小时:datetime.utcnow() 弃用实测与漂移复盘
数据库·python·标准库·datetime·时区·弃用
东莞市云毅网络有限公司7 小时前
AI 引用句逐条回指原文:让回答可回溯的校验实现
python·数据清洗·rag·企业知识库·文档解析
数字融合9 小时前
透明化视频三维矿山井下照明重建技术
人工智能·python·数码相机
yi0119 小时前
LeetCode 219:存在重复元素 II——哈希表记录“最近一次出现的位置”
数据结构·人工智能·笔记·python·算法·leetcode·哈希表
Marst Code10 小时前
上位机开发日记 · 第 2 篇 · 架构先行:六层分层与边界
python