使用tf.function编写张量处理函数时遇'false_fn' must be callable错误求助
问题分析与解决
核心错误原因
你遇到的'false_fn' must be callable错误,是因为tf.case的default参数要求传入可调用对象(比如无参函数),但你直接传了张量x,不符合API要求。
代码中的其他问题
除了这个错误,你的代码还有几处需要修正:
logp_dash未定义,应该调用TensorFlow的对数函数tf.math.log(p_dash),注意要处理输入为0的情况(避免log(0)报错)Ku=K.sum(Ku)语法错误,tf.Tensor.sum()不需要传入自身作为参数,直接调用K.sum()即可- 冗余的类型转换:没必要先把值转成Python float再转TensorFlow张量,直接用TensorFlow的API处理更高效
修正后的代码
import tensorflow as tf from tensorflow.math import log @tf.function def Spa(x): # 直接转换输入为tf.float32张量,无需先转Python float x = tf.convert_to_tensor(x, dtype=tf.float32) p = tf.constant(0.05, dtype=tf.float32) p_dash = x # 处理log(0)的情况,添加极小值避免报错 K = p * log(tf.maximum(p_dash, 1e-10)) Ku = K.sum() y = tf.constant(0.0, dtype=tf.float32) # 定义返回0的可调用函数 def return_zero(): return tf.constant(0.0, dtype=tf.float32) # 定义返回x的可调用函数作为default def return_x(): return x # tf.case的每个分支和default都必须是可调用对象 r = tf.case( [(tf.less(x, y), return_zero), (tf.greater(x, Ku), return_zero)], default=return_x, exclusive=False ) return r
验证说明
- 给
tf.math.log传入tf.maximum(p_dash, 1e-10),避免输入为0时触发对数函数的数值错误 - 所有分支和default都包装成了无参函数,符合
tf.case的API要求 - 添加了
@tf.function装饰器,满足你需要的图模式执行需求
内容的提问来源于stack exchange,提问作者p200401Samuel
相关产品推荐
相关产品推荐

