TensorFlow2.4.0-rc0报OperatorNotAllowedInGraphError错误排查
错误位置与原因
报错触发点为代码中的if kpt[2] == tf.constant(0.0):行。
被@tf.function装饰的函数会被编译为静态计算图执行,Python原生if/else分支要求判断条件为Python布尔值,但此处kpt是张量类型,kpt[2] == tf.constant(0.0)的返回值为布尔张量,无法直接作为Python分支的判断依据,因此触发类型错误。
解决方法
使用TensorFlow静态图兼容的条件分支算子tf.cond替换原生Python的if/else逻辑即可,修改后代码如下:
@tf.function def keypoint_distance(self, kpt): condition = tf.equal(kpt[2], 0.0) # true_fn对应条件成立时的执行逻辑,false_fn对应条件不成立时的执行逻辑 return tf.cond( pred=condition, true_fn=lambda: tf.ones((self.LABEL_HEIGHT, self.LABEL_WIDTH), dtype=tf.float32), false_fn=lambda: tf.linalg.norm(self.grid - kpt[0:2], axis=-1) )
注意事项
如果你的业务场景不需要强静态图优化,也可以在@tf.function中添加参数关闭严格执行校验,2.4.0-rc0版本可写为@tf.function(experimental_compile=False),但更推荐使用tf.cond适配静态图逻辑,保证执行效率和版本兼容性。
内容的提问来源于stack exchange,提问作者ge1mina023
相关产品推荐
相关产品推荐

