You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.06 18:55:06