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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 18:18:38