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

TensorFlow自定义损失函数报InvalidArgumentError: BatchMatMulV2输入类型不匹配

报错原因

InvalidArgumentError触发的核心是tf.matmul运算要求两个输入张量的数据类型完全一致:你的标签张量y_true默认是int64整型,而tf.math.log(y_pred)输出是浮点型(float32/float64,也就是报错里提到的double张量),类型不匹配导致运算失败。
额外说明:你当前手写的二元交叉熵损失用tf.matmul属于多余操作,不仅容易触发维度、类型类报错,还会增加不必要的计算开销,交叉熵本身只需逐元素计算即可。

解决方案

优先推荐用两种写法修复,同时规避后续数值不稳定问题:

  • 方案1:直接调用TensorFlow官方内置的二元交叉熵接口,避免手写踩坑
def loss_function(y_pred, y_true):
    # 先将标签转为和预测值一致的浮点类型
    y_true = tf.cast(y_true, dtype=y_pred.dtype)
    return tf.reduce_mean(tf.keras.losses.binary_crossentropy(y_true, y_pred))
  • 方案2:保留手写逻辑,修正类型+替换矩阵乘法为逐元素运算,同时增加数值裁剪避免log(0)出现无穷值
def loss_function(y_pred, y_true):
    # 统一两个张量的数据类型
    y_true = tf.cast(y_true, dtype=y_pred.dtype)
    # 裁剪预测值,防止出现log(0)导致的数值异常
    y_pred = tf.clip_by_value(y_pred, 1e-7, 1 - 1e-7)
    # 逐元素计算交叉熵后取均值
    return -tf.reduce_mean(y_true * tf.math.log(y_pred) + (1 - y_true) * tf.math.log(1 - y_pred))

内容的提问来源于stack exchange,提问作者S M Abrar Jahin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 02:57:03