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

如何构建采用Keras神经网络作为二元分类器的Classifier Chain?

用Keras神经网络构建Classifier Chain的实现示例

刚好我之前折腾过这个需求,Scikit-learn官方文档里确实没给出用Keras神经网络作为基分类器的Classifier Chain示例,我来给你一步步拆解实现思路,附上可运行的代码。

核心逻辑回顾

Classifier Chain的本质是把多标签分类任务拆解成一系列二元分类任务:每个后续的分类器会把前面所有分类器的预测结果作为额外特征,拼接到原始输入里一起训练。要让Keras模型能和Scikit-learn的ClassifierChain配合,关键是把Keras模型包装成符合Scikit-learn接口的分类器(也就是实现fit、predict和predict_proba方法)。

步骤1:把Keras模型包装成Scikit-learn兼容的分类器

我们需要写一个自定义类,继承Scikit-learn的BaseEstimator和ClassifierMixin,这样就能无缝接入ClassifierChain。这个类会帮我们处理Keras模型的构建、训练和预测逻辑:

import numpy as np
from sklearn.base import BaseEstimator, ClassifierMixin
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense, Dropout
from tensorflow.keras.optimizers import Adam

class KerasBinaryClassifier(BaseEstimator, ClassifierMixin):
    def __init__(self, hidden_units=64, dropout_rate=0.2, epochs=20, batch_size=32):
        # 这里不提前指定输入维度,训练时会自动从数据中获取
        self.hidden_units = hidden_units
        self.dropout_rate = dropout_rate
        self.epochs = epochs
        self.batch_size = batch_size
        self.model = None
        self.input_dim = None  # 动态存储当前分类器的输入维度

    def _build_model(self):
        # 构建一个简单的全连接神经网络,你可以根据需求改成CNN/RNN等
        model = Sequential([
            Dense(self.hidden_units, activation='relu', input_shape=(self.input_dim,)),
            Dropout(self.dropout_rate),
            Dense(self.hidden_units//2, activation='relu'),
            Dropout(self.dropout_rate),
            Dense(1, activation='sigmoid')
        ])
        model.compile(
            optimizer=Adam(learning_rate=0.001),
            loss='binary_crossentropy',
            metrics=['accuracy']
        )
        return model

    def fit(self, X, y, **kwargs):
        # 从训练数据中获取输入维度,动态构建模型
        self.input_dim = X.shape[1]
        self.model = self._build_model()
        # 确保标签是二维数组,适配Keras的输入要求
        y = y.reshape(-1, 1)
        # 训练模型,关闭日志输出避免刷屏
        self.model.fit(X, y, epochs=self.epochs, batch_size=self.batch_size, verbose=0, **kwargs)
        return self

    def predict_proba(self, X):
        # ClassifierChain依赖概率输出,必须返回(n_samples, 2)的格式(负类、正类概率)
        assert X.shape[1] == self.input_dim, "输入特征维度和训练时不匹配!"
        prob_positive = self.model.predict(X, verbose=0)
        prob_negative = 1 - prob_positive
        return np.hstack([prob_negative, prob_positive])

    def predict(self, X):
        # 基于概率输出做二元分类预测
        prob = self.predict_proba(X)[:, 1]
        return (prob >= 0.5).astype(int)

步骤2:构建并训练Classifier Chain

现在我们可以用Scikit-learn的ClassifierChain,把上面的Keras分类器作为基模型,配合多标签数据集进行训练和测试。这里用Scikit-learn自带的模拟多标签数据来演示:

from sklearn.datasets import make_multilabel_classification
from sklearn.model_selection import train_test_split
from sklearn.multioutput import ClassifierChain
from sklearn.metrics import hamming_loss, classification_report

# 生成模拟多标签数据:1000样本,20个特征,5个标签类别,每个样本平均2个标签
X, y = make_multilabel_classification(
    n_samples=1000, n_features=20, n_classes=5, n_labels=2, random_state=42
)

# 划分训练集和测试集
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

# 初始化Keras基分类器
base_keras_clf = KerasBinaryClassifier(hidden_units=128, epochs=30, batch_size=64)

# 构建Classifier Chain:随机打乱标签顺序,你也可以手动指定order参数
chain = ClassifierChain(base_estimator=base_keras_clf, order='random', random_state=42)

# 训练Chain
chain.fit(X_train, y_train)

# 生成预测结果
y_pred = chain.predict(X_test)
y_pred_proba = chain.predict_proba(X_test)

# 评估模型性能
print(f"Hamming Loss: {hamming_loss(y_test, y_pred):.4f}")
print("\n分类报告:")
print(classification_report(y_test, y_pred))

几个关键注意事项

  • 必须实现predict_proba:Classifier Chain在训练后续分类器时,需要前面分类器的概率输出作为特征,所以这个方法是核心,不能省略,而且输出格式要严格符合Scikit-learn的要求(n_samples × 2的数组)。
  • 动态输入维度:链中每个分类器的输入维度是原始特征数加上前面所有标签的数量,所以我们不能提前固定输入维度,而是在fit时动态获取,这样每个分类器都能适配自己的输入。
  • 模型可定制:你可以完全替换_build_model里的网络结构,比如换成卷积神经网络处理图像,或者循环神经网络处理序列数据,只要保证最终输出是单个sigmoid单元的二元分类结果就行。
  • 链的顺序影响性能:不同的标签顺序会影响Chain的最终效果,你可以尝试order='random'多次训练取平均,或者根据标签之间的相关性手动指定顺序来优化。

内容的提问来源于stack exchange,提问作者Tuong Nguyen Minh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.09 09:58:16