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

TensorFlow BERT调用cross_val_predict报predict_classes属性错误

问题背景

计划基于TensorFlow框架使用BERT模型完成5分类任务,预设训练配置如下:

  • 采用5折交叉验证
  • 设置batch_size = 64、epoch = 100
  • 搭配早停(early stopping)机制
  • 所用数据集共包含5个类别

用于复现问题的示例代码如下:

import tensorflow_hub as hub
import tensorflow_text as text
...


bert_preprocess = hub.KerasLayer("https://tfhub.dev/tensorflow/bert_en_uncased_preprocess/3")
bert_encoder = hub.KerasLayer("https://tfhub.dev/tensorflow/bert_en_uncased_L-12_H-768_A-12/4")

def create_model():

  model = Sequential()
  text_input = tf.keras.layers.Input(shape=(), dtype=tf.string, name='text')
  preprocessed_text = bert_preprocess(text_input)
  outputs = bert_encoder(preprocessed_text)

  l = tf.keras.layers.Dropout(0.1, name="dropout")(outputs['pooled_output'])
  l = tf.keras.layers.Dense(5, activation='softmax', name="output")(l)

  model = tf.keras.Model(inputs=[text_input], outputs = [l])


  model.summary()
  model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['categorical_accuracy'])
  return model

estimator = KerasClassifier(build_fn = create_model,
                                  epochs=1,
                                  batch_size=64, 
                                  verbose=1
                                  )


kf = KFold(n_splits=5) 
pred = cross_val_predict(estimator, X_train, Y_train,
                          cv=kf,                          
                          verbose=1)

运行代码时抛出如下错误:

Layer (type)                    Output Shape         Param #     Connected to                     
==================================================================================================
text (InputLayer)               [(None,)]            0                                            
__________________________________________________________________________________________________
keras_layer (KerasLayer)        {'input_type_ids': ( 0           text[0][0]                       
__________________________________________________________________________________________________
keras_layer_1 (KerasLayer)      {'encoder_outputs':  109482241   keras_layer[12][0]               
                                                                 keras_layer[12][1]               
                                                                 keras_layer[12][2]               
__________________________________________________________________________________________________
dropout (Dropout)               (None, 768)          0           keras_layer_1[12][13]            
__________________________________________________________________________________________________
output (Dense)                  (None, 5)            3845        dropout[0][0]                    
==================================================================================================
Total params: 109,486,086
Trainable params: 3,845
Non-trainable params: 109,482,241
__________________________________________________________________________________________________
38/38 [==============================] - 226s 6s/step - loss: 1.6111 - categorical_accuracy: 0.2434


---------------------------------------------------------------------------
AttributeError                            Traceback (most recent call last)
<ipython-input-26-55c3162a6b6d> in <module>()
     16                          
     17 
---> 18                           verbose=1)
     19 
     20 

10 frames
/usr/local/lib/python3.7/dist-packages/tensorflow/python/keras/wrappers/scikit_learn.py in predict(self, x, **kwargs)
    239     """
    240     kwargs = self.filter_sk_params(Sequential.predict_classes, kwargs)
---> 241     classes = self.model.predict_classes(x, **kwargs)
    242     return self.classes_[classes]
    243 

AttributeError: 'Functional' object has no attribute 'predict_classes'

移除代码中的model = Sequential()语句后重新运行,仍然触发相同错误,需要修复该问题以顺利落地预设的BERT分类训练流程。

错误原因

该报错和是否删除Sequential()初始化语句无关,核心原因有两点:

  1. 代码最终构建的是tf.keras.Model类型的函数式(Functional)模型,并非Sequential模型
  2. TensorFlow 2.6及以上版本已经彻底移除了所有Keras模型的predict_classes方法,而当前使用的tensorflow.keras.wrappers.scikit_learn.KerasClassifier是已停止维护的旧版适配代码,内部默认调用已经被删除的predict_classes方法做分类预测,因此会抛出属性不存在的错误。
修复方案

方案1:替换为官方推荐的新版适配层(推荐)

旧的Keras sklearn封装器已经废弃,官方推荐使用独立维护的scikeras库作为Keras和scikit-learn的适配层,该库已经修复了predict_classes的兼容问题,步骤如下:

  1. 安装依赖:
pip install scikeras
  1. 修改导入语句,替换旧的KerasClassifier,同时清理无效代码、补充早停配置匹配预设训练要求:
# 删除旧导入:from tensorflow.keras.wrappers.scikit_learn import KerasClassifier
from scikeras.wrappers import KerasClassifier
from sklearn.model_selection import KFold, cross_val_predict
import tensorflow as tf

# 配置早停机制
early_stop = tf.keras.callbacks.EarlyStopping(
    monitor='val_loss',
    patience=5,
    restore_best_weights=True
)

def create_model():
  # 删除无效的model = Sequential()语句
  text_input = tf.keras.layers.Input(shape=(), dtype=tf.string, name='text')
  preprocessed_text = bert_preprocess(text_input)
  outputs = bert_encoder(preprocessed_text)

  l = tf.keras.layers.Dropout(0.1, name="dropout")(outputs['pooled_output'])
  l = tf.keras.layers.Dense(5, activation='softmax', name="output")(l)

  model = tf.keras.Model(inputs=[text_input], outputs = [l])
  model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['categorical_accuracy'])
  return model

# 初始化estimator,传入预设训练参数和早停回调
estimator = KerasClassifier(
    model=create_model,
    epochs=100,
    batch_size=64, 
    verbose=1,
    callbacks=[early_stop]
)

# 交叉验证建议开启shuffle,避免数据分布偏移影响结果
kf = KFold(n_splits=5, shuffle=True, random_state=42) 
pred = cross_val_predict(estimator, X_train, Y_train, cv=kf, verbose=1)

方案2:手动兼容旧版封装器(临时使用)

如果不想额外安装依赖,可以在模型创建时手动为实例绑定predict_classes方法,模拟旧版本行为:

def create_model():
  text_input = tf.keras.layers.Input(shape=(), dtype=tf.string, name='text')
  preprocessed_text = bert_preprocess(text_input)
  outputs = bert_encoder(preprocessed_text)

  l = tf.keras.layers.Dropout(0.1, name="dropout")(outputs['pooled_output'])
  l = tf.keras.layers.Dense(5, activation='softmax', name="output")(l)

  model = tf.keras.Model(inputs=[text_input], outputs = [l])
  model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['categorical_accuracy'])
  
  # 手动实现predict_classes逻辑
  def predict_classes(x, **kwargs):
    pred_proba = model.predict(x, **kwargs)
    return pred_proba.argmax(axis=-1)
  model.predict_classes = predict_classes

  return model

注意:该方案仅做临时兼容,旧版封装器不会再更新维护,长期使用建议迁移到scikeras。

额外优化提示

从模型summary可以看到当前BERT编码器的参数全部被冻结,可训练参数只有最后分类层的3845个,如果你的数据集和通用英文语料差异较大,建议训练几轮后解冻BERT顶层2-3层,用更小的学习率微调,分类效果会有明显提升。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.31 08:33:20