如何设置TensorFlow Hub中BERT预处理层的output_shape参数?
如何设置TensorFlow Hub中BERT预处理层的输出序列长度?
我正在用TensorFlow Hub构建文本分类的简易BERT模型,代码如下:
import tensorflow as tf import tensorflow_hub as tf_hub bert_preprocess = tf_hub.KerasLayer("https://tfhub.dev/tensorflow/bert_en_uncased_preprocess/3") bert_encoder = tf_hub.KerasLayer("https://tfhub.dev/tensorflow/bert_en_uncased_L-12_H-768_A-12/4") text_input = tf.keras.layers.Input(shape=(), dtype=tf.string, name='text') preprocessed_text = bert_preprocess(text_input) encoded_input = bert_encoder(preprocessed_text) l1 = tf.keras.layers.Dropout(0.3, name="dropout1")(encoded_input['pooled_output']) l2 = tf.keras.layers.Dense(1, activation='sigmoid', name="output")(l1) model = tf.keras.Model(inputs=[text_input], outputs = [l2]) model.summary()
观察模型摘要后发现,bert_preprocess的输出序列长度固定为128,但我的文本平均长度远短于此,想把输出长度改成40左右。尝试通过output_shape参数传递给tf_hub.KerasLayer:
bert_preprocess = tf_hub.KerasLayer("https://tfhub.dev/tensorflow/bert_en_uncased_preprocess/3", output_shape=(64,)) # 或者 bert_preprocess = tf_hub.KerasLayer("https://tfhub.dev/tensorflow/bert_en_uncased_preprocess/3", output_shape=[64])
但调用预处理层时会抛出ValueError,请问正确的设置方式是什么?
正确设置方法
output_shape参数仅用于手动覆盖模型输出形状,无法配置BERT预处理层的序列长度。要自定义预处理后的序列长度,需通过arguments参数传递配置:
- 直接在构建
tf_hub.KerasLayer时,添加arguments={"max_seq_length": 40}参数,指定你需要的序列长度:
bert_preprocess = tf_hub.KerasLayer( "https://tfhub.dev/tensorflow/bert_en_uncased_preprocess/3", arguments={"max_seq_length": 40} )
- 测试输出是否符合预期:
test_text = ["we have a very sunny day today don't you think so?"] output = bert_preprocess(test_text) print(output['input_word_ids'].shape) # 输出为(1, 40),即批量大小1,序列长度40
- 将修改后的预处理层整合到原模型中:
import tensorflow as tf import tensorflow_hub as tf_hub bert_preprocess = tf_hub.KerasLayer( "https://tfhub.dev/tensorflow/bert_en_uncased_preprocess/3", arguments={"max_seq_length": 40} ) bert_encoder = tf_hub.KerasLayer("https://tfhub.dev/tensorflow/bert_en_uncased_L-12_H-768_A-12/4") text_input = tf.keras.layers.Input(shape=(), dtype=tf.string, name='text') preprocessed_text = bert_preprocess(text_input) encoded_input = bert_encoder(preprocessed_text) l1 = tf.keras.layers.Dropout(0.3, name="dropout1")(encoded_input['pooled_output']) l2 = tf.keras.layers.Dense(1, activation='sigmoid', name="output")(l1) model = tf.keras.Model(inputs=[text_input], outputs = [l2]) model.summary()
此时查看模型摘要,bert_preprocess的输出序列长度会变为40,满足需求。
补充说明
TF Hub提供的BERT预处理层内置了分词、截断和填充逻辑,max_seq_length是控制输出序列长度的核心参数。通过arguments参数传递该值,会覆盖预处理层的默认配置(默认128),从而生成指定长度的序列输入。
内容的提问来源于stack exchange,提问作者lazarea
相关产品推荐
相关产品推荐

