TensorFlow训练AraBART调用model.fit时出现SelectV2类型不匹配报错
问题描述
参考Hugging Face官方summarization-tf示例,使用AraBART模型、xlsum阿拉伯语数据集开展文本摘要任务,执行如下训练代码时触发报错:
model.fit( train_dataset, validation_data=validation_dataset, epochs=1 )
核心报错
完整错误栈如下:
TypeError: in user code: /opt/conda/lib/python3.7/site-packages/keras/engine/training.py:853 train_function * return step_function(self, iterator) /opt/conda/lib/python3.7/site-packages/transformers/modeling_tf_utils.py:1279 run_call_with_unpacked_inputs * return func(self, **unpacked_inputs) /opt/conda/lib/python3.7/site-packages/transformers/models/mbart/modeling_tf_mbart.py:1300 call * labels = tf.where( /opt/conda/lib/python3.7/site-packages/tensorflow/python/util/dispatch.py:206 wrapper ** return target(*args, **kwargs) /opt/conda/lib/python3.7/site-packages/tensorflow/python/ops/array_ops.py:4716 where_v2 return gen_math_ops.select_v2(condition=condition, t=x, e=y, name=name) /opt/conda/lib/python3.7/site-packages/tensorflow/python/ops/gen_math_ops.py:8912 select_v2 "SelectV2", condition=condition, t=t, e=e, name=name) /opt/conda/lib/python3.7/site-packages/tensorflow/python/framework/op_def_library.py:558 _apply_op_helper inferred_from[input_arg.type_attr])) TypeError: Input 'e' of 'SelectV2' Op has type int64 that does not match type int32 of argument 't'.
报错本质为tf.where(对应SelectV2算子)要求两个分支输入的张量类型完全一致,当前传入的e参数为int64类型,t参数为int32类型,类型不匹配触发算子校验失败。
排查方向
- 检查数据集预处理完成后,训练集、验证集返回的
labels字段张量类型,确认是否为int64 - 核对当前环境安装的
transformers与tensorflow版本,确认是否存在跨版本适配问题:部分版本的TF MBart实现中,默认填充的标签忽略索引为硬编码的int32类型,和预处理默认生成的int64类型标签天然冲突 - 检查数据预处理逻辑,确认标签生成环节是否缺少类型转换步骤
解决方案
- 预处理阶段强制将所有模型输入张量转为int32类型,改动最小、生效最快。在tokenize处理函数末尾增加类型转换逻辑,参考代码:
import tensorflow as tf def preprocess_function(examples): # 原有输入、标签tokenize逻辑保持不变 inputs = [doc for doc in examples["text"]] model_inputs = tokenizer(inputs, max_length=1024, truncation=True) labels = tokenizer(text_target=examples["summary"], max_length=128, truncation=True) model_inputs["labels"] = labels["input_ids"] # 新增类型转换代码 model_inputs = {k: tf.cast(v, dtype=tf.int32) for k, v in model_inputs.items()} return model_inputs
- 固定依赖版本:将
transformers降级到4.28.0版本,搭配TensorFlow 2.7~2.10版本使用,该版本组合不存在上述类型硬编码不匹配问题,无需修改业务代码 - 本地源码修改:直接修改本地安装路径下的
modeling_tf_mbart.py文件第1300行附近的tf.where逻辑,将填充的忽略索引值转为和labels一致的类型,该方案侵入性强,后续升级依赖会导致修改失效,不推荐使用。
内容的提问来源于stack exchange,提问作者stu_dent
相关产品推荐
相关产品推荐

