TensorFlow中实现等价于scipy.truncnorm.rvs的截断正态分布采样问询
在TensorFlow中实现等价于
scipy.stats.truncnorm.rvs()的截断正态分布采样 我刚巧在从NumPy/Scipy迁移到TensorFlow做大规模计算时遇到过一模一样的问题,拒绝采样在静态计算图里的局限性太头疼了——毕竟没法提前预知要多少次拒绝循环,完全没法静态构建计算图。所以当时我直接放弃了拒绝采样,用逆CDF变换法完美解决了问题,这方法不管是静态图还是动态图都能顺畅运行,和scipy的结果也完全对齐。
核心思路回顾
截断正态分布本质是把标准正态分布限制在[a, b](标准化后的截断边界)之间,再重新归一化概率。逆CDF法的逻辑就是:
- 先生成[0,1]区间的均匀分布样本
- 将这个均匀样本映射到截断正态分布的CDF值域范围内
- 对映射后的数值求标准正态CDF的逆,得到截断后的正态样本
- 如果需要非标准正态(带均值μ、标准差σ),再做线性变换
TensorFlow代码实现
下面是完整的实现函数,参数和scipy.stats.truncnorm.rvs()完全对应,方便直接替换:
import tensorflow as tf def truncnorm_rvs(loc=0.0, scale=1.0, a=-2.0, b=2.0, shape=(1,), dtype=tf.float32): """ 生成截断正态分布样本,等价于scipy.stats.truncnorm.rvs(a, b, loc, scale, size=shape) 参数: loc: 正态分布均值 scale: 正态分布标准差 a: 标准化后的下截断点 (对应实际边界: loc + a*scale) b: 标准化后的上截断点 (对应实际边界: loc + b*scale) shape: 输出样本的形状 dtype: 数据类型 """ # 计算标准正态CDF在a和b处的值 def norm_cdf(x): return 0.5 * (1.0 + tf.math.erf(x / tf.sqrt(2.0))) cdf_a = norm_cdf(a) cdf_b = norm_cdf(b) # 生成[0,1]均匀分布样本 uniform_samples = tf.random.uniform(shape=shape, minval=0.0, maxval=1.0, dtype=dtype) # 映射到截断正态的CDF区间 [cdf_a, cdf_b] u = cdf_a + (cdf_b - cdf_a) * uniform_samples # 求标准正态CDF的逆(分位数函数) def norm_ppf(u): return tf.sqrt(2.0) * tf.math.erfinv(2.0 * u - 1.0) truncated_norm_samples = norm_ppf(u) # 转换为非标准正态分布(如果需要) truncated_norm_samples = loc + scale * truncated_norm_samples return truncated_norm_samples
验证与使用示例
你可以用scipy的结果来对比验证,比如:
import scipy.stats as stats import numpy as np # TensorFlow生成样本 tf_samples = truncnorm_rvs(loc=5.0, scale=2.0, a=-1.0, b=1.0, shape=(10000,)).numpy() # Scipy生成样本 scipy_samples = stats.truncnorm.rvs(a=-1.0, b=1.0, loc=5.0, scale=2.0, size=10000) # 对比均值和方差 print(f"TF均值: {np.mean(tf_samples):.4f}, Scipy均值: {np.mean(scipy_samples):.4f}") print(f"TF方差: {np.var(tf_samples):.4f}, Scipy方差: {np.var(scipy_samples):.4f}")
输出应该几乎完全一致,误差来自随机采样的随机性。
静态计算图适配说明
这个方法完全不需要动态循环,所有操作都是纯张量运算,不管你是用TensorFlow 1.x的静态图构建,还是TensorFlow 2.x中用tf.function装饰成静态图函数,都能正常运行,完美解决了拒绝采样在静态图里的痛点。
内容的提问来源于stack exchange,提问作者tcquinn
相关产品推荐
相关产品推荐

