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

不平衡回归中自定义SERA评估指标返回NaN的问题排查求助

问题根源与解决方案

核心原因

你遇到的NaN问题,本质是自定义评分函数中传递的x_relevance与交叉验证时的验证集样本不匹配:

  • 你在make_scorer中直接传入了整个数据集的X['relevance'],但交叉验证的每个fold里,y_true和y_pred仅对应验证集的子集样本,两者长度不一致。
  • 拼接DataFrame时,索引不匹配导致部分行的true/pred或phi为NaN,计算平方误差时产生NaN,最终积分结果为NaN。

修复步骤与优化代码

1. 修改SERA计算函数

调整函数逻辑,通过y_true的索引筛选对应验证集的relevance,确保样本匹配;同时简化间隔值生成、增加空子集处理:

import pandas as pd
import numpy as np
from scipy.integrate import simps
from sklearn.metrics import make_scorer

def calc_sera(y_true, y_pred, x_relevance=None):
    # 将输入转换为带索引的Series,确保样本匹配
    y_true = pd.Series(y_true)
    y_pred = pd.Series(y_pred, index=y_true.index)
    # 提取当前验证集对应的相关性数据
    rel_subset = x_relevance.loc[y_true.index]
    
    # 生成0到1的均匀间隔值(步长0.001,共1001个点),替代手动列表生成
    epsilon = np.linspace(0, 1, 1001)
    
    relevance_vals = []
    ser_vals = []
    
    for phi in epsilon:
        relevance_vals.append(phi)
        # 筛选相关性≥phi的样本
        mask = rel_subset >= phi
        # 处理空子集情况,避免潜在计算问题
        if not mask.any():
            ser_vals.append(0.0)
            continue
        # 计算平方误差和
        error_sq_sum = (y_true[mask] - y_pred[mask]).pow(2).sum()
        ser_vals.append(error_sq_sum)
    
    # 数值积分计算SERA
    sera = simps(ser_vals, relevance_vals)
    return sera

2. 正确创建评分器

保持传入整个数据集的x_relevance,但现在函数会自动筛选验证集对应的部分:

sera_scorer = make_scorer(calc_sera, x_relevance=X['relevance'], greater_is_better=False)

3. 保持原有调优代码不变

你的CatBoost和GridSearchCV逻辑无需修改,只需替换评分器名称:

from catboost import CatBoostRegressor
from sklearn.model_selection import KFold, GridSearchCV

model = CatBoostRegressor(random_state=0)
cv = KFold(n_splits=5, shuffle=True, random_state=42)

parameters = {
    'depth': [6,8,10],
    'learning_rate': [0.01, 0.05, 0.1],
    'iterations': [100, 200, 500,1000]
}

clf = GridSearchCV(
    estimator=model, 
    param_grid=parameters,
    scoring=sera_scorer, 
    verbose=0,
    cv=cv
)

clf.fit(
    X=X.drop(columns=['relevance']),
    y=y,
    sample_weight=X['relevance']
)

print("Best parameters:", clf.best_params_)
print("Lowest SERA: ", clf.best_score_)

额外注意事项

  • 确保X['relevance']是带索引的pandas Series,这样才能通过y_true.index正确筛选验证集数据。
  • 新增的空子集判断(if not mask.any())可以避免极端情况下(如所有样本相关性都小于某个phi)的计算异常。
  • 使用np.linspace替代手动生成间隔值,减少了round操作带来的精度误差,代码更简洁可靠。

内容的提问来源于stack exchange,提问作者Daniel Aben-Athar Bemerguy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 10:10:31