如何在HuggingFace Trainer类中设置恒定学习率?含get_constant_schedule用法
在HuggingFace Trainer中配置恒定学习率(使用get_constant_schedule)
恒定学习率指训练全程保持初始学习率不变,无任何衰减逻辑。以下是用transformers库的get_constant_schedule实现该配置的完整步骤:
1. 基础组件准备
先加载模型、分词器并处理数据集(以IMDB分类任务为例):
from transformers import AutoModelForSequenceClassification, AutoTokenizer, Trainer, TrainingArguments from transformers import get_constant_schedule from datasets import load_dataset import torch # 加载预训练模型和分词器 model = AutoModelForSequenceClassification.from_pretrained("bert-base-uncased", num_labels=2) tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased") # 处理数据集 def tokenize_func(examples): return tokenizer(examples["text"], padding="max_length", truncation=True) dataset = load_dataset("imdb") tokenized_ds = dataset.map(tokenize_func, batched=True) train_ds = tokenized_ds["train"].shuffle(seed=42).select(range(1000)) # 取小批量数据演示 eval_ds = tokenized_ds["test"].shuffle(seed=42).select(range(1000))
2. 手动创建优化器与恒定调度器
先定义带初始学习率的优化器,再用get_constant_schedule生成无衰减调度器:
# 初始化优化器,指定初始学习率 optimizer = torch.optim.AdamW(model.parameters(), lr=2e-5) # 生成恒定学习率调度器 scheduler = get_constant_schedule(optimizer)
3. 配置Trainer并启动训练
将自定义的优化器和调度器传入Trainer,同时设置训练参数:
training_args = TrainingArguments( output_dir="./constant_lr_results", per_device_train_batch_size=8, num_train_epochs=3, logging_dir="./logs", logging_steps=10, evaluation_strategy="epoch" ) # 初始化Trainer,传入自定义的优化器和调度器 trainer = Trainer( model=model, args=training_args, train_dataset=train_ds, eval_dataset=eval_ds, optimizers=(optimizer, scheduler) ) # 开始训练 trainer.train()
4. 可选:验证学习率是否恒定
添加回调函数,每10步打印当前学习率,确认数值始终等于初始值:
from transformers import TrainerCallback class LRTracker(TrainerCallback): def on_step_end(self, args, state, control, **kwargs): if state.global_step % 10 == 0: current_lr = kwargs["optimizer"].param_groups[0]["lr"] print(f"Step {state.global_step}: Current LR = {current_lr}") # 给Trainer添加回调 trainer.add_callback(LRTracker())
替代方案:使用scheduler_func简化配置
如果不想手动创建优化器,可让Trainer自动生成优化器,通过scheduler_func指定调度器生成逻辑:
def constant_scheduler_factory(optimizer): return get_constant_schedule(optimizer) training_args = TrainingArguments( output_dir="./constant_lr_results", per_device_train_batch_size=8, num_train_epochs=3, logging_dir="./logs", logging_steps=10, learning_rate=2e-5 # 指定初始学习率 ) trainer = Trainer( model=model, args=training_args, train_dataset=train_ds, eval_dataset=eval_ds, scheduler_func=constant_scheduler_factory ) trainer.train()
关键说明
get_constant_schedule的核心作用是返回一个无衰减的调度器,全程维持优化器的初始学习率。- 若直接使用
TrainingArguments的lr_scheduler_type="constant"也能实现恒定学习率,但手动调用get_constant_schedule更灵活,方便后续扩展调度器逻辑。
内容的提问来源于stack exchange,提问作者Ramraj Chandradevan
相关产品推荐
相关产品推荐

