使用TF-Hub BERT构建分类模型时tf.data训练报错如何解决?
错误根源
报错的核心原因是你构造的tf.data数据集没有添加批次维度,TF-Hub提供的BERT预处理器要求输入为带批次维度的字符串张量(shape为(None,)),而你当前未做batch处理的数据集每个样本的输入是shape为()的单个字符串标量,和预处理器要求的输入规格不匹配。
小提示:你原有代码中train_dataset的定义存在括号嵌套语法错误,将test_dataset的定义包裹进了train_dataset的赋值括号内,修正时请注意调整括号位置。
解决方案
方案1(推荐):调整数据集构造格式,增加batch处理
直接将数据集构造为(输入, 标签)的元组格式,同时添加batch、shuffle、prefetch等常规数据流水线配置:
# 可根据自身显存大小调整批次大小 batch_size = 16 # 训练集构造 train_dataset = tf.data.Dataset.from_tensor_slices( ( tf.cast(corpus_train.values, tf.string), tf.cast(labels_train, tf.int32) ) ).shuffle(buffer_size=len(corpus_train)) # 打乱训练集顺序 .batch(batch_size) # 增加批次维度,匹配预处理器输入要求 .prefetch(tf.data.AUTOTUNE) # 预加载数据提升训练效率 # 测试集构造 test_dataset = tf.data.Dataset.from_tensor_slices( ( tf.cast(corpus_test.values, tf.string), tf.cast(labels_test, tf.int32) ) ).batch(batch_size).prefetch(tf.data.AUTOTUNE)
修改完成后直接调用原有fit代码即可正常训练。
方案2:保留原字典数据集格式,调整fit参数
如果你不想修改原有的字典结构数据集,也可以先给数据集加batch,再在fit时手动指定输入和标签的取值字段:
batch_size = 16 # 先给数据集增加batch维度 train_dataset = train_dataset.batch(batch_size) test_dataset = test_dataset.batch(batch_size) # 训练时手动拆分输入和标签 classifier_model.fit( x=train_dataset.map(lambda item: item["features"]), y=train_dataset.map(lambda item: item["labels"]), validation_data=( test_dataset.map(lambda item: item["features"]), test_dataset.map(lambda item: item["labels"]) ), epochs=2 )
内容的提问来源于stack exchange,提问作者An old man in the sea.
相关产品推荐
相关产品推荐

