Pyro 杂项算子库(pyro.ops)完全指南:从 HMC 数值工具到高斯收缩与流式统计

发布时间:2026/9/25 6:50:02
Pyro 杂项算子库(pyro.ops)完全指南:从 HMC 数值工具到高斯收缩与流式统计 人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载Pyro 的pyro.ops模块实现了一整套与概率编程主体解耦的张量数值工具是 HMC/NUTS 采样器、高斯消息传递、时序模型与诊断统计的底层支撑。本文将基于 docs/source/ops.rst 的模块划分逐一讲解每一类算子的核心接口、数学原理与典型应用场景并结合仓库源码给出可直接运行的使用示例帮助你把这些高性能数值原语用于自己的贝叶斯建模与研究工作中。模块总览一张 Pyro 数值原语的地图官方文档 ops.rst 将pyro.ops划分为 10 个主题对应仓库目录 pyro/ops 下的 10 个模块文件外加一个einsum子包文档小节对应源码模块定位Utilities for HMCdual_averaging.py、integrator.py、welford.pyHMC/NUTS 的步长自适应、辛积分与质量矩阵估计Newton Optimizersnewton.py可微分的牛顿优化步Laplace 近似Special Functionsspecial.py数值稳定的特殊函数log-Beta、log-二项式系数、修正贝塞尔函数、Gauss–Hermite 积分等Tensor Utilitiestensor_utils.py张量级数、FFT 卷积、DCT/Haar 变换、安全 Cholesky 等Tensor Indexingindexing.py嵌套元组索引与向量化广播索引vindexTensor Contractioneinsum 子包、contract.py基于opt_einsum的收缩路径缓存与带 plate 语义的einsum/ubersumGaussian Contractiongaussian.py非归一化高斯信息形式的收缩、条件、边缘化Statistical Utilitiesstats.pyMCMC 诊断R-hat、ESS、自相关、WAIC、CRPS 等Streaming Statisticsstreaming.py可合并跨链聚合的流式统计量State Space Model and GP Utilitiesssm_gp.py状态空间模型与高斯过程的离散化工具模块的设计哲学在文档开头已点明这些工具“mostly independent of the rest of Pyro”大部分独立于 Pyro 的其他部分这意味着你可以在自己的 PyTorch 项目里直接 import 使用而不必引入整套概率编程框架。Utilities for HMCHamiltonian 采样的三个数值支柱对偶平均Dual Averaging步长自适应pyro/ops/dual_averaging.py 中的DualAveraging实现了 Nesterov 的对偶平均方案用于在 HMC/NUTS 采样过程中自适应地调节 leapfrog 步长使其逼近目标接受率。其核心思路是普通次梯度方法中新次梯度权重递减如 Nesterov 论文 [1] 所述而对偶平均以相等权重累加对偶空间中的统计量从而保证收敛。构造函数与参数默认值直接来自 dual_averaging.py#L43prox_center0prox 中心把原始序列拉向该点t010稳定方案初始步的自由参数来自 Hoffman Gelman 的 NUTS 论文 [2]kappa0.75控制每步权重的参数取值范围(0.5, 1]取值越小方案越快遗忘早期状态gamma0.05控制收敛速度的自由参数。使用方式非常简单——每获得一个新统计量就调用一次step随时可用get_state取出最新值from pyro.ops.dual_averaging import DualAveraging adapt_scheme DualAveraging(prox_center0, t010, kappa0.75, gamma0.05) for t in range(500): stat some_hmc_acceptance_statistic() # 例如 log(step_size * accept_prob) adapt_scheme.step(stat) x_t, x_avg adapt_scheme.get_state() # 步长自适应常取 x_avg平均序列它对早期状态更鲁棒从实现看step内部维护了对偶序列均值_g_avg权重1/(tt0)、原始序列_x_t prox_center - sqrt(t)/gamma * g_avg以及带权重t^{-kappa}的滑动平均_x_avg。该类的实际消费者是 pyro/infer/mcmc/adaptation.py#L47即 HMC/NUTS 的StepSizeAdaptation。Velocity Verlet 辛积分器pyro/ops/integrator.py 的velocity_verlet实现了二阶辛积分是 HMC 中 leapfrog 数值积分的中枢z_next, r_next, z_grads, potential_energy velocity_verlet( z, # dict采样点名字 - 位置张量 r, # dict采样点名字 - 动量张量 potential_fn, # 势能函数如负对数后验 kinetic_grad, # 动能对动量的梯度 step_size, # 步长 num_steps1, # 积分步数 z_gradsNone, # 可选的当前梯度缓存避免重复计算 )每个单步_single_step_verlet按标准三分量更新先半步更新动量r(n1/2)再整步更新位置z(n1)最后半步更新动量r(n1)实现相空间体积守恒辛性。potential_grad负责在z上开requires_grad_、计算grad(potential_energy, z_nodes)并返回(梯度 dict, 势能标量)。值得一提的细节是异常处理机制模块维护了一个全局注册表_EXCEPTION_HANDLERS默认注册了torch_singular处理器见 integrator.py#L119它会捕获矩阵奇异/不正定等RuntimeError把势能置为nan并返回零梯度从而避免 HMC 轨迹因数值奇异性直接崩溃。你还可以用register_exception_handler注册自己的处理器。该函数被 hmc.py#L191 与 nuts.py#L199 直接调用构成了 Pyro HMC/NUTS 后端的物理引擎。Welford 在线协方差估计pyro/ops/welford.py 提供两个类用于在采样过程中在线估计 HMC 的质量矩阵mass matrixWelfordCovariance(diagonalTrue)经典 Welford 在线算法Knuth《计算机程序设计艺术》[1]用增量方式更新均值与二阶中心矩避免两遍扫描。update(sample)每来一个样本更新一次get_covariance(regularizeTrue)返回协方差。注意少于 2 个样本时抛RuntimeError。WelfordArrowheadCovariance(head_size0)箭形arrowhead结构协方差head_size指定头部大小返回(top, bottom_diag)两部分用于块对角/箭形质量矩阵的快速表示。regularizeTrue时采用与 Stan 一致的正则化scaled_cov n/(n5) * cov并在对角上加1e-3 * 5/(n5)的收缩项保证协方差正定。from pyro.ops.welford import WelfordCovariance adapt WelfordCovariance(diagonalTrue) for z in samples: adapt.update(z) cov adapt.get_covariance() # 对角协方差n 2该实现被 adaptation.py#L301 用于 HMC 的对角/稠密质量矩阵自适应也被 streaming.py 的CountMeanVarianceStats复用streaming.py#L11。Newton Optimizers可微分的牛顿步与 Laplace 近似pyro/ops/newton.py 的newton_step(loss, x, trust_radiusNone)对一批小维数变量执行一步牛顿更新返回(mode, cov)loss是x的二阶可微标量函数把loss解释为负对数密度时(mode, cov)可直接构造 Laplace 近似MultivariateNormal(mode, cov)x形状为(N, D)D只支持 1、2、3分别派发到newton_step_1d/2d/3dcov形状为x.shape[:-1] (D, D)trust_radius可选用于把更新约束在信任域球内2D/3D 通过最小特征值正则化 Hessian 实现1D 通过 clamp 实现由于牛顿迭代的二次收敛性最终解对输入可微——即使中间步骤全部detach只要loss是2d阶可微的返回值就是d阶可微的这是该实现可微分优化用法的理论根基见源码 docstring 引用的 Christianson 1994。文档给出的优化循环示例强调一个关键陷阱——迭代中间必须 detach否则反向传播会贯穿整个迭代过程x torch.zeros(1000, 2) # 任意初始值 for step in range(100): x x.detach() # 阻断对上一轮梯度的传播 x.requires_grad True loss my_loss_function(x) x newton_step(loss, x, trust_radius1.0) # 最终的 x 仍然是可微的实现细节上1D 版本直接clamp(min1e-8)保证 Hessian 逆非负2D/3D 版本用pyro.ops.linalg.rinverse对称矩阵伪逆求逆并用warn_if_nan监控梯度与 Hessian 的数值健康。3D 的最小特征值来自 linalg.py 的eig_3d闭式解。Special Functions数值稳定的特殊函数集合pyro/ops/special.py 为贝叶斯计算中常见的数值困难场景提供稳定实现safe_log(x)与torch.log等价但把log(0)处的梯度钳制在至多1/finfo.eps避免反向传播中出现无穷梯度自定义torch.autograd.Function实现见 special.py#L15。log_beta(x, y, tol0.0)log-Beta 函数。当tol 0.02时直接退化为torch.lgamma组合更便宜当tol 0.02时使用移位 Stirling 近似迭代ceil(0.082/tol)次把绝对误差压到tol以内。log_binomial(n, k, tol0.0)log 二项式系数同样支持高容差下的近似模式小容差模式为n_plus_1.lgamma() - (k1).lgamma() - (n_plus_1-k).lgamma()。注意其被torch.no_grad()修饰。log_I1(orders, value, terms250)第一类修正贝塞尔函数的前orders阶对数截断到terms项求和用于 Von Mises 等环形分布的归一化计算。get_quad_rule(num_quad, prototype_tensor)基于numpy.polynomial.hermite.hermgauss的 Gauss–Hermite 求积点与对数权重返回张量会继承prototype_tensor的dtype/device。文档自带示例quad_points, log_weights get_quad_rule(32, prototype_tensor) quad_points * 4.0 # 变换到 N(0, 4.0) variance torch.logsumexp(quad_points.pow(2.0).log() log_weights, axis0).exp() assert (variance - 16.0).abs().item() 1.0e-6sparse_multinomial_likelihood(total_count, nonzero_logits, nonzero_value)稀疏多项式对数似然只对非零位置求值等价于稠密Multinomial(logitslogits).log_prob(value).sum()但可避免构造超大的全 logits 向量内部用带weakref的缓存_log_factorial_cache记忆(x1).lgamma().sum()避免重复计算。Tensor UtilitiesFFT、DCT、Haar 与安全线性代数pyro/ops/tensor_utils.py 是纯张量层工具服务于时序模型、高斯过程与概率线性代数FFT 卷积next_fast_len(size)返回不小于size的快速长度素因子仅为 2/3/5等价scipy.fftpack.next_fast_lenconvolve(signal, kernel, mode)用rfft/irfft实现 1D 卷积支持full/valid/same三种模式并自动做零填充对齐。周期特征periodic_repeat、periodic_cumsum分别支持静态季节性与漂移季节性的时间序列构造periodic_features(duration, max_period, min_period)生成(duration, 2*ceil(max_period/min_period)-2)形状的 sin/cos 回归特征归一化到[-1,1]。文档给出的组合用法示例回归年季节性时设max_period365.25、min_period7短时间尺度交给periodic_repeat/periodic_cumsum。正交变换dct/idct是缩放为正交的 II 型离散余弦变换等价scipy.fftpack.dctwithnormorthohaar_transform/inverse_haar_transform是沿最后一维的 Haar 小波变换。这些变换支撑了 reparam/haar.py 等重参数化策略。安全线性代数safe_cholesky依据cholesky_relative_jitter设置可通过 settings.py 的cholesky_relative_jitter调节默认 4.0 倍finfo.eps在 Cholesky 分解前加自适应抖动safe_normalize(x, p2)把零向量映射到[1,0,...,0]以避免球面投影的奇异点precision_to_scale_tril(P)从精度矩阵求尺度下三角矩阵triangular_solve等函数统一在事件维为 1 时退化为逐元素运算避免不必要的矩阵分支。张量组装block_diag_embed/block_diagonal完成块对角矩阵的嵌入与还原repeated_matmul(M, n)用对数并行的倍增法一次性返回M, M^2, ..., M^nas_complex是torch.view_as_complex的 stride 安全版本broadcast_tensors_without_dim在保持指定维尺寸不变的前提下广播其余维度。Tensor Indexing兼容标量/向量/枚举语义的索引pyro/ops/indexing.py 解决概率编程中一个非常实际的问题同一份索引代码要同时兼容标量求值、向量化求值和 reshape。index(tensor, args)与Index包装器把嵌套元组索引展平并合并连续的Ellipsis。文档例子要泛化x[..., t]其中t可能是标量1、切片slice(None)或 reshape 操作(Ellipsis, None)等价x.unsqueeze(-1)。Index(x)[..., i, j, :]与index(x, (Ellipsis, i, j, slice(None)))等价。vindex(tensor, args)与Vindex包装器带广播语义的向量化高级索引特别适合从离散随机变量中选择混合成分。与 NumPy NEP-21 建议略有不同Pyro 约定Ellipsis只能出现在最左侧表示未知 batch 维。例如x事件维为 3 时xij Vindex(x)[..., i, :, j] # ... 表示未知的 batch 形状 # new_batch_shape broadcast_shape(old_batch_shape, i.shape, j.shape) # new_event_shape (x.size(1),)约束条件源码明确声明每个参数只能是Ellipsis、slice(None)、整数或带空事件维的torch.LongTensor不支持非平凡切片与BoolTensormask非前导的Ellipsis直接抛NotImplementedError。当所有参数都不是多维张量时vindex与标准索引完全一致。该工具被离散枚举推理如 enum.py广泛使用是混合模型向量化求值的关键。Tensor Contraction带 plate 语义的 opt_einsumeinsum 子包带缓存的收缩路径pyro/ops/einsum/init.py 提供contract(equation, *operands)与contract_expression(equation, *shapes)是opt_einsum的薄封装默认开启收缩路径缓存cache_pathTrue全局_PATH_CACHE同一 equationshape 组合只计算一次最优路径。子包还包含torch_log.pylog-space einsum、torch_map.pymap-reduce、torch_sample.py采样等后端实现。contract.pyeinsum与ubersumpyro/ops/contract.py 在opt_einsum之上叠加了 Pyro 的 plateplate/iarange语义einsum(equation, *operands)标准 einsum但对每个操作数接受一个可选的ordinalfrozenset 的 plate 帧集合元参数用于描述张量所在的最小 plate 上下文从而在收缩时自动广播。输出形状还受dims求和维集合与target_dims需保留的求和维参数控制。ubersum(equation, *operands)ubersumuber einsum支持稠密与稀疏枚举两种模式可以在枚举与 vectorized plate 之间转换。还有naive_ubersum作为参考实现。内部流程_partition_terms把项与求和维建成二分图并按连通分量分组避免不必要的广播_contract_component通过消息传递把树上的张量逐步降维。_check_plates_are_sensible保证保留 plate 维时必须保留其全部 plate的语义正确性_check_tree_structure拒绝非树形嵌套的 plate 依赖。底层代数由 rings.py 的LogRing等环结构提供sumproduct/product/inv后端映射见BACKEND_TO_RING。这些函数由 enum.py 与 traceenum ELBO 在离散变分推断/精确边缘化中调用是 Pyro 枚举plate混合推理的数学引擎。Gaussian Contraction信息形式的非归一化高斯pyro/ops/gaussian.py 的Gaussian类用信息形式表示任意半正定二次函数即秩亏的缩放高斯分布Gaussian(log_normalizer, info_vec, precision)其中info_vec precision mean信息向量precision是精度矩阵。之所以不用(mean, cov)而用(info_vec, precision)是因为精度矩阵可以有零特征值秩亏此时协方差根本不存在但信息形式下收缩、条件化等操作仍然快速且数值稳定注释NB: using info_vec instead of mean to deal with rank-deficient problem见 gaussian.py#L38。核心 API 一览形状工具dim()、batch_shape三个字段广播、expand、reshape、__getitem__索引 batch 维组装静态方法cat(parts, dim)沿 batch 维拼接、event_pad(left, right)沿事件维填充、event_permute(perm)置换事件维代数__add__/__sub__在信息空间叠加二次型即贝叶斯相乘、log_density(value)、rsample()、condition(value)/left_condition条件化、marginalize(left, right)边缘化、event_logsumexp工厂函数mvn_to_gaussian、matrix_and_mvn_to_gaussian、gaussian_tensordot、sequential_gaussian_tensordot线性高斯序列收缩用于 HMM 前向滤波、sequential_gaussian_filter_sampleKalman 平滑采样。文档对该类启用了:special-members: __add__,__getitem__说明这两个魔术方法是官方 API 的一部分。这些工具被 pyro/distributions/hmm.py 的线性高斯 HMM 与 contrib/timeseries 的时序模型作为底层代数引擎使用。Statistical UtilitiesMCMC 诊断与预测评分pyro/ops/stats.py 实现贝叶斯分析的标准诊断与评分指标chain_dim/sample_dim参数支持负索引收敛诊断gelman_rubin(input, chain_dim0, sample_dim1)计算 R-hat要求两维都 ≥2split_gelman_rubin把每个链切成两半再算 R-hat要求sample_dim 4autocorrelation/autocovariance用 FFT 加速effective_sample_size基于 Geyer 的初始单调序列估计计算 ESS。后验汇总quantile、pi百分位区间、hpdi最高后验密度区间、resample重采样、weighed_quantile带对数权重的分位数。模型选择/预测评分waic(input, log_weights, pointwiseFalse, dim0)Widely Applicable Information Criterion用_weighted_mean/_weighted_variance计算crps_empirical(pred, truth)连续排序概率分数energy_score_empirical能量分数支持pred_batch_size分块与自定义cdistfit_generalized_pareto广义 Pareto 拟合用于重要性采样诊断。这些指标被 MCMC 后处理实际消费pyro.infer.mcmc.util的get_model_chain_options与_print_summary计算每个 site 的n_eff stats.effective_sample_size(...)与r_hat stats.split_gelman_rubin(...)见 mcmc/util.py#L525并在MCMC.summary()的 docstring 中推荐使用effective_sample_size与split_gelman_rubinapi.py#L635。from pyro.ops import stats rhat stats.gelman_rubin(samples) # samples: (chain, sample, ...) ess stats.effective_sample_size(samples)Streaming Statistics可跨链合并的流式统计pyro/ops/streaming.py 定义StreamingStats抽象基类用于对张量树做流式统计聚合核心是三个抽象方法update(sample)从单个样本更新状态原地修改、要求样本可交换顺序无关merge(other)合并两个聚合统计例如来自不同 MCMC 链纯函数——返回新对象、不改动self与otherget()返回聚合结果。内置实现类统计内容get()返回CountStats样本计数{count: n}StackStats样本堆叠{count: n, last: tensor}CountMeanStats计数与均值{count, mean}CountMeanVarianceStats计数、均值、方差内部用WelfordCovariance{count, mean, variance}StatsOfDict字典聚合器每个 key 绑定一个子统计类型types参数default指定未知 key 的类型字典StatsOfDict是 MCMC 并行链汇总的核心MCMC在summary()时用StatsOfDict(types{...}, defaultCountMeanVarianceStats)收集各 site 的样本统计再通过merge把不同链的结果合并见 api.py#L774然后交给stats.py计算 R-hat 与 ESS。from pyro.ops.streaming import CountMeanVarianceStats acc CountMeanVarianceStats() for sample in samples: acc.update(sample) summary acc.get() # {count: n, mean: ..., variance: ...}State Space Model and GP Utilities高斯过程的 SDE 离散化pyro/ops/ssm_gp.py 的MaternKernel类把 Matérn 核的高斯过程转成线性状态空间模型SSM从而把 GP 的时间复杂度从立方降为线性。构造函数MaternKernel(nu1.5, num_gps1, length_scale_initNone, kernel_scale_initNone)nuMatérn 平滑度参数取值需为半整数如0.5, 1.5, 2.5, ...决定状态维数q nu 0.5num_gps并行高斯过程数量length_scale_init/kernel_scale_init长度尺度与核幅度的初始值。核心方法transition_matrix(dt)给定时间步长dt的状态转移矩阵A指数矩阵stationary_covariance()平稳协方差process_covariance(A)给定转移矩阵的扩散协方差transition_matrix_and_covariance(dt)一次返回(A, Q)。该工具由 contrib/gp 的变分 GP 与时序模型如 contrib/timeseries/lgssmgp.py使用把连续时间 Matérn 过程离散化为线性高斯 SSM再与 gaussian.py 的sequential_gaussian_tensordot配合做精确滤波与平滑。在 Pyro 之外使用这些原语由于pyro.ops的设计目标是与框架解耦你可以把任意几个模块单独摘出来做自定义 HMC 采样器时直接用DualAveragingvelocity_verletWelfordCovariance组装出步长自适应 质量矩阵估计的完整管线做 Laplace 近似与可微分优化时直接调用newton_step处理时序数据时用periodic_features/periodic_repeat/periodic_cumsum构造季节特征用convolve/dct/haar_transform做谱分析需要跨链合并统计时用StreamingStats家族 stats的gelman_rubin/effective_sample_size完成 MCMC 收敛诊断实现自己的消息传递/因子图推理时contract.ubersum与Gaussian提供了 log-space 与信息形式两种收缩代数。如需进一步验证接口细节可直接阅读对应源码模块及配套测试例如 tests/ops/test_gaussian.py、tests/ops/test_contract.py、tests/ops/test_stats.py、tests/ops/test_integrator.py 覆盖了上述算子的数值正确性HMC 相关工具在 tests/infer/mcmc 下有端到端验证。将这些高性能原语与 Pyro 的 poutine 效应处理器、infer 推理算法组合即可搭建从采样、诊断到预测的完整贝叶斯工作流。赞分享人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载相关推荐Pyro 的 MCMC 推理引擎从 MCMC/NUTS/HMC 内核到诊断与流式采样Pyro 的 MCMC 推理引擎从 MCMC/NUTS/HMC 内核到诊断与流式采样 PyroDeep universal probabilistic pr人工智能机器学习深度学习概率编程Scala数值计算新范式Breeze库完全指南Scala数值计算新范式Breeze库完全指南 还在为Scala中的数值计算和线性代数操作而烦恼吗面对复杂的矩阵运算、优化算法和统计计算你是否渴望一个既高科学计算机器学习NumCpp完全实战指南从零基础到高效数值计算NumCpp完全实战指南从零基础到高效数值计算 在C开发中你是否曾经羡慕Python开发者能够轻松使用NumPy进行复杂的数值计算现在有了NumCp上一篇Cronos Rootkit 安装与使用指南下一篇StringManipulation 插件使用教程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询