如何在TFP的DistributionLambda的convert_to_tensor_fn中使用Keras上层值
问题描述
我在构建Keras/TensorFlow Probability(TFP)模型时,希望在DistributionLambda层的convert_to_tensor_fn参数中使用前一层的输出值。尝试编写了如下代码:
from functools import partial import tensorflow as tf from tensorflow.keras import layers, Model import tensorflow_probability as tfp from typing import Union tfd = tfp.distributions zero_buffer = 1e-5 def quantile(s: tfd.Distribution, q: Union[tf.Tensor, float]) -> Union[tf.Tensor, float]: return s.quantile(q) # 4 records (1st value represents CDF value, # 2nd represents location, # 3rd represents scale) sample_input = tf.constant([[0.25, 0.0, 1.0], [0.5, 1.0, 0.5], [0.75, -1.0, 2.0], [0.95, 3.0, 2.5]], dtype=tf.float32) # Build toy model for demonstration input_layer = layers.Input(3) dist = tfp.layers.DistributionLambda( make_distribution_fn=lambda t: tfd.Normal(loc=t[..., 1], scale=zero_buffer + tf.nn.softplus(t[..., 2])), convert_to_tensor_fn=lambda t, s: partial(quantile, q=t[..., 0])(s) )(input_layer) model = Model(input_layer, dist)
但根据TFP官方文档,convert_to_tensor_fn仅允许接收tfd.Distribution作为输入,上述带两个参数的lambda写法无法运行。我想知道如何在convert_to_tensor_fn中访问前一层的数据,推测可以通过partial函数或类似方法实现。
在Keras模型框架外,该需求很容易实现,示例代码如下:
# input data in Tensor Constant form cdf_data = tf.constant([0.25, 0.5, 0.75, 0.95], dtype=tf.float32) norm_mu = tf.constant([0.0, 1.0, -1.0, 3.0], dtype=tf.float32) norm_scale = tf.constant([1.0, 0.5, 2.0, 2.5], dtype=tf.float32) quant = partial(quantile, q=cdf_data) norm = tfd.Normal(loc=norm_mu, scale=norm_scale) quant(norm)
输出结果:
<tf.Tensor: shape=(4,), dtype=float32, numpy=array([-0.6744898, 1. , 0.3489796, 7.112134 ], dtype=float32)>
解决方案
核心思路是让make_distribution_fn返回的分布对象携带前一层的额外数据(即CDF值q),这样convert_to_tensor_fn就能从分布实例中获取到所需的q值,以下提供两种可行实现方式:
方式一:利用Distribution的额外属性(推荐)
TFP的Distribution类提供了experimental_extra_properties方法,可以给分布对象附加自定义属性,无需自定义子类:
from functools import partial import tensorflow as tf from tensorflow.keras import layers, Model import tensorflow_probability as tfp from typing import Union tfd = tfp.distributions zero_buffer = 1e-5 def quantile(s: tfd.Distribution, q: Union[tf.Tensor, float]) -> Union[tf.Tensor, float]: return s.quantile(q) sample_input = tf.constant([[0.25, 0.0, 1.0], [0.5, 1.0, 0.5], [0.75, -1.0, 2.0], [0.95, 3.0, 2.5]], dtype=tf.float32) input_layer = layers.Input(3) dist = tfp.layers.DistributionLambda( # 在创建Normal分布时附加q值属性 make_distribution_fn=lambda t: tfd.Normal( loc=t[..., 1], scale=zero_buffer + tf.nn.softplus(t[..., 2]) ).experimental_extra_properties(q=t[..., 0]), # 从分布的额外属性中获取q值并计算分位数 convert_to_tensor_fn=lambda s: quantile(s, q=s.experimental_extra_properties['q']) )(input_layer) model = Model(input_layer, dist) # 测试输出 print(model(sample_input))
方式二:自定义Distribution子类
通过继承tfd.Normal创建自定义分布类,显式存储q值:
from functools import partial import tensorflow as tf from tensorflow.keras import layers, Model import tensorflow_probability as tfp from typing import Union tfd = tfp.distributions zero_buffer = 1e-5 class NormalWithQuantile(tfd.Normal): """自定义Normal分布类,附加q值属性""" def __init__(self, loc, scale, q, **kwargs): super().__init__(loc=loc, scale=scale, **kwargs) self.q = q def quantile(s: tfd.Distribution, q: Union[tf.Tensor, float]) -> Union[tf.Tensor, float]: return s.quantile(q) sample_input = tf.constant([[0.25, 0.0, 1.0], [0.5, 1.0, 0.5], [0.75, -1.0, 2.0], [0.95, 3.0, 2.5]], dtype=tf.float32) input_layer = layers.Input(3) dist = tfp.layers.DistributionLambda( # 创建自定义分布实例,传入q值 make_distribution_fn=lambda t: NormalWithQuantile( loc=t[..., 1], scale=zero_buffer + tf.nn.softplus(t[..., 2]), q=t[..., 0] ), # 直接访问分布的q属性计算分位数 convert_to_tensor_fn=lambda s: s.quantile(s.q) )(input_layer) model = Model(input_layer, dist) # 测试输出 print(model(sample_input))
两种方式运行后都会得到和模型外测试一致的输出:
tf.Tensor([-0.6744898 1. 0.3489796 7.112134 ], shape=(4,), dtype=float32)
内容的提问来源于stack exchange,提问作者Jed
相关产品推荐
相关产品推荐

