ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

Pyro 概率分布系统全解:PyTorch 封装、自定义分布、变换与约束的完整指南

Pyro 概率分布系统全解:PyTorch 封装、自定义分布、变换与约束的完整指南 人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载PyroDeep universal probabilistic programming with Python and PyTorch为贝叶斯建模与概率推理提供了一套完整的分布生态。本文以官方文档 docs/source/distributions.rst 为骨架结合仓库源码逐层拆解 Pyro 的分布系统从 PyTorch 分布的薄封装、Pyro 自研分布与扩展接口到可学习参数的变换Transform、变换工厂Transform Factories与约束Constraints体系。读完本文你将掌握 Pyro 分布模块的完整脉络能够在模型与推理代码中正确选用、组合甚至自行扩展分布与变换。一、分布系统总体架构三层结构从 pyro/distributions/init.py 的导入关系可以清楚看到 Pyro 分布体系分为三层PyTorch Distributions薄封装层大多数 Pyro 分布是对torch.distributions的轻量封装通过 pyro/distributions/torch.py 程序化地加载所有 PyTorch 分布并混入TorchDistributionMixin以兼容 Pyro 的接口约定Pyro Distributions自研分布层Pyro 在pyro/distributions/目录下实现的数十个专属分布覆盖 HMM 时序族、共轭族、零膨胀族、方向统计、稳定分布、拒绝采样等多个领域Transforms / Transform Modules / Constraints变换与约束层提供torch.distributions.transforms之上的扩展变换、带可学习参数的流式变换模块以及自定义约束。三个层在pyro.distributions命名空间下统一对外导出用户通过from pyro.distributions import dist即pyro.distributions模块即可访问全部分布、变换与约束。二、PyTorch Distributions薄封装层官方文档开篇即明确Pyro 中大多数分布是围绕 PyTorch 分布的薄封装两者接口的差异体现在TorchDistributionMixin中。从源码看pyro/distributions/torch.py 在模块加载时遍历torch.distributions.__dict__凡是torch.distributions.Distribution的子类都会被动态包装若 Pyro 已在locals()中定义了同名增强版本如Beta、Binomial、Categorical、Dirichlet、Gamma等则直接使用该版本否则动态创建type(_name, (_Dist, TorchDistributionMixin), {})即继承 PyTorch 分布 混入 Pyro Mixin并拼接两边的 docstring。__all__也因此包含Bernoulli、Normal、MultivariateNormal、StudentT、RelaxedBernoulli、VonMises、Wishart等全部 PyTorch 分布。增强点几个被 Pyro 覆写的关键分布虽然大部分分布只是薄封装但torch.py中若干分布做了实质性增强这些正是文档.. automodule:: pyro.distributions.torch会渲染出的内容Beta/Dirichlet/Gamma实现了实验性的conjugate_update()可将两个共轭分布融合为一个后验分布并给出对数归一化常数。例如Beta的实现满足concentration1 other.concentration1 - 1这样的参数合并规则且满足恒等式f.log_prob(x) g.log_prob(x) fg.log_prob(x) log_normalizerBinomial提供两个实验性阈值类属性approx_sample_thresh与approx_log_prob_tol。前者用于超大群体抽样时以矩匹配的截断 Poisson 近似替代精确二项抽样见sample()中对total_count approx_sample_thresh的分支后者用于log_prob()中启用移位 Stirling 近似把 3 次lgamma()计算降为 4 次log()推荐取值在 0.1~0.01 之间。这两个阈值在 pyro/settings.py 中注册为全局设置binomial_approx_sample_thresh、binomial_approx_log_prob_tolCategorical覆写log_prob()与enumerate_support()当枚举变量携带_pyro_categorical_support标记时直接在logits上做reshape/transpose而完全跳过torch.gather极大加速枚举enumeration场景下的对数概率计算LogNormal构造时以 Pyro 的Normal而非 PyTorch 的Normal作为基分布保证整个变换链都在 Pyro 生态内Poisson支持is_sparseTrue参数对稀疏观测做稀疏化的log_prob计算Uniform保留未广播的low/high从而给出准确的support interval(low, high)Independent实现conjugate_update()把基分布的共轭更新通过to_event()/sum_rightmost正确对齐。三、Pyro 分布基类体系3.1 抽象基类Distribution文档用autoclass列出 pyro/distributions.Distribution它是所有 Pyro 分布的抽象基类基于ABCMeta。核心约定如下分布即随机函数对象d dist.Bernoulli(param); x d(); p d.log_prob(x)。__call__只是sample(*args, **kwargs)的别名抽象方法派生类必须实现sample()与log_prob()score_parts(x)返回ScoreParts(log_prob, score_function, entropy_term)三元组是 SVI 等推理引擎计算 ELBO 随机梯度估计的成分。默认实现区分两种情形当has_rsample True时走重参数化路径score_function0entropy_termlog_prob否则走 score function 估计器score_functionlog_probentropy_term0。推理引擎如SVI正是依据.has_rsample决定使用重参数化采样器还是 score function 估计器enumerate_support()仅离散分布实现返回按第一个维度排列的支撑集注意它返回的是所有批量化随机变量锁步的支撑值而非笛卡尔积conjugate_update(other)实验性 API只有少数共轭分布支持返回(updated, log_normalizer)对has_rsample_(value)在单个实例上强制开启/关闭重参数化采样可用于指示推理算法对不连续决定下游控制流的变量避免重参数化梯度.rv属性实验性的随机变量 DSL 入口返回pyro.contrib.randomvariable.RandomVariable支持链式操作或运算符重载例如Uniform(0, 1).rv.log().neg().dist等价于一个Exponential分布。另外DistributionMeta元类在__call__时依次尝试全局COERCIONS钩子为未来扩展参数自动转换预留了机制。3.2TorchDistributionMixin与TorchDistribution这是文档中单列的两大核心类pyro/distributions/torch_distribution.pyTorchDistributionMixin给 PyTorch 分布提供 Pyro 兼容性的 Mixin主要用于包装既有 PyTorch 分布新分布类应当优先继承TorchDistribution。它带来以下 Pyro 专属能力__call__(sample_shape)能重参数化就调rsample否则调sampleshape(sample_shape)返回sample_shape batch_shape event_shapeevent_dimlen(event_shape)expand(batch_shape)与expand_by(sample_shape)前者把 batch 维从 1 扩到更大后者在 batch_shape 左侧追加 sample 维to_event(reinterpreted_batch_ndims)把最右侧 n 个 batch 维重新解释为 event 维负值可剥离Independent的维度旧的.reshape()已被拆分为.expand_by(...).to_event(...)mask(mask)返回MaskedDistribution这是 Pyro 实现pyro.mask的基础设施infer_shapes(**arg_shapes)类方法根据构造参数形状推断batch_shape与event_shape。TorchDistributiontorch.distributions.Distribution TorchDistributionMixin的组合基类文档明确这应当成为几乎所有新 Pyro 分布的基类。其 docstring 同时给出了实现新分布的完整契约派生类必须实现sample或rsample当has_rsample True与log_prob必须实现batch_shape、event_shape属性离散类可额外实现enumerate_support并设置has_enumerate_support True。3.3 形状语义sample / batch / eventTorchDistribution的 docstring 用三条规则定义了与 PyTorch 完全一致的形状语义这是理解 Pyro 一切分布的基础sample shapeiid 样本的维度由sample()的参数决定batch shape同一分布的不同独立参数化由参数形状推断对一个分布实例是固定的event shape单次事件的内在维度对分布类是固定的log_prob评分时事件维会被坍缩。三者满足恒等式assert d.shape(sample_shape) sample_shape d.batch_shape d.event_shape且向量化log_prob的返回形状为sample_shape d.batch_shape。文档中的示例同样验证了to_event的逐步转移d0.to_event(2)后batch_shape从[2,3,4,5]变为[2,3]event_shape变为[4,5]。四、Pyro 专属分布全景文档Pyro Distributions一节以autoclass形式列出了全部自研分布。按功能族归类如下全部可在 pyro/distributions/init.py 中找到对应导入与源文件4.1 时序与隐马尔可夫族HMMpyro/distributions/hmm.py提供一整套面向时序建模的分布DiscreteHMM离散状态隐马尔可夫模型initial_logitstransition_logits提供向量化前向算法与采样GaussianHMM高斯观测 HMM基于 pyro/ops/gaussian.py 的高斯对象与sequential_gaussian_filter_sample等运算实现线性高斯状态空间模型的滤波、采样GammaGaussianHMM观测为 Gamma 分布的共轭 HMM底层依赖 pyro/ops/gamma_gaussian.pyIndependentHMM时间步之间条件独立的 HMM 退化情形LinearHMM显式线性高斯状态转移的 HMMinit/trans/shift/covGaussianMRF高斯马尔可夫随机场分布按精度矩阵形式定义局部依赖结构。源码中_logmatmulexp、_sequential_logmatmulexp等内部函数展示了其在 log 空间数值稳定地做转移矩阵连乘的策略。对应测试可见 tests/distributions/test_hmm.py。4.2 共轭分布族BetaBinomial、DirichletMultinomial、GammaPoissonpyro/distributions/conjugate.py分别把 Beta、Dirichlet、Gamma 先验与二项、多项、Poisson 似然积分得到的边缘分布是层次贝叶斯模型中常用的一步到位分布InverseGammainverse_gamma.py当 PyTorch 尚无该分布时启用与 Gamma 互为倒数变换关系ExtendedBinomial、ExtendedBetaBinomialextended.py支持分数/连续化 total_count 的扩展二项分布LogNormalNegativeBinomiallog_normal_negative_binomial.py对数正态与负二项复合常用于过度离散计数建模。4.3 零膨胀Zero-Inflated族pyro/distributions/zero_inflated.py 提供ZeroInflatedDistribution通用零膨胀包装器接受任意单变量base_dist通过gate零膨胀概率或gate_logits其 logits二者必须二选一构造log_prob在value 0处合并结构零 基分布自身产生的 0两部分的概率sample先按bernoulli(gate)决定是否置零ZeroInflatedPoisson、ZeroInflatedNegativeBinomial以 Poisson、负二项为基分布的常用特化广泛用于保险、医学等含过量零的计数数据。4.4 混合分布族MixtureOfDiagNormalsdiag_normal_mixture.py对角高斯混合log_prob用 log-sum-exp 数值稳定计算MixtureOfDiagNormalsSharedCovariancediag_normal_mixture_shared_cov.py各分量共享协方差的对角高斯混合变体MaskedMixturemixture.py以布尔掩码选择两个分布之一用于if-else 分支的软建模GaussianScaleMixturegaussian_scale_mixture.py与GroupedNormalNormalgrouped_normal_normal.py高斯尺度混合与分组正态-正态层级结构服务于厚尾先验与多组方差建模。4.5 方向统计与循环分布VonMises3Dvon_mises_3d.py三维单位球面上的 von Mises–Fisher 分布用于方向数据SineBivariateVonMisessine_bivariate_von_mises.py双变量循环分布的 sine 参数化变体SineSkewedsine_skewed.py给任意圆上分布施加 sine 偏斜的包装ProjectedNormalprojected_normal.py把高维高斯投影到球面上得到的分布比 von Mises 更适合作为可重参数化的方向先验has_rsample True。4.6 稳定分布与重尾分布Stable与StableWithLogProbstable.pyα 稳定分布族支持重尾数据建模StableWithLogProb额外提供log_prob评分能力对应 stable_log_prob.py 的实现SoftLaplacesoftlaplace.py、AsymmetricLaplace/SoftAsymmetricLaplaceasymmetriclaplace.pyLaplace 及其软化处处可微变体常用于鲁棒回归Logistic/SkewLogisticlogistic.py、MultivariateStudentTmultivariate_studentt.py、AffineBetaaffine_beta.py亦属此列覆盖 S 型变换与多元重尾分布需求。4.7 匹配与组合优化分布OneOneMatchingone_one_matching.py1-1 完美匹配分布支持枚举/概率计算用于指派问题、数据关联OneTwoMatchingone_two_matching.py允许 1-2 匹配的扩展变体在 pyro/ops/arrowhead.py 基础上实现。 对应测试见 tests/distributions/test_one_one_matching.py、test_one_two_matching.py。4.8 退化、经验与特殊支撑分布Deltadelta.py退化点质量分布对应确定性变量的建模Empiricalempirical.py由一组样本及其对数权重构成的经验分布。其形状约定为log_weights的形状必须等于samples最左侧的形状样本沿log_weights的最右侧维aggregation_dim聚合sample_size返回样本数mean/variance为加权统计量整型样本会报错提示先转浮点向量化样本不能被log_prob评分——这正是pyro.infer.Predictive等场景的底层支撑ImproperUniformimproper_uniform.py非正常不可归一化均匀分布用于无信息先验Unitunit.py单元素支撑的平凡分布FoldedDistributionfolded.py把基分布的支撑折叠到非负区间如折叠正态OrderedLogisticordered_logistic.py有序类别逻辑回归分布CoalescentTimes/CoalescentTimesWithRatecoalescent.py群体遗传学中的溯祖时间分布SpanningTreespanning_tree.py随机生成树的分布配合 spanning_tree.cpp 的 C 扩展见 tests/distributions/test_spanning_tree.pyTruncatedPolyaGammapolya_gamma.py截断 Polya-Gamma 分布服务于贝叶斯 logistic 回归的数据增强LKJ/LKJCorrCholeskylkj.py相关矩阵与其 Cholesky 因子上的 LKJ 先验。4.9 梯度友好的采样包装Rejectorrejector.py通用拒绝采样分布给定提议分布propose、接受对数概率log_prob_accept与总接受对数概率log_scalehas_rsample True内部用 LRU(1) 缓存共享多次调用的工作量OMTMultivariateNormalomt_mvn.py基于 OMTOptimal Mass Transport的多元正态梯度方差通常更低代价是 Cholesky 因子梯度计算为 O(D³)AVFMultivariateNormalavf_mvn.py基于低秩扰动逼近的多元正态均摊方差因子AVF梯度估计RelaxedBernoulliStraightThrough/RelaxedOneHotCategoricalStraightThroughrelaxed_straight_through.pyGumbel-Softmax 的 straight-through 变体前向用离散样本、反向用松弛梯度MaskedDistribution与ExpandedDistributiontorch_distribution.py前者是.mask()的返回类型mask is False时log_prob、score_parts、kl_divergence全部短路为常量零值从而在效果上裁剪无关数据后者是.expand()/.expand_by()的返回类型精确记录扩张维度与插值维度保证采样与评分形状正确。4.10 条件分布Conditionalpyro/distributions/conditional.py 为以上下文为条件的分布提供抽象ConditionalDistribution抽象基类要求实现condition(context)返回一个普通torch.distributions.DistributionConditionalTransform/ConditionalTransformModule条件变换的抽象与带可学习参数版本ConditionalTransformedDistribution对条件基分布施加一系列条件变换condition(context)后即得到一个普通TransformedDistribution文件内的ConditionalFlowStack示例展示了如何组合多层conditional_planar流构建条件归一化流并用-cond_dist.condition(context).log_prob(data)计算负对数似然。五、Transforms变换系统文档Transforms一节列出的类位于 pyro/distributions/transforms/init.py。该模块同样采用from torch.distributions.transforms import * 自研扩展的策略PyTorch 的AffineTransform、ExpTransform、SigmoidTransform、ComposeTransform、StickBreakingTransform、CumulativeDistributionTransform等全部可用同时 Pyro 补充了CholeskyTransform/CorrMatrixCholeskyTransformcholesky.py下三角与相关矩阵的 Cholesky 变换DiscreteCosineTransformdiscrete_cosine.pyDCT 正交变换HaarTransformhaar.pyHaar 小波正交变换用于图像等结构化变量见 tests/distributions/test_haar.pyELUTransform/LeakyReLUTransformbasic.pyELU / LeakyReLU 双射LowerCholeskyAffinelower_cholesky_affine.py仿射 下三角耦合Normalizenormalize.py向量归一化到球面OrderedTransformordered.py把实向量映射为严格递增向量Permutepermute.py、PositivePowerTransformpower.py、SimplexToOrderedTransformsimplex_to_ordered.py、SoftplusTransform/SoftplusLowerCholeskyTransformsoftplus.py、UnitLowerCholeskyTransformunit_cholesky.py。此外transforms/__init__.py底部用transform_to.register/biject_to.register把自定义约束绑定到默认变换见下节例如constraints.sphere - Normalize()、constraints.corr_matrix - ComposeTransform([CorrCholeskyTransform(), CorrMatrixCholeskyTransform().inv])、constraints.ordered_vector - OrderedTransform()等。六、Transform Modules带可学习参数的流式变换文档TransformModules一节聚焦归一化流Normalizing Flows。其基类是 pyro/distributions/torch_transform.py 中的TransformModuletorch.distributions.Transform torch.nn.Module让变换参数可被 PyTorch 优化器与 Pyro 参数存储自动管理ComposeTransformModuleComposeTransform torch.nn.ModuleList让一系列TransformModule的参数在PyroModule中自动注册iterated()工厂即基于它构造深度流。具体模块包括每个都可在pyro/distributions/transforms/下找到源文件AffineAutoregressiveaffine_autoregressive.py逆自回归流IAF默认采用 Kingma et al. 2016 的式 (10)y μ σ⊙xstableTrue时改用y σ⊙x (1-σ)⊙μ以提升数值稳定性参数log_scale_min_clip、log_scale_max_clip、sigmoid_bias控制尺度裁剪。其 docstring 中给出了与AutoRegressiveNNpyro/nn/auto_reg_nn.py配合构建流式变分后验的完整示例base_dist dist.Normal(zeros(10), ones(10)); transform AffineAutoregressive(AutoRegressiveNN(10, [40])); flow_dist dist.TransformedDistribution(base_dist, [transform])AffineCouplingaffine_coupling.py、BlockAutoregressiveblock_autoregressive.py、BatchNormbatchnorm.py、Householderhouseholder.py、MatrixExponentialmatrix_exponential.py、NeuralAutoregressiveneural_autoregressive.py、Planarplanar.py、Radialradial.py、Splinespline.py、SplineAutoregressivespline_autoregressive.py、SplineCouplingspline_coupling.py、Sylvestersylvester.py、GeneralizedChannelPermutegeneralized_channel_permute.py、Polynomialpolynomial.py以及对应的Conditional*版本ConditionalAffineAutoregressive、ConditionalAffineCoupling、ConditionalHouseholder、ConditionalMatrixExponential、ConditionalNeuralAutoregressive、ConditionalPlanar、ConditionalRadial、ConditionalSpline、ConditionalSplineAutoregressive、ConditionalGeneralizedChannelPermute。七、Transform Factories小写辅助工厂函数文档Transform Factories一节解释了 Pyro 的一个重要设计每个Transform/TransformModule都配有对应的小写工厂函数其最低输入是变换的输入维度input_dim并可接收直观的附加参数。这些工厂函数的目的原文档表述向用户隐藏变换是否需要构建 hypernet超网络以及 hypernet 的输入/输出维度。例如spline(input_dim, ...)内部可能选择Spline无需超网络或SplineAutoregressive需要超网络用户无需关心区别。完整清单均位于 pyro/distributions/transforms/init.py 的__all__iterated、affine_autoregressive、affine_coupling、batchnorm、block_autoregressive、conditional_affine_autoregressive、conditional_affine_coupling、conditional_generalized_channel_permute、conditional_householder、conditional_matrix_exponential、conditional_neural_autoregressive、conditional_planar、conditional_radial、conditional_spline、conditional_spline_autoregressive、elu、generalized_channel_permute、householder、leaky_relu、matrix_exponential、neural_autoregressive、permute、planar、polynomial、radial、spline、spline_autoregressive、spline_coupling、sylvester。其中iterated(repeats, base_fn, *args, **kwargs)的实现transforms/__init__.py内即返回ComposeTransformModule([base_fn(*args, **kwargs) for _ in range(repeats)])用于快速堆叠多层可学习变换组成深度归一化流。八、Constraints约束系统文档末节.. automodule:: pyro.distributions.constraints指向 pyro/distributions/constraints.py。该模块extends torch.distributions.constraintsPyTorch 的real、positive、simplex、interval、lower_cholesky、corr_cholesky、independent、real_vector等全部可用Pyro 在此基础上新增了六类约束integer整数约束is_discrete Truesphere任意维欧氏球面check用相对容差10.0 * finfo.eps * size**0.5检验范数误差corr_matrix相关矩阵对角全 1 且正定ordered_vector沿 event 维严格递增的实向量positive_ordered_vector递增且元素为正的向量softplus_positive/softplus_lower_cholesky/unit_lower_cholesky分别对应 softplus 正数、softplus 下三角、单位对角下三角。这些约束与第五节末尾的注册逻辑联动当某个参数被声明为上述约束时Pyro 的transform_to/biject_to会自动选择默认变换把它映射到无约束空间供 HMC/NUTS 等推理算法使用。九、实战建议如何在模型中使用将上述体系落到实际建模中几个高频组合如下定义随机变量在模型函数内用pyro.sample(name, dist.Normal(loc, scale).to_event(1))声明多维变量to_event用于把 batch 维转为事件维使plate语义与评分正确利用mask与枚举dist.MaskedDistribution支撑pyro.mask实现缺失数据/观测掩码Categorical的enumerate_support快速路径与has_enumerate_support配合pyro.infer.enum做离散穷举参见 tests/infer/test_enum.py构造流式变分后验base_dist dist.Normal(...).to_event(1) 一组 Transform Module或工厂函数如dist.transforms.spline_autoregressive(input_dim)dist.TransformedDistribution即可作为AutoGuide之外的自定义变分族验证模式通过 pyro.distributions.util 导出的validation_enabled()/enable_validation()开关分布参数校验调试形状错误时尤其有用如Empirical.log_prob会拒绝带 sample_shape 的向量化输入新增分布优先继承TorchDistribution并实现sample/log_prob/batch_shape/event_shape离散分布再实现enumerate_support并置has_enumerate_support True需要低方差梯度时置has_rsample True并实现rsample。十、延伸阅读分布 API 骨架文档docs/source/distributions.rst分布实现与导出pyro/distributions/init.py、pyro/distributions/torch.py基类与 Mixinpyro/distributions/distribution.py、pyro/distributions/torch_distribution.py变换系统pyro/distributions/transforms/init.py、pyro/distributions/torch_transform.py约束系统pyro/distributions/constraints.py测试基准tests/distributions/、tests/distributions/dist_fixture.py涵盖形状、log_prob、KL、均值方差等一致性验证以及 tests/distributions/test_distributions.py、tests/distributions/test_transforms.py官方教程中的分布应用示例tutorial/source/gmm.ipynb、tutorial/source/hmm.rst、tutorial/source/stable.ipynb综上Pyro 的分布系统以 PyTorch 为底座、以TorchDistribution为统一接口、以变换与约束为建模翼构成了从简单先验到流式变分后验、从离散枚举到连续重参数化的完整概率建模工具箱。理解本指南中的分层结构与关键类契约即可在 Pyro 中游刃有余地选择与组合分布。赞分享人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载相关推荐PyMC 概率分布完全指南模型构建、自定义分布与自动变换PyMC 概率分布完全指南模型构建、自定义分布与自动变换 本文以 docs/source/guides/Probability_Distributions.r人工智能机器学习科学计算PyMC 分布变换Transforms完全指南约束空间与无约束空间的桥梁PyMC 分布变换Transforms完全指南约束空间与无约束空间的桥梁 导读 PyMC 的变换Transform机制是连接概率分布定义域约束空间人工智能机器学习科学计算掌握Android Sunflower布局约束Jetpack Compose中的权重分布终极指南掌握Android Sunflower布局约束Jetpack Compose中的权重分布终极指南 Sunflower是一款展示Android开发最佳实践的园艺移动开发示例工程上一篇Carota核心功能解析探索HTML Canvas上的富文本渲染技术下一篇Burn 的 LibTorch 后端burn-tch详解LibTorch 安装、CUDA/MPS 配置与张量实现原理创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表