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

TensorFlow中实现等价于scipy.truncnorm.rvs的截断正态分布采样问询

在TensorFlow中实现等价于scipy.stats.truncnorm.rvs()的截断正态分布采样

我刚巧在从NumPy/Scipy迁移到TensorFlow做大规模计算时遇到过一模一样的问题,拒绝采样在静态计算图里的局限性太头疼了——毕竟没法提前预知要多少次拒绝循环,完全没法静态构建计算图。所以当时我直接放弃了拒绝采样,用逆CDF变换法完美解决了问题,这方法不管是静态图还是动态图都能顺畅运行,和scipy的结果也完全对齐。

核心思路回顾

截断正态分布本质是把标准正态分布限制在[a, b](标准化后的截断边界)之间,再重新归一化概率。逆CDF法的逻辑就是:

  1. 先生成[0,1]区间的均匀分布样本
  2. 将这个均匀样本映射到截断正态分布的CDF值域范围内
  3. 对映射后的数值求标准正态CDF的逆,得到截断后的正态样本
  4. 如果需要非标准正态(带均值μ、标准差σ),再做线性变换

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 09:00:30