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

自定义scikit-learn新奇检测分类器通过check_estimator测试遇阻

解决scikit-learn自定义新奇检测分类器的check_estimator测试问题

首先,你遇到的问题根源在于:check_classifiers_classes()测试是为多类分类器设计的,而你的新奇检测模型本质是二类(正常/异常)的单类检测任务,不符合该测试的预设条件。下面提供两种针对性的解决方案:

方案1:适配scikit-learn的异常检测接口(推荐)

scikit-learn专门为新奇/异常检测任务提供了OutlierMixin基类,这类模型的设计逻辑就是区分"正常样本"和"异常样本",完全适配你的场景:

  • 让你的自定义分类器继承OutlierMixin和BaseEstimator,而非常规的ClassifierMixin。
  • 调整接口细节:
    • fit方法可以只接收正常样本的X(如果训练时没有异常样本),或者接收带标记的y(比如用1表示正常,-1表示异常,这是scikit-learn的惯例)。
    • predict方法返回1(正常)或-1(异常),符合scikit-learn异常检测模型的规范。
  • 此时调用check_estimator时,会自动适配异常检测相关的测试用例,不会触发check_classifiers_classes()这类针对多类分类器的测试。

示例代码框架:

from sklearn.base import BaseEstimator, OutlierMixin
from sklearn.utils.validation import check_array

class CustomNoveltyDetector(BaseEstimator, OutlierMixin):
    def fit(self, X, y=None):
        # 处理训练逻辑,比如学习正常样本的分布
        X = check_array(X)
        # 你的训练代码...
        self.classes_ = [1, -1]  # 遵循异常检测的标签惯例
        return self
    
    def predict(self, X):
        X = check_array(X)
        # 你的预测逻辑,返回1(正常)或-1(异常)
        predictions = ...
        return predictions

方案2:跳过特定测试(保留ClassifierMixin时使用)

如果你因为业务需求必须继承ClassifierMixin(比如要融入分类器流水线),可以在调用check_estimator时显式跳过check_classifiers_classes测试:

from sklearn.utils.estimator_checks import check_estimator

check_estimator(YourCustomClassifier, exclude=["check_classifiers_classes"])

同时需要确保你的分类器满足二类分类器的核心规范:

  • 在fit方法中正确设置classes_属性为np.array([0, 1])。
  • predict方法仅返回0或1两个标签。
  • 如果实现了predict_proba,需返回对应两个类别的概率值。

这种方式虽然能通过测试,但要注意:部分scikit-learn的分类器工具(比如多类交叉验证)可能不会完全适配你的二类新奇检测场景,使用时需要额外留意。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 09:25:59