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

自定义Perceptron模型无法使用scikit-learn的K-Fold交叉验证?

解决自定义Keras模型无法使用Scikit-learn交叉验证的问题

问题原因

Scikit-learn的交叉验证函数(如cross_val_score、cross_val_predict)要求传入的估计器必须实现get_params和set_params方法,而你的自定义Perceptron类继承自tf.keras.Model,默认没有实现这些方法,导致无法被克隆,从而触发报错。

解决方案

方案1:用KerasClassifier包装模型(推荐)

Scikit-learn提供了专门的包装类,可将Keras模型转换为符合sklearn规范的估计器,无需手动实现接口方法。

修改后的完整代码:

import tensorflow as tf
from sklearn.model_selection import cross_val_predict, KFold
from tensorflow.keras.wrappers.scikit_learn import KerasClassifier

def build_perceptron():
    # 定义模型构建函数,交叉验证时会每次重新初始化模型
    model = tf.keras.Sequential()
    model.add(tf.keras.layers.Dense(units=1, activation='sigmoid', input_shape=(x_train.shape[1],)))
    model.compile(optimizer='sgd', loss='binary_crossentropy', metrics=['accuracy'])
    return model

# 包装模型,传入训练参数
model = KerasClassifier(build_fn=build_perceptron, epochs=100, batch_size=1, verbose=0)

kfold = KFold(n_splits=5, shuffle=True, random_state=42)
scores = cross_val_predict(model, x_train, y_train, cv=kfold)

方案2:手动实现Scikit-learn估计器接口

若不想使用包装类,可以让自定义模型继承sklearn的基类,手动实现必要的接口方法:

import tensorflow as tf
from sklearn.base import BaseEstimator, ClassifierMixin
from sklearn.model_selection import cross_val_predict, KFold

class Perceptron(BaseEstimator, ClassifierMixin, tf.keras.Model):
    def __init__(self, units=1, activation='sigmoid', optimizer='sgd', loss='binary_crossentropy'):
        super().__init__()
        # 定义可被get_params识别的参数
        self.units = units
        self.activation = activation
        self.optimizer = optimizer
        self.loss = loss
        self.dense = tf.keras.layers.Dense(units=self.units, activation=self.activation)

    def call(self, inputs):
        return self.dense(inputs)

    def fit(self, X, y, epochs=100, batch_size=1, verbose=0):
        self.compile(optimizer=self.optimizer, loss=self.loss, metrics=['accuracy'])
        super().fit(X, y, epochs=epochs, batch_size=batch_size, verbose=verbose)
        return self

    def predict(self, X):
        # 将概率输出转换为二分类类别(以0.5为阈值)
        return (self.predict(X) > 0.5).astype(int)

# 使用模型进行交叉验证
model = Perceptron()
kfold = KFold(n_splits=5, shuffle=True, random_state=42)
scores = cross_val_predict(model, x_train, y_train, cv=kfold)

注意事项

  • 方案1中必须传入模型构建函数而非已实例化的模型,这样交叉验证的每个折都会重新初始化模型,避免数据泄露。
  • 若处理回归任务,需使用KerasRegressor替代KerasClassifier。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 22:35:34