超参数调优进阶:Optuna/Bayesian/Early Stopping

超参数调优进阶:Optuna/Bayesian/Early Stopping

1. 调优方法对比

复制代码
超参数调优方法:
├── 网格搜索(Grid Search):穷举所有组合,慢但全面
├── 随机搜索(Random Search):随机采样,快但不保证最优
├── 贝叶斯优化(Bayesian):基于历史结果智能搜索
└── 早停法(Early Stopping):训练中动态停止

2. Optuna 调优

python 复制代码
import optuna
from sklearn.ensemble import RandomForestClassifier
from sklearn.model_selection import cross_val_score

def objective(trial):
    params = {
        'n_estimators': trial.suggest_int('n_estimators', 50, 300),
        'max_depth': trial.suggest_int('max_depth', 3, 15),
        'min_samples_split': trial.suggest_int('min_samples_split', 2, 20),
        'min_samples_leaf': trial.suggest_int('min_samples_leaf', 1, 10),
        'max_features': trial.suggest_categorical('max_features', ['sqrt', 'log2', None]),
    }
    
    model = RandomForestClassifier(**params, random_state=42)
    scores = cross_val_score(model, X_train, y_train, cv=5, scoring='accuracy')
    return scores.mean()

study = optuna.create_study(direction='maximize')
study.optimize(objective, n_trials=100, show_progress_bar=True)

print(f"最佳参数: {study.best_params}")
print(f"最佳分数: {study.best_value:.4f}")

3. XGBoost + Optuna

python 复制代码
import optuna
import xgboost as xgb

def objective_xgb(trial):
    params = {
        'n_estimators': trial.suggest_int('n_estimators', 50, 500),
        'max_depth': trial.suggest_int('max_depth', 3, 12),
        'learning_rate': trial.suggest_float('learning_rate', 0.01, 0.3, log=True),
        'subsample': trial.suggest_float('subsample', 0.6, 1.0),
        'colsample_bytree': trial.suggest_float('colsample_bytree', 0.6, 1.0),
        'reg_alpha': trial.suggest_float('reg_alpha', 1e-8, 10.0, log=True),
        'reg_lambda': trial.suggest_float('reg_lambda', 1e-8, 10.0, log=True),
    }
    
    model = xgb.XGBClassifier(**params, random_state=42, use_label_encoder=False)
    scores = cross_val_score(model, X_train, y_train, cv=5, scoring='accuracy')
    return scores.mean()

study = optuna.create_study(direction='maximize')
study.optimize(objective_xgb, n_trials=200)

4. Early Stopping

python 复制代码
import lightgbm as lgb

train_data = lgb.Dataset(X_train, label=y_train)
val_data = lgb.Dataset(X_val, label=y_val, reference=train_data)

params = {
    'objective': 'binary',
    'metric': 'binary_logloss',
    'learning_rate': 0.05,
    'num_leaves': 31,
}

callbacks = [
    lgb.early_stopping(stopping_rounds=50),
    lgb.log_evaluation(period=10),
]

model = lgb.train(
    params, train_data,
    valid_sets=[val_data],
    num_boost_round=1000,
    callbacks=callbacks,
)

总结

方法 速度 精度 推荐场景
Grid Search 小参数空间
Random Search 快速探索
Optuna 复杂参数空间
Early Stopping 训练中使用
相关推荐
JarmanYuo几秒前
YOLO 涨点研究(六):网络结构改进之小目标增强篇1——无人机视角下的车辆与行人检测
人工智能·pytorch·python·yolo·计算机视觉·无人机
格林威3 分钟前
C# 相机图像阴影校正:使用OpenCvSharp实现工业相机阴影平场校正功能
人工智能·数码相机·opencv·计算机视觉·c#·机器视觉·工业相机
桃西西呀6 分钟前
9月1日起AI图要被标记了,机器怎么一眼认出哪张是 AI 画的?
人工智能·机器学习·llm
xian_wwq8 分钟前
【学习笔记】深度认知系列-第10讲AI绘画与设计——从Stable Diffusion到Midjourney
笔记·学习·ai作画
凯尔萨厮9 分钟前
Java学习笔记十五(GUI)
java·笔记·学习
招财小梗10 分钟前
沈阳AI企业咨询可定制数字化方案吗?
大数据·人工智能·python
其实防守也摸鱼11 分钟前
教育信息技术应用创新---基础软件信息赛(题库)
大数据·运维·人工智能·web安全·自动化
Zzj_tju12 分钟前
Prompt Injection 防御:隔离不可信上下文的最小复现
人工智能·深度学习·机器学习·自然语言处理·prompt
合合技术团队12 分钟前
论文解读|合合信息与上海交通大学打造DocIQ模型,AI为文档图像质量“做体检”
人工智能·计算机视觉
Superzhangaa12 分钟前
太希智能全栈自研背后的机器人逻辑
人工智能·机器人·开源