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
相关产品推荐
相关产品推荐

