You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.16 23:20:29