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
相关产品推荐
相关产品推荐

