终端运行TFP代码报错:无法将KerasTensor解析为数据类型
问题表现
使用TensorFlow Probability工具包通过bijector从简单分布(如高斯分布)构建可训练分布时,Jupyter Notebook中可正常运行的代码迁移到终端执行.py文件时抛出如下错误:
TypeError: Cannot interpret '<KerasTensor: shape=(None, 3) dtype=float32 (created by layer 'input_1')>' as a data type
已知条件:
- 输入数据
X_data维度为(m,3),m为样本数量 trainable_distribution已通过tensorflow_probability.distributions.TransformedDistribution(base_dist, bijector)构建完成- 出错代码片段:
def train_dist_routine(X_data, trainable_distribution, n_epochs=200, batch_size=None): x_ = tensorflow.keras.layers.Input(shape=(3,), dtype=tf.float32) print(x_) log_prob_ = trainable_distribution.log_prob(x_) model = tensorflow.keras.models.Model(x_, log_prob_) model.compile(optimizer=tf.optimizers.Adam(), loss=lambda _, log_prob: -log_prob) ns = X_data.shape[0] if batch_size is None: batch_size = ns history = model.fit(x=X_data, y=np.zeros((ns, 0), dtype=np.float32), batch_size=batch_size, epochs=n_epochs, validation_split=0.2, shuffle=True, verbose=False) return history
错误原因
Jupyter Notebook环境默认开启TensorFlow Eager Execution即时执行模式,KerasTensor可以被自动识别转换为普通张量传入TFP算子计算;终端执行.py脚本时,默认的图执行模式下,裸调用的TFP分布方法无法识别KerasTensor类型,无法完成dtype推断就会抛出该类型错误。
修复方案
按优先级选择以下任意一种方案即可:
- 对齐运行模式:在所有TensorFlow、TFP相关导入语句之后,加入一行代码强制开启即时执行,和Notebook运行环境保持一致:
tf.config.run_functions_eagerly(True) - 包装为Keras兼容层:不要直接在Keras计算图中裸调用
trainable_distribution.log_prob,使用TFP提供的DistributionLambda层包装计算逻辑,自动完成KerasTensor的类型适配:from tensorflow_probability.layers import DistributionLambda x_ = tf.keras.layers.Input(shape=(3,), dtype=tf.float32) log_prob_ = DistributionLambda(lambda inp: trainable_distribution.log_prob(inp))(x_) model = tf.keras.models.Model(x_, log_prob_) - 检查版本匹配:如果上述方案都无效,确认安装的TensorFlow和TensorFlow Probability版本为官方匹配的对应版本,版本不兼容也会导致KerasTensor和TFP算子的类型适配失效。
内容的提问来源于stack exchange,提问作者Negar Erfanian
相关产品推荐
相关产品推荐

