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

如何在XGBoost中实现类似ReduceLROnPlateau的学习率衰减功能

XGBoost实现类似TensorFlow ReduceLROnPlateau的动态降学习率功能

XGBoost本身没有内置同名回调,你可以通过自定义训练回调实现完全一致的效果,当验证集指标停止提升时自动下调学习率:

实现方法

  • 自定义回调适配XGBoost原生训练接口和Scikit-learn包装接口,可直接复用以下代码:
import xgboost as xgb
import numpy as np

class ReduceLROnPlateau(xgb.callback.TrainingCallback):
    def __init__(self, metric: str, patience: int = 5, factor: float = 0.5, min_lr: float = 1e-7, verbose: int = 1):
        """
        验证集指标无提升时自动降低学习率的回调
        :param metric: 要监控的评估指标,需和训练时指定的评估指标完全一致
        :param patience: 容忍指标无提升的训练轮数,超过该轮数则触发学习率下调
        :param factor: 学习率衰减系数,每次调整后新学习率 = 旧学习率 * factor
        :param min_lr: 学习率下限,低于该值不再调整
        :param verbose: 是否打印学习率调整日志,1为打印,0为不打印
        """
        self.metric = metric
        self.patience = patience
        self.factor = factor
        self.min_lr = min_lr
        self.verbose = verbose
        # 初始化最优指标,损失/误差类越小越好,AUC/准确率类越大越好
        self.best_score = np.inf if any(k in metric for k in ('loss', 'error', 'mae', 'mse')) else -np.inf
        self.wait_count = 0
        self.current_lr = None

    def before_training(self, model):
        self.current_lr = model.get_param('learning_rate')
        return model

    def after_iteration(self, model, epoch, evals_log):
        # 获取当前轮验证集指标值
        current_score = evals_log['validation_0'][self.metric][-1]
        # 判断指标是否提升
        is_improved = False
        if any(k in self.metric for k in ('loss', 'error', 'mae', 'mse')):
            if current_score < self.best_score:
                is_improved = True
        else:
            if current_score > self.best_score:
                is_improved = True

        if is_improved:
            self.best_score = current_score
            self.wait_count = 0
        else:
            self.wait_count += 1
            if self.wait_count >= self.patience:
                new_lr = max(self.current_lr * self.factor, self.min_lr)
                if new_lr < self.current_lr:
                    self.current_lr = new_lr
                    model.set_param('learning_rate', self.current_lr)
                    if self.verbose:
                        print(f"第{epoch}轮训练:验证集{self.metric}无提升,调整学习率为{self.current_lr:.8f}")
                self.wait_count = 0
        # 返回False表示继续训练,返回True会终止训练
        return False
  • 调用示例(原生XGBoost接口):
# 构造数据集
dtrain = xgb.DMatrix(X_train, label=y_train)
dval = xgb.DMatrix(X_val, label=y_val)

# 训练参数
params = {
    "objective": "binary:logistic",
    "learning_rate": 0.1,
    "max_depth": 6,
    "eval_metric": "error" # 这里指定的评估指标要和回调里的metric参数一致
}

# 启动训练,传入自定义回调
model = xgb.train(
    params,
    dtrain,
    num_boost_round=1000,
    evals=[(dval, "validation_0")],
    callbacks=[ReduceLROnPlateau(metric="error", patience=3, factor=0.3, min_lr=1e-6)]
)
  • 调用示例(Scikit-learn接口):
from xgboost import XGBClassifier

model = XGBClassifier(
    objective="binary:logistic",
    learning_rate=0.1,
    max_depth=6,
    n_estimators=1000,
    eval_metric="error"
)

model.fit(
    X_train, y_train,
    eval_set=[(X_val, y_val)],
    callbacks=[ReduceLROnPlateau(metric="error", patience=3, factor=0.3)]
)

注意事项

  • 回调中监控的metric参数必须和训练时指定的eval_metric完全一致
  • 可以和XGBoost自带的EarlyStopping回调组合使用,同时实现学习率调整和早停功能
  • 学习率衰减系数建议设置在0.1~0.5之间,避免学习率下降过快导致模型欠拟合

内容的提问来源于stack exchange,提问作者Sticky

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 07:42:00