TensorFlow自定义损失函数运行model.fit时报期望bool却得到float错误如何解决
TypeError: Expected bool, got 0.0 of type 'float' instead.
报错原因
该错误出现在tf.cast操作的第一个入参位置,预期接收布尔类型值,实际传入了浮点类型的0.0,具体触发逻辑如下:
- 手动传入
tf.constant调用损失函数时,执行动态图逻辑,y_true == 0.0会直接返回布尔类型张量,符合tf.cast入参要求,运行正常。 - 调用
model.fit时,损失函数被@tf.function装饰器编译为静态图执行,若输入的y_true张量类型在静态图推断阶段不明确,y_true == 0.0的判断结果会被错误推断为浮点类型(1.0对应True、0.0对应False),不符合tf.cast第一个参数的类型要求,触发报错。
修复方案
任选以下一种方法即可解决:
- 方法1:显式将判断结果转为布尔类型后再做cast:
mask_0 = tf.cast(tf.cast(y_true == 0.0, tf.bool), tf.float32) mask_1 = tf.cast(tf.cast(y_true == 1.0, tf.bool), tf.float32)
- 方法2:使用TensorFlow官方提供的相等判断API避免语法糖的类型推断问题:
mask_0 = tf.cast(tf.equal(y_true, 0.0), tf.float32) mask_1 = tf.cast(tf.equal(y_true, 1.0), tf.float32)
- 方法3:在训练数据流水线中提前固定
y_true的类型为浮点型,消除静态图推断的不确定性:
比如在tf.data.Dataset的预处理逻辑中添加:
y_true = tf.cast(y_true, tf.float32)
原有代码问题
原有代码未兼容TensorFlow静态图的类型推断规则:直接使用Python原生==运算符做张量值判断,在静态图编译场景下容易出现类型推断异常,同时没有对判断结果做显式的类型约束,最终导致类型不匹配报错。
内容的提问来源于stack exchange,提问作者user5178150
相关产品推荐
相关产品推荐

