☰
Pyro 杂项算子库(pyro.ops)完全指南:从 HMC 数值工具到高斯收缩与流式统计
2026/9/25 6:49:56 网站建设 项目流程
  • 人工智能
  • 机器学习
  • 深度学习
  • 概率编程

【免费下载链接】pyro

Deep 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嵌套元组索引与向量化广播索引vindex
Tensor Contractioneinsum 子包、contract.py基于opt_einsum的收缩路径缓存与带 plate 语义的einsum/ubersum
Gaussian 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 HMC:Hamiltonian 采样的三个数值支柱

对偶平均(Dual Averaging)步长自适应

pyro/ops/dual_averaging.py 中的DualAveraging实现了 Nesterov 的对偶平均方案,用于在 HMC/NUTS 采样过程中自适应地调节 leapfrog 步长,使其逼近目标接受率。其核心思路是:普通次梯度方法中"新次梯度权重递减"(如 Nesterov 论文 [1] 所述),而对偶平均以相等权重累加对偶空间中的统计量,从而保证收敛。

构造函数与参数(默认值直接来自 dual_averaging.py#L43):

  • prox_center=0:prox 中心,把原始序列拉向该点;
  • t0=10:稳定方案初始步的自由参数(来自 Hoffman & Gelman 的 NUTS 论文 [2]);
  • kappa=0.75:控制每步权重的参数,取值范围(0.5, 1],取值越小方案越快遗忘早期状态;
  • gamma=0.05:控制收敛速度的自由参数。

使用方式非常简单——每获得一个新统计量就调用一次step,随时可用get_state取出最新值:

from pyro.ops.dual_averaging import DualAveraging adapt_scheme = DualAveraging(prox_center=0, t0=10, kappa=0.75, gamma=0.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/(t+t0))、原始序列_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_steps=1, # 积分步数 z_grads=None, # 可选的当前梯度缓存,避免重复计算 )

每个单步_single_step_verlet按标准三分量更新:先半步更新动量r(n+1/2),再整步更新位置z(n+1),最后半步更新动量r(n+1),实现相空间体积守恒(辛性)。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 matrix):

  • WelfordCovariance(diagonal=True):经典 Welford 在线算法(Knuth《计算机程序设计艺术》[1]),用增量方式更新均值与二阶中心矩,避免两遍扫描。update(sample)每来一个样本更新一次;get_covariance(regularize=True)返回协方差。注意少于 2 个样本时抛RuntimeError。
  • WelfordArrowheadCovariance(head_size=0):箭形(arrowhead)结构协方差,head_size指定头部大小,返回(top, bottom_diag)两部分,用于块对角/箭形质量矩阵的快速表示。

regularize=True时采用与 Stan 一致的正则化:scaled_cov = n/(n+5) * cov,并在对角上加1e-3 * 5/(n+5)的收缩项,保证协方差正定。

from pyro.ops.welford import WelfordCovariance adapt = WelfordCovariance(diagonal=True) 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_radius=None)对一批小维数变量执行一步牛顿更新,返回(mode, cov):

  • loss是x的二阶可微标量函数;把loss解释为负对数密度时,(mode, cov)可直接构造 Laplace 近似MultivariateNormal(mode, cov);
  • x形状为(N, D),D只支持 1、2、3(分别派发到newton_step_1d/2d/3d),cov形状为x.shape[:-1] + (D, D);
  • trust_radius可选,用于把更新约束在信任域球内(2D/3D 通过最小特征值正则化 Hessian 实现,1D 通过 clamp 实现);
  • 由于牛顿迭代的二次收敛性,最终解对输入可微——即使中间步骤全部detach,只要loss是2+d阶可微的,返回值就是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_radius=1.0) # 最终的 x 仍然是可微的

实现细节上,1D 版本直接clamp(min=1e-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, tol=0.0):log-Beta 函数。当tol < 0.02时直接退化为torch.lgamma组合(更便宜);当tol >= 0.02时使用移位 Stirling 近似,迭代ceil(0.082/tol)次把绝对误差压到tol以内。
  • log_binomial(n, k, tol=0.0):log 二项式系数,同样支持高容差下的近似模式;小容差模式为n_plus_1.lgamma() - (k+1).lgamma() - (n_plus_1-k).lgamma()。注意其被@torch.no_grad()修饰。
  • log_I1(orders, value, terms=250):第一类修正贝塞尔函数的前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, axis=0).exp() assert (variance - 16.0).abs().item() < 1.0e-6
  • sparse_multinomial_likelihood(total_count, nonzero_logits, nonzero_value):稀疏多项式对数似然,只对非零位置求值,等价于稠密Multinomial(logits=logits).log_prob(value).sum(),但可避免构造超大的全 logits 向量;内部用带weakref的缓存_log_factorial_cache记忆(x+1).lgamma().sum(),避免重复计算。

Tensor Utilities:FFT、DCT、Haar 与安全线性代数

pyro/ops/tensor_utils.py 是纯张量层工具,服务于时序模型、高斯过程与概率线性代数:

  • FFT 卷积:next_fast_len(size)返回不小于size的"快速长度"(素因子仅为 2/3/5,等价scipy.fftpack.next_fast_len),convolve(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_period=365.25、min_period=7,短时间尺度交给periodic_repeat/periodic_cumsum。
  • 正交变换:dct/idct是缩放为正交的 II 型离散余弦变换(等价scipy.fftpack.dctwithnorm="ortho");haar_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, p=2)把零向量映射到[1,0,...,0]以避免球面投影的奇异点;precision_to_scale_tril(P)从精度矩阵求尺度下三角矩阵;triangular_solve等函数统一在事件维为 1 时退化为逐元素运算,避免不必要的矩阵分支。
  • 张量组装:block_diag_embed/block_diagonal完成块对角矩阵的嵌入与还原;repeated_matmul(M, n)用对数并行的倍增法一次性返回M, M^2, ..., M^n;as_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_einsum

einsum 子包:带缓存的收缩路径

pyro/ops/einsum/init.py 提供contract(equation, *operands)与contract_expression(equation, *shapes),是opt_einsum的薄封装,默认开启收缩路径缓存(cache_path=True,全局_PATH_CACHE),同一 equation+shape 组合只计算一次最优路径。子包还包含torch_log.py(log-space einsum)、torch_map.py(map-reduce)、torch_sample.py(采样)等后端实现。

contract.py:einsum与ubersum

pyro/ops/contract.py 在opt_einsum之上叠加了 Pyro 的 plate(plate/iarange)语义:

  • einsum(equation, *operands):标准 einsum,但对每个操作数接受一个可选的ordinal(frozenset 的 plate 帧集合)元参数,用于描述张量所在的最小 plate 上下文,从而在收缩时自动广播。输出形状还受dims(求和维集合)与target_dims(需保留的求和维)参数控制。
  • ubersum(equation, *operands):ubersum("uber 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_sample(Kalman 平滑采样)。

文档对该类启用了:special-members: __add__,__getitem__,说明这两个魔术方法是官方 API 的一部分。这些工具被 pyro/distributions/hmm.py 的线性高斯 HMM 与 contrib/timeseries 的时序模型作为底层代数引擎使用。

Statistical Utilities:MCMC 诊断与预测评分

pyro/ops/stats.py 实现贝叶斯分析的标准诊断与评分指标,chain_dim/sample_dim参数支持负索引:

  • 收敛诊断:gelman_rubin(input, chain_dim=0, sample_dim=1)计算 R-hat(要求两维都 ≥2);split_gelman_rubin把每个链切成两半再算 R-hat(要求sample_dim >= 4);autocorrelation/autocovariance用 FFT 加速;effective_sample_size基于 Geyer 的初始单调序列估计计算 ESS。
  • 后验汇总:quantile、pi(百分位区间)、hpdi(最高后验密度区间)、resample(重采样)、weighed_quantile(带对数权重的分位数)。
  • 模型选择/预测评分:waic(input, log_weights, pointwise=False, dim=0)(Widely Applicable Information Criterion,用_weighted_mean/_weighted_variance计算);crps_empirical(pred, truth)(连续排序概率分数);energy_score_empirical(能量分数,支持pred_batch_size分块与自定义cdist);fit_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_rubin(api.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与other;
  • get():返回聚合结果。

内置实现:

类统计内容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={...}, default=CountMeanVarianceStats)收集各 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(nu=1.5, num_gps=1, length_scale_init=None, kernel_scale_init=None):

  • nu:Matérn 平滑度参数(取值需为半整数,如0.5, 1.5, 2.5, ...),决定状态维数q = nu + 0.5;
  • num_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 采样器时,直接用DualAveraging+velocity_verlet+WelfordCovariance组装出"步长自适应 + 质量矩阵估计"的完整管线;
  • 做 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 推理算法组合,即可搭建从采样、诊断到预测的完整贝叶斯工作流。

  • 人工智能
  • 机器学习
  • 深度学习
  • 概率编程

【免费下载链接】pyro

Deep universal probabilistic programming with Python and PyTorch

项目地址:https://gitcode.com/gh_mirrors/py/pyro
点击查看免费下载
上一篇:Cronos Rootkit 安装与使用指南
下一篇:StringManipulation 插件使用教程

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询