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

如何在Optuna中为随机森林回归模型实现剪枝?

问题

我正在开发机器学习模型,使用Optuna进行超参数调优,希望尝试剪枝功能,但不知如何实现。目前我使用RandomForestRegressor,其余功能运行正常,现有目标函数代码如下:

def objective(trial):
    n_estimators = trial.suggest_int('n_estimators', 100, 1000)
    max_depth = trial.suggest_int('max_depth', 5, 50)
    min_samples_split = trial.suggest_int('min_samples_split', 2, 30)
    min_samples_leaf = trial.suggest_int('min_samples_leaf', 1, 10)
    max_samples = trial.suggest_float('max_samples', 0.5, 1.0)
    max_features = trial.suggest_int('max_features', 5, 30)
    max_leaf_nodes = trial.suggest_int('max_leaf_nodes', 100, 200)

    model = RandomForestRegressor(n_estimators=n_estimators,
                              max_depth=max_depth,
                              min_samples_split=min_samples_split,
                              min_samples_leaf=min_samples_leaf,
                              max_samples=max_samples,
                              max_features=max_features,
                              max_leaf_nodes=max_leaf_nodes)

    kFold = KFold(n_splits=5)
    scores = cross_val_score(model, X_train_transformed, y_train, cv=kFold, scoring='r2', n_jobs=-1)
    mean_score = np.mean(scores)

    return mean_score


study = optuna.create_study(direction = 'maximize',
                        sampler=optuna.samplers.TPESampler(multivariate=True))
study.optimize(objective, n_trials=300)

请问如何为我的目标函数实现剪枝?

实现剪枝的步骤

Optuna剪枝需要在训练过程中定期向trial报告中间结果,让剪枝器判断当前trial是否有继续的价值。由于cross_val_score无法中途输出结果,需要拆分交叉验证流程,具体操作如下:

  1. 导入剪枝依赖:引入Optuna的剪枝器和剪枝异常类。
  2. 手动拆分交叉验证:遍历KFold的每个折,训练后记录分数并报告给trial。
  3. 添加剪枝判断:每次报告后检查是否需要终止当前trial,若需要则抛出剪枝异常。
  4. 关联剪枝器到Study:创建Study时指定剪枝器,定义剪枝规则。
修改后的完整代码
import optuna
from optuna.pruners import MedianPruner
from optuna.exceptions import TrialPruned
from sklearn.ensemble import RandomForestRegressor
from sklearn.model_selection import KFold
import numpy as np

def objective(trial):
    # 超参数采样
    n_estimators = trial.suggest_int('n_estimators', 100, 1000)
    max_depth = trial.suggest_int('max_depth', 5, 50)
    min_samples_split = trial.suggest_int('min_samples_split', 2, 30)
    min_samples_leaf = trial.suggest_int('min_samples_leaf', 1, 10)
    max_samples = trial.suggest_float('max_samples', 0.5, 1.0)
    max_features = trial.suggest_int('max_features', 5, 30)
    max_leaf_nodes = trial.suggest_int('max_leaf_nodes', 100, 200)

    model = RandomForestRegressor(n_estimators=n_estimators,
                                  max_depth=max_depth,
                                  min_samples_split=min_samples_split,
                                  min_samples_leaf=min_samples_leaf,
                                  max_samples=max_samples,
                                  max_features=max_features,
                                  max_leaf_nodes=max_leaf_nodes)

    kFold = KFold(n_splits=5)
    scores = []
    
    for fold_idx, (train_idx, val_idx) in enumerate(kFold.split(X_train_transformed, y_train)):
        # 拆分当前折的训练/验证集
        X_fold_train, X_fold_val = X_train_transformed[train_idx], X_train_transformed[val_idx]
        y_fold_train, y_fold_val = y_train[train_idx], y_train[val_idx]
        
        # 训练并计算当前折分数
        model.fit(X_fold_train, y_fold_train)
        score = model.score(X_fold_val, y_fold_val)
        scores.append(score)
        
        # 向trial报告当前折的结果,step标记当前是第几个折
        trial.report(score, step=fold_idx)
        
        # 判断是否需要剪枝,是则抛出异常终止当前trial
        if trial.should_prune():
            raise TrialPruned()
    
    mean_score = np.mean(scores)
    return mean_score


# 创建Study时绑定剪枝器,n_warmup_steps表示前2个折不触发剪枝,给模型基础训练空间
study = optuna.create_study(
    direction='maximize',
    sampler=optuna.samplers.TPESampler(multivariate=True),
    pruner=MedianPruner(n_warmup_steps=2)
)
study.optimize(objective, n_trials=300)
关键说明
  • 剪枝器选择:MedianPruner是回归任务的常用选项,会终止表现低于已完成trial中位数的任务;也可尝试SuccessiveHalvingPruner或HyperbandPruner,后者更适合大规模调优场景。
  • n_warmup_steps参数:设置为2是为了避免模型还未完成基础训练就被剪枝,可根据交叉验证折数调整。
  • 手动交叉验证的必要性:必须替换cross_val_score为手动遍历,才能在每一步输出中间结果,这是实现剪枝的核心前提。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 15:56:31