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

使用Optuna批量调优CatBoost时第二轮试验出现NaN错误

问题分析

核心问题有两个:

  1. 生成器只能迭代一次:第一次试验耗尽了训练/验证生成器,后续试验中循环无法获取任何数据,导致rmse为空列表,计算均值时得到nan。
  2. 批量训练逻辑错误:每次处理一个batch时,都调用fit并训练完整的iterations轮次,这会导致模型在每个batch上重复训练几百到几千次,不仅效率极低,还会引发过拟合或训练异常。
修复方案

1. 每次试验重新生成批量生成器

将生成器的创建逻辑移到objective_regressor函数内部,确保每个试验都能获取到完整的训练/验证批次数据。

2. 调整增量训练的迭代策略

CatBoost的init_model用于增量训练时,每次fit的iterations参数代表本次新增的训练轮次,而非总轮次。因此需要将总迭代次数分配到每个batch中,或者设置每batch训练固定的少量轮次。

3. 优化验证逻辑

避免每个batch都评估一次,可改为每N个batch评估一次,或者在所有batch训练完成后用完整验证集评估,减少计算开销。

4. 添加异常防护

处理rmse为空的情况,返回一个极大值让Optuna自动跳过无效参数组合。

修改后的代码
import numpy as np
import math
from sklearn.model_selection import train_test_split
from sklearn.metrics import mean_squared_error
import optuna
from catboost import CatBoostRegressor
from tqdm import tqdm

def expand_embeddings(df, embedding_col="embeddings"):
    embeddings = np.array(df[embedding_col].to_list(), dtype=np.float32)
    other_features = df.drop(columns=[embedding_col]).to_numpy(dtype=np.float32)
    return np.hstack([other_features, embeddings])

def batch_generator(df, target_col, batch_size):
    for i in range(0, len(df), batch_size):
        batch = df.iloc[i:i + batch_size]
        y = batch[target_col].to_numpy(dtype=np.float32)
        X = batch.drop(columns=[target_col])
        X = expand_embeddings(X)
        yield X, y

# 提前拆分数据集(无需每次试验重复拆分)
train_data, val_data = train_test_split(result, test_size=0.1, random_state=42)
num_batches = 1300
batch_size_train = math.ceil(train_data.shape[0] / num_batches)
batch_size_test = math.ceil(val_data.shape[0] / num_batches)

def objective_regressor(trial):
    # 每次试验重新生成批次生成器
    train_batches = batch_generator(train_data, target_col="weight", batch_size=batch_size_train)
    val_batches = batch_generator(val_data, target_col="weight", batch_size=batch_size_test)
    
    params = {
        'per_batch_iterations': trial.suggest_int('per_batch_iterations', 1, 10),  # 每批次训练轮次
        'depth': trial.suggest_int('depth', 4, 10),
        'learning_rate': trial.suggest_float('learning_rate', 0.01, 0.1),
        'l2_leaf_reg': trial.suggest_float('l2_leaf_reg', 1, 10),
        'bagging_temperature': trial.suggest_float('bagging_temperature', 0, 1),
        'random_strength': trial.suggest_float('random_strength', 0.1, 10),
        'eval_metric': 'RMSE'
    }

    model = CatBoostRegressor(
        iterations=params['per_batch_iterations'],
        depth=params['depth'],
        learning_rate=params['learning_rate'],
        l2_leaf_reg=params['l2_leaf_reg'],
        bagging_temperature=params['bagging_temperature'],
        random_strength=params['random_strength'],
        eval_metric=params['eval_metric'],
        task_type='CPU',
        random_seed=42,
        verbose=0
    )

    rmse_scores = []
    
    for X_batch, y_batch in tqdm(train_batches):
        # 处理验证批次耗尽的情况
        try:
            X_val_batch, y_val_batch = next(val_batches)
        except StopIteration:
            val_batches = batch_generator(val_data, target_col="weight", batch_size=batch_size_test)
            X_val_batch, y_val_batch = next(val_batches)
        
        # 增量训练:第一次训练无需init_model
        model.fit(
            X_batch, y_batch,
            eval_set=(X_val_batch, y_val_batch),
            init_model=model if rmse_scores else None,
            verbose=0
        )
        
        # 计算当前批次RMSE(直接用squared=False得到RMSE)
        y_pred = model.predict(X_val_batch)
        rmse_scores.append(mean_squared_error(y_val_batch, y_pred, squared=False))
    
    # 异常防护:无有效分数时返回极大值
    if not rmse_scores:
        return float('inf')
    
    return np.mean(rmse_scores)

study_regressor = optuna.create_study(direction='minimize')
study_regressor.optimize(objective_regressor, n_trials=20)
额外优化建议
  • 使用CatBoost Pool:将批次数据包装成catboost.Pool对象,能更好地兼容CatBoost的训练逻辑,提升效率。
  • 调整批量大小:如果num_batches=1300导致每个batch数据量过小,可以适当减少批次数量,避免模型在过小的数据集上训练导致不稳定。
  • 验证集复用:可以在所有batch训练完成后,用完整的验证集进行一次评估,结果更可靠。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 00:42:01