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

TensorFlow Probability正态分布prob方法形状不兼容问题咨询

问题:TensorFlow Probability中正态分布prob方法的广播适配

我正在学习TensorFlow Probability教程,针对正态分布的μ和σ进行二维网格近似,想了解下述代码中广播机制的适用范围:

import tensorflow as tf
import tensorflow_probability as tfp
tfd = tfp.distributions

mu_list = tf.linspace(start=150, stop=160, num=100)
sigma_list = tf.linspace(start=7, stop=9, num=100)
mesh = tf.meshgrid(mu_list, sigma_list)
mu = tf.cast(tf.reshape(mesh[0], -1), tf.float32)
sig = tf.cast(tf.reshape(mesh[1], -1), tf.float32)
dists = tfd.Normal(loc=mu, scale=sig)
heights = tfd.Normal(loc=150, scale=20).sample(352)
dists.prob(heights)

运行代码时出现错误:

Incompatible shapes: [352] vs. [10000]

我确定可以通过tf.map_fn或tf.vectorized_map解决此问题,但好奇能否在.prob/.log_prob调用中直接生成形状为[10000, 352]的张量。


解决方法

可以通过调整张量维度对齐广播规则,直接生成目标形状的张量,无需循环类方法。

TensorFlow的广播规则要求从最后一个维度开始匹配,当某一维度长度为1时,会自动扩展至对应维度的长度。只需让mu/sig和heights的维度满足广播条件即可:

方法1:扩展mu和sig的维度

给mu和sig增加一个尾部维度,使其形状变为[10000, 1],这样与形状为[352]的heights广播后,就能得到[10000, 352]的结果:

import tensorflow as tf
import tensorflow_probability as tfp
tfd = tfp.distributions

mu_list = tf.linspace(start=150, stop=160, num=100)
sigma_list = tf.linspace(start=7, stop=9, num=100)
mesh = tf.meshgrid(mu_list, sigma_list)
mu = tf.cast(tf.reshape(mesh[0], -1), tf.float32)
sig = tf.cast(tf.reshape(mesh[1], -1), tf.float32)

# 增加尾部维度,适配广播
mu = tf.expand_dims(mu, axis=-1)
sig = tf.expand_dims(sig, axis=-1)

dists = tfd.Normal(loc=mu, scale=sig)
heights = tfd.Normal(loc=150, scale=20).sample(352)
prob_tensor = dists.prob(heights)
print(prob_tensor.shape)  # 输出 (10000, 352)

方法2:扩展heights的维度

保持mu和sig形状不变,给heights增加头部维度,使其形状变为[1, 352],同样能触发广播得到目标形状:

import tensorflow as tf
import tensorflow_probability as tfp
tfd = tfp.distributions

mu_list = tf.linspace(start=150, stop=160, num=100)
sigma_list = tf.linspace(start=7, stop=9, num=100)
mesh = tf.meshgrid(mu_list, sigma_list)
mu = tf.cast(tf.reshape(mesh[0], -1), tf.float32)
sig = tf.cast(tf.reshape(mesh[1], -1), tf.float32)

dists = tfd.Normal(loc=mu, scale=sig)
heights = tfd.Normal(loc=150, scale=20).sample(352)
# 增加头部维度,适配广播
heights = tf.expand_dims(heights, axis=0)

prob_tensor = dists.prob(heights)
print(prob_tensor.shape)  # 输出 (10000, 352)

广播原理说明

调整维度后,mu/sig的形状[10000, 1]与heights的形状[352](或[1, 352])会按照广播规则自动扩展:

  • heights会被扩展为[1, 352]
  • mu/sig会被扩展为[10000, 352]
    最终每个分布(共10000个)都会对每个样本(共352个)计算概率值,得到目标形状的张量。

内容的提问来源于stack exchange,提问作者ChrisS

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 16:20:55