Pyro 分布测试工具指南:Goodness-of-Fit 拟合优度检验(gof 模块深度解析)
发布时间:2026/9/25 3:46:22来源:尧图网络
人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载本篇技术指南围绕 Pyro 官方文档中的Testing Utilities一节docs/source/testing.rst展开系统讲解 Pyro 用于验证分布实现正确性的拟合优度检验工具pyro.distributions.testing.gof。你将掌握如何在 Pyro 中为自定义分布编写“样本与密度一致”的统计检验、理解 Pearson 卡方检验与近邻距离检验的底层原理并学会参考仓库真实测试用例为自己的分布做质量把关。一、为什么 Pyro 需要拟合优度检验概率编程库的核心资产是大量可采样、可求密度的分布对象。任何一个新实现的分布其sample()与log_prob()必须自洽从sample()得到的样本代入log_prob()后对应的概率密度应与样本的聚集模式一致。若两者不一致后续的变分推断、MCMC 或期望计算都会得出错误结果。Pyro 在pyro/distributions/testing子包中提供了专门的测试工具集见init.py其定位是“测试新分布的工具以及用于测试新推断算法的分布”。其中核心模块 gof.py 实现了一整套拟合优度Goodness of Fit检验用于检查分布的sample()与log_prob()方法之间的一致性。该模块的统计思想是对任意一个实现正确的分布检验返回的 p 值gof在好数据下应当服从Uniform(0,1)分布在坏数据样本与密度不符下则应趋近于 0。这样测试者只需设定一个允许的“误报率”阈值并断言gof TEST_FAILURE_RATE即可在大量随机测试中稳定地捕捉实现错误。二、核心用法一份可复制的检验模板模块文档给出的标准用法如下摘自 gof.py 模块 docstringTEST_FAILURE_RATE 1 / 20 # For 1 in 20 chance of spurious failure. def test_my_distribution(): d MyDistribution() samples d.sample([10000]) probs d.log_prob(samples).exp() gof auto_goodness_of_fit(samples, probs) assert gof TEST_FAILURE_RATE要点拆解samples与probs是形状一致的张量probs由log_prob(samples).exp()得到即样本点上的真实概率密度auto_goodness_of_fit会自动根据样本的“环境维度”ambient dimension选择适当的检验策略TEST_FAILURE_RATE是一个全局配置建议设为小于测试套件总数分之一从而把整批测试的“误报率”控制在可接受范围内断言条件gof TEST_FAILURE_RATE的含义是p 值足够大样本与密度显著一致才通过测试否则判定分布实现有问题。2.1 仓库真实阈值2e-5Pyro 仓库自身在 tests/common.py 中定义了统一的失败率阈值TEST_FAILURE_RATE 2e-5 # For all goodness-of-fit tests.这个值远小于1/20意味着仓库内所有依赖拟合优度的测试合计允许约两万分之一的发生概率误报属于相当严格的设定。在 tests/distributions/test_stable_log_prob.py 中还有针对特定分布的更严格阈值TEST_FAILURE_RATE 5e-4。实践建议阈值的选择应与测试的样本量、测试次数相协调——样本越多p 值对偏差越敏感可以放心使用更小的阈值。三、检验工具族六大函数全景gof模块的检验函数构成了一个“递进式”的流水线从最底层的多项分布卡方检验到一维/多维连续密度检验最终汇入自动分派的auto_goodness_of_fit。3.1 multinomial_goodness_of_fit最底层的 Pearson 卡方检验这是整个模块的统计内核见 gof.py。它实现的是Pearson 卡方检验且支持“截断数据”truncated data参数类型说明probstorch.Tensor概率向量一维countstorch.Tensor各类别观测频数向量一维形状与probs相同total_countint或None可选的总计数。为None表示数据未截断总数为counts.sum()提供时表示真实试验总数total_count counts.sum()必须成立plotbool是否打印直方图默认False核心计算流程源码层面对每个类别计算期望均值mean total_count * p与方差variance total_count * p * (1 - p)累计统计量chi_squared (c - mean) ** 2 / variance自由度dof 1方差约束若variance 1会抛出InvalidTest异常并提示 “Goodness of fit is inaccurate; use more samples” —— 即样本太少时检验本身不可信拒绝给出结果若某个类别概率p ≈ 1abs(p - 1) 1e-8直接返回1 if c total_count else 0若某个类别概率为 0 却观测到计数c 0警告并返回math.inf该数据在该分布下几乎不可能出现p 值视为无穷大——仍然能通过gof TEST_FAILURE_RATE的断言但实际含义是“极端不符”未截断时自由度要减 1dof - 1因为总计数已知会消耗一个自由度最终survival chi2sf(chi_squared, dof)返回卡方分布的生存函数值作为 p 值。3.2 unif01_goodness_of_fit均匀样本的自动分箱检验见 gof.py。它接收一个应服从Uniform(0,1)的样本向量将其自动分箱后调用multinomial_goodness_of_fit分箱数取round(len(samples) ** 0.333)立方根启发式若少于 7 箱则抛出InvalidTest(imprecise test, use more samples)通过samples.mul(bin_count).long().clamp(min0, maxbin_count - 1)分箱并用scatter_add_统计频数由于箱子等宽各箱先验概率相等即probs ones(bin_count) / bin_count。3.3 exp_goodness_of_fit指数样本的变换见 gof.py。利用概率积分变换将服从Exponential(1)的样本经samples.neg().exp()变换为Uniform(0,1)样本再交由unif01_goodness_of_fit处理。这是连接“任意连续分布”与“均匀检验”的桥梁。3.4 density_goodness_of_fit一维连续密度检验见 gof.py。接收一维实值样本samples及其密度probs核心思路是把“密度不均匀”的样本转换为“密度均匀”的指数样本若样本数不足 100len(samples) 100抛出InvalidTest(imprecision; use more samples)对样本排序计算相邻样本间距gaps用相邻点密度的均值近似“稀疏度”sparsity 0.5 * (1/probs[1:] 1/probs[:-1])进而得到局部密度估计density len(samples) / sparsity构造“指数样本”exp_samples density * gaps若原分布实现正确这些间距经局部密度缩放后应服从Exponential(1)交给exp_goodness_of_fit完成最终检验。3.5 vector_density_goodness_of_fit多维连续密度检验近邻法多维情形无法直接排序取间距见 gof.py。该函数基于最近邻距离分布Nearest Neighbour Distribution理论将多维样本变换到一维指数分布样本量下限更严格len(samples) 1000 * dim时抛出InvalidTest其中dim为数据所在流形维度默认取samples.shape[-1]调用get_nearest_neighbor_distances计算每个点到其最近邻的距离见 gof.py若环境安装了scipy使用scipy.spatial.cKDTree获得 O(N log N) 复杂度否则回退到torch.cdist的 O(N²) 暴力实现用球体积公式volume_of_sphere(dim, radii) radius**dim * pi**(0.5*dim) / gamma(0.5*dim 1)将最近邻距离转换为体积构造指数样本exp_samples density * volume其中density len(samples) * probs最终交由exp_goodness_of_fit检验。该方法的理论依据是 Bickel Breiman (1983) 关于最近邻距离函数和的矩界与大数定律以及 Mike Williams (2010) 在高能物理中提出的非分箱多维拟合优度检验思路。3.6 auto_goodness_of_fit自动分派入口见 gof.py。这是绝大多数测试场景的首选入口根据样本的“环境维度”自动路由samples samples.reshape(samples.shape[0], -1) ambient_dim samples.shape[1:].numel() if dim is None: dim ambient_dim if ambient_dim 0: return 1.0 if ambient_dim 1: return density_goodness_of_fit(samples, probs, plotplot) return vector_density_goodness_of_fit(samples, probs, dimdim, plotplot)参数类型说明samplestorch.Tensor样本张量约定堆叠在最左侧batch维度上probstorch.Tensor各样本处的概率密度向量形状与samples.shape[:1]一致dimint或None数据所在的流形维度默认取samples.shape[1:].numel()环境维度plotbool是否打印直方图默认False特殊处理逻辑值得注意纯标量样本ambient_dim 0直接返回1.0表示恒通过一维样本展平后走density_goodness_of_fit多维样本dim参数支持“流形维度”与“环境维度”分离——例如ProjectedNormal这类定义在高维球面上的分布样本实际上位于维度减一的子流形上此时可显式传入dim samples.size(-1) - 1提高检验精度仓库测试 test_distributions.py 正是这么做的。四、实战用 gof 校验自定义分布假设你为 Pyro 实现了一个新的分布MyDistribution希望验证其采样与密度的一致性import torch from pyro.distributions.testing.gof import auto_goodness_of_fit TEST_FAILURE_RATE 2e-5 # 与仓库 tests/common.py 保持一致 def test_my_distribution(): d MyDistribution(alphatorch.tensor(1.5)) num_samples 50000 samples d.sample(torch.Size([num_samples])) probs d.log_prob(samples).exp() # 多维分布按 batch 逐个检验 probs probs.reshape(num_samples, -1) samples samples.reshape(probs.shape d.event_shape) for b in range(probs.size(-1)): gof auto_goodness_of_fit(samples[:, b], probs[:, b]) assert gof TEST_FAILURE_RATE上述写法直接对应仓库中对所有连续分布执行的通用检验 test_gof仓库以num_samples 50000采样对每个 batch 独立检验并针对特殊分布做了适配——例如 Dirichlet 分布截掉最后一个自由度samples samples[..., :-1]LKJ 系列因“子流形缩放不正确”而标记为xfail。调试小技巧将plotTrue传入任意 gof 函数会打印一张 ASCII 直方图Prob/Count两列条宽上限为 60见 gof.py便于肉眼观察样本-密度失配发生在哪个概率区间。4.1 离散分布的检验模板对于离散分布仓库使用multinomial_goodness_of_fit直接做卡方检验示例见 test_gof.pyimport torch import torch.distributions as dist from pyro.distributions.testing.gof import multinomial_goodness_of_fit N, K 100000, 20 logits torch.randn(K) probs (logits - logits.logsumexp(-1)).exp() d dist.Categorical(probs) samples d.sample((N,)) counts torch.zeros(K, dtypetorch.long) counts.scatter_add_(0, samples, torch.ones(N, dtypetorch.long)) gof multinomial_goodness_of_fit(probs, counts, plotTrue) assert gof 0.1注意此处采样量N100000、类别数K20每个类别的期望频数total_count * p足够大能通过函数内部的方差约束variance 1检验结果才可靠。五、底层原理卡方生存函数的纯 Python 实现multinomial_goodness_of_fit返回的 p 值来自 special.py 中自实现的卡方分布生存函数chi2sf不依赖 scipychi2sf(x, s)见 special.py计算1 - CDF其中 CDF 由不完全伽马函数与伽马函数的比值给出F(x; s) gamma(x/2, s/2) / Gamma(s/2)生存概率即1 - F(x; s)incomplete_gamma(x, s)见 special.py采用下不完全伽马函数的级数展开gamma(x, s) x^s * Gamma(s) * e^{-x} * Σ x^k / Gamma(s k 1)求和 100 项。由于伽马函数阶乘级增长级数强烈收敛为避免数值溢出全程在对数域进行中间运算log_gamma_s、log_num - log_denom后取exp并在x 1e3时直接返回math.gamma(s)作为极限值源码注释声明该实现与 scipy 的数值结果一致到机器精度matches the results from scipy to numerical precision。之所以要自实现而不是依赖 scipy是为了让 gof 模块保持轻量依赖仅torch与标准库这也使得整个检验工具在无 scipy 环境下依然可运行——唯一受影响的是多维近邻距离计算会退化为 O(N²) 的torch.cdist实现。六、局限性与使用注意事项使用本模块时以下几点是源码中明确给出的边界条件样本量不足时拒绝给出结论一维密度检验要求样本数大于 100多维要求大于1000 * dim均匀分箱检验要求箱子数不少于 7多项检验要求各类别方差大于 1。不满足时统一抛出InvalidTestValueError子类见 gof.py这是“拒绝检验”而非“检验失败”提示测试者应增大采样量p 值的统计语义gof在实现正确时服从Uniform(0,1)因此即便实现无误小概率也会出现低 p 值导致误报——这正是TEST_FAILURE_RATE需要按测试总数合理设定的原因流形维度参数对定义在低维子流形上的分布如球面分布、单纯形分布务必正确传递dim否则检验会因“环境维度”与“流形维度”不一致而产生偏差。仓库对 LKJ 系列分布显式标记xfail“incorrect submanifold scaling”正是这一类问题的典型代表概率为 0 的类别观测到计数此时函数返回math.infp 值语义为“该数据在模型下极端不可能”。它仍能通过gof TEST_FAILURE_RATE断言需要结合业务判断这是合理的拒绝采样行为还是实现缺陷。七、在 Pyro 仓库中的真实应用拟合优度检验并非文档中的“纸上谈兵”而是 Pyro 回归测试体系的有机组成部分tests/distributions/test_distributions.py 中的test_gof对所有连续分布逐一执行auto_goodness_of_fit是分布正确性的第一道防线tests/distributions/test_stable_log_prob.py 用它同时校验 Pyro 采样与 scipy 参考采样的自洽性tests/distributions/test_von_mises.py 对环形分布dim1与球面分布dim2显式传入流形维度做检验tests/distributions/test_spanning_tree.py 对生成树离散分布使用multinomial_goodness_of_fit阈值由 tests/common.py 统一为2e-5。配合 docs/source/testing.rst 文档页该页通过 Sphinx 的automodule指令自动呈现pyro.distributions.testing.gof的全部成员读者可直接查看文档与源码互相对照学习。结语Pyro 的拟合优度检验工具以“p 值应服从 Uniform(0,1)”这一简洁思想为核心构建了一条从离散卡方检验、一维密度检验到多维近邻距离检验的完整链路。无论是为 Pyro 贡献新的分布还是排查已有分布在极端参数下的实现缺陷auto_goodness_of_fit与TEST_FAILURE_RATE这套组合都提供了统计上严谨、工程上易用的标准答案。掌握它就等于掌握了 Pyro 分布质量验证的通行证。赞分享人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载相关推荐Haystack CacheChecker 元数据命中检测从 items 到 Document Store 过滤的调用链拆解Haystack CacheChecker 元数据命中检测从 items 到 Document Store 过滤的调用链拆解 索引管道每次重跑都在重复转换、拆人工智能大模型预训练微调LoRARLHF强化学习分布式训练模型推理服务推理引擎模型量化模型压缩本地部署NLPInsightFace人脸检测模块深度解析InsightFace人脸检测模块深度解析 InsightFace项目集成了多种先进的人脸检测算法包括RetinaFace、SCRFD和BlazeFace等人工智能计算机视觉深度学习SciPy 卡方检验实战指南用 scipy.stats.chisquare 做拟合优度检验SciPy 卡方检验实战指南用 scipy.stats.chisquare 做拟合优度检验 卡方检验Chi square test是统计学中最常用的分类数科学计算数据科学高性能计算上一篇HFS2模板系统详解自定义界面与高级功能扩展指南下一篇XO与Next.js 16路线图未来Next.js支持的完整指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网