如何在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
相关产品推荐
相关产品推荐

