news 2026/9/25 3:41:54

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Pyro 概率分布系统全解:PyTorch 封装、自定义分布、变换与约束的完整指南
  • 人工智能
  • 机器学习
  • 深度学习
  • 概率编程

【免费下载链接】pyro

Deep universal probabilistic programming with Python and PyTorch

项目地址:https://gitcode.com/gh_mirrors/py/pyro
点击查看免费下载

Pyro(Deep universal probabilistic programming with Python and PyTorch)为贝叶斯建模与概率推理提供了一套完整的分布生态。本文以官方文档 docs/source/distributions.rst 为骨架,结合仓库源码逐层拆解 Pyro 的分布系统:从 PyTorch 分布的薄封装、Pyro 自研分布与扩展接口,到可学习参数的变换(Transform)、变换工厂(Transform Factories)与约束(Constraints)体系。读完本文,你将掌握 Pyro 分布模块的完整脉络,能够在模型与推理代码中正确选用、组合甚至自行扩展分布与变换。

一、分布系统总体架构:三层结构

从 pyro/distributions/init.py 的导入关系可以清楚看到 Pyro 分布体系分为三层:

  1. PyTorch Distributions(薄封装层):大多数 Pyro 分布是对torch.distributions的轻量封装,通过 pyro/distributions/torch.py 程序化地加载所有 PyTorch 分布,并混入TorchDistributionMixin以兼容 Pyro 的接口约定;
  2. Pyro Distributions(自研分布层):Pyro 在pyro/distributions/目录下实现的数十个专属分布,覆盖 HMM 时序族、共轭族、零膨胀族、方向统计、稳定分布、拒绝采样等多个领域;
  3. 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_normalizer;
  • Binomial:提供两个实验性阈值类属性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_tol;
  • Categorical:覆写log_prob()与enumerate_support(),当枚举变量携带_pyro_categorical_support标记时,直接在logits上做reshape/transpose而完全跳过torch.gather,极大加速枚举(enumeration)场景下的对数概率计算;
  • LogNormal:构造时以 Pyro 的Normal而非 PyTorch 的Normal作为基分布,保证整个变换链都在 Pyro 生态内;
  • Poisson:支持is_sparse=True参数,对稀疏观测做稀疏化的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_function=0,entropy_term=log_prob);否则走 score function 估计器(score_function=log_prob,entropy_term=0)。推理引擎(如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.py):

  • TorchDistributionMixin:给 PyTorch 分布提供 Pyro 兼容性的 Mixin,主要用于包装既有 PyTorch 分布;新分布类应当优先继承TorchDistribution。它带来以下 Pyro 专属能力:
    • __call__(sample_shape):能重参数化就调rsample,否则调sample;
    • shape(sample_shape):返回sample_shape + batch_shape + event_shape;
    • event_dim:len(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。
  • TorchDistribution:torch.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 / event

TorchDistribution的 docstring 用三条规则定义了与 PyTorch 完全一致的形状语义,这是理解 Pyro 一切分布的基础:

  • sample shape:iid 样本的维度,由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 时序与隐马尔可夫族(HMM)

pyro/distributions/hmm.py提供一整套面向时序建模的分布:

  • DiscreteHMM:离散状态隐马尔可夫模型,initial_logits+transition_logits,提供向量化前向算法与采样;
  • GaussianHMM:高斯观测 HMM,基于 pyro/ops/gaussian.py 的高斯对象与sequential_gaussian_filter_sample等运算实现线性高斯状态空间模型的滤波、采样;
  • GammaGaussianHMM:观测为 Gamma 分布的共轭 HMM,底层依赖 pyro/ops/gamma_gaussian.py;
  • IndependentHMM:时间步之间条件独立的 HMM 退化情形;
  • LinearHMM:显式线性高斯状态转移的 HMM(init/trans/shift/cov);
  • GaussianMRF:高斯马尔可夫随机场分布,按精度矩阵形式定义局部依赖结构。

源码中_logmatmulexp、_sequential_logmatmulexp等内部函数展示了其在 log 空间数值稳定地做转移矩阵连乘的策略。对应测试可见 tests/distributions/test_hmm.py。

4.2 共轭分布族

  • BetaBinomial、DirichletMultinomial、GammaPoisson(pyro/distributions/conjugate.py):分别把 Beta、Dirichlet、Gamma 先验与二项、多项、Poisson 似然积分得到的边缘分布,是层次贝叶斯模型中常用的一步到位分布;
  • InverseGamma(inverse_gamma.py):当 PyTorch 尚无该分布时启用,与 Gamma 互为倒数变换关系;
  • ExtendedBinomial、ExtendedBetaBinomial(extended.py):支持分数/连续化 total_count 的扩展二项分布;
  • LogNormalNegativeBinomial(log_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 混合分布族

  • MixtureOfDiagNormals(diag_normal_mixture.py):对角高斯混合,log_prob用 log-sum-exp 数值稳定计算;
  • MixtureOfDiagNormalsSharedCovariance(diag_normal_mixture_shared_cov.py):各分量共享协方差的对角高斯混合变体;
  • MaskedMixture(mixture.py):以布尔掩码选择两个分布之一,用于"if-else 分支"的软建模;
  • GaussianScaleMixture(gaussian_scale_mixture.py)与GroupedNormalNormal(grouped_normal_normal.py):高斯尺度混合与分组正态-正态层级结构,服务于厚尾先验与多组方差建模。

4.5 方向统计与循环分布

  • VonMises3D(von_mises_3d.py):三维单位球面上的 von Mises–Fisher 分布,用于方向数据;
  • SineBivariateVonMises(sine_bivariate_von_mises.py):双变量循环分布的 sine 参数化变体;
  • SineSkewed(sine_skewed.py):给任意圆上分布施加 sine 偏斜的包装;
  • ProjectedNormal(projected_normal.py):把高维高斯"投影"到球面上得到的分布,比 von Mises 更适合作为可重参数化的方向先验(has_rsample = True)。

4.6 稳定分布与重尾分布

  • Stable与StableWithLogProb(stable.py):α 稳定分布族,支持重尾数据建模;StableWithLogProb额外提供log_prob评分能力(对应 stable_log_prob.py 的实现);
  • SoftLaplace(softlaplace.py)、AsymmetricLaplace/SoftAsymmetricLaplace(asymmetriclaplace.py):Laplace 及其"软化"(处处可微)变体,常用于鲁棒回归;
  • Logistic/SkewLogistic(logistic.py)、MultivariateStudentT(multivariate_studentt.py)、AffineBeta(affine_beta.py)亦属此列,覆盖 S 型变换与多元重尾分布需求。

4.7 匹配与组合优化分布

  • OneOneMatching(one_one_matching.py):1-1 完美匹配分布(支持枚举/概率计算),用于指派问题、数据关联;
  • OneTwoMatching(one_two_matching.py):允许 1-2 匹配的扩展变体(在 pyro/ops/arrowhead.py 基础上实现)。 对应测试见 tests/distributions/test_one_one_matching.py、test_one_two_matching.py。

4.8 退化、经验与特殊支撑分布

  • Delta(delta.py):退化点质量分布,对应确定性变量的建模;
  • Empirical(empirical.py):由一组样本及其对数权重构成的经验分布。其形状约定为:log_weights的形状必须等于samples最左侧的形状,样本沿log_weights的最右侧维(aggregation_dim)聚合;sample_size返回样本数,mean/variance为加权统计量(整型样本会报错提示先转浮点);向量化样本不能被log_prob评分——这正是pyro.infer.Predictive等场景的底层支撑;
  • ImproperUniform(improper_uniform.py):非正常(不可归一化)均匀分布,用于无信息先验;
  • Unit(unit.py):单元素支撑的平凡分布;
  • FoldedDistribution(folded.py):把基分布的支撑折叠到非负区间(如折叠正态);
  • OrderedLogistic(ordered_logistic.py):有序类别逻辑回归分布;
  • CoalescentTimes/CoalescentTimesWithRate(coalescent.py):群体遗传学中的溯祖时间分布;
  • SpanningTree(spanning_tree.py):随机生成树的分布(配合 spanning_tree.cpp 的 C++ 扩展,见 tests/distributions/test_spanning_tree.py);
  • TruncatedPolyaGamma(polya_gamma.py):截断 Polya-Gamma 分布,服务于贝叶斯 logistic 回归的数据增强;
  • LKJ/LKJCorrCholesky(lkj.py):相关矩阵与其 Cholesky 因子上的 LKJ 先验。

4.9 梯度友好的采样包装

  • Rejector(rejector.py):通用拒绝采样分布,给定提议分布propose、接受对数概率log_prob_accept与总接受对数概率log_scale,has_rsample = True;内部用 LRU(1) 缓存共享多次调用的工作量;
  • OMTMultivariateNormal(omt_mvn.py):基于 OMT(Optimal Mass Transport)的多元正态,梯度方差通常更低,代价是 Cholesky 因子梯度计算为 O(D³);
  • AVFMultivariateNormal(avf_mvn.py):基于低秩扰动逼近的多元正态,均摊方差因子(AVF)梯度估计;
  • RelaxedBernoulliStraightThrough/RelaxedOneHotCategoricalStraightThrough(relaxed_straight_through.py):Gumbel-Softmax 的 straight-through 变体,前向用离散样本、反向用松弛梯度;
  • MaskedDistribution与ExpandedDistribution(torch_distribution.py):前者是.mask()的返回类型,mask is False时log_prob、score_parts、kl_divergence全部短路为常量零值,从而在效果上"裁剪"无关数据;后者是.expand()/.expand_by()的返回类型,精确记录扩张维度与插值维度,保证采样与评分形状正确。

4.10 条件分布(Conditional)

pyro/distributions/conditional.py 为"以上下文为条件的分布"提供抽象:

  • ConditionalDistribution:抽象基类,要求实现condition(context),返回一个普通torch.distributions.Distribution;
  • ConditionalTransform/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/CorrMatrixCholeskyTransform(cholesky.py):下三角与相关矩阵的 Cholesky 变换;
  • DiscreteCosineTransform(discrete_cosine.py):DCT 正交变换;
  • HaarTransform(haar.py):Haar 小波正交变换,用于图像等结构化变量(见 tests/distributions/test_haar.py);
  • ELUTransform/LeakyReLUTransform(basic.py):ELU / LeakyReLU 双射;
  • LowerCholeskyAffine(lower_cholesky_affine.py):仿射 + 下三角耦合;
  • Normalize(normalize.py):向量归一化到球面;
  • OrderedTransform(ordered.py):把实向量映射为严格递增向量;
  • Permute(permute.py)、PositivePowerTransform(power.py)、SimplexToOrderedTransform(simplex_to_ordered.py)、SoftplusTransform/SoftplusLowerCholeskyTransform(softplus.py)、UnitLowerCholeskyTransform(unit_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 中的:

  • TransformModule:torch.distributions.Transform + torch.nn.Module,让变换参数可被 PyTorch 优化器与 Pyro 参数存储自动管理;
  • ComposeTransformModule:ComposeTransform + torch.nn.ModuleList,让一系列TransformModule的参数在PyroModule中自动注册,iterated()工厂即基于它构造深度流。

具体模块包括(每个都可在pyro/distributions/transforms/下找到源文件):

  • AffineAutoregressive(affine_autoregressive.py):逆自回归流(IAF),默认采用 Kingma et al. 2016 的式 (10)y = μ + σ⊙x,stable=True时改用y = σ⊙x + (1-σ)⊙μ以提升数值稳定性;参数log_scale_min_clip、log_scale_max_clip、sigmoid_bias控制尺度裁剪。其 docstring 中给出了与AutoRegressiveNN(pyro/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]);
  • AffineCoupling(affine_coupling.py)、BlockAutoregressive(block_autoregressive.py)、BatchNorm(batchnorm.py)、Householder(householder.py)、MatrixExponential(matrix_exponential.py)、NeuralAutoregressive(neural_autoregressive.py)、Planar(planar.py)、Radial(radial.py)、Spline(spline.py)、SplineAutoregressive(spline_autoregressive.py)、SplineCoupling(spline_coupling.py)、Sylvester(sylvester.py)、GeneralizedChannelPermute(generalized_channel_permute.py)、Polynomial(polynomial.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.constraints":PyTorch 的real、positive、simplex、interval、lower_cholesky、corr_cholesky、independent、real_vector等全部可用,Pyro 在此基础上新增了六类约束:

  • integer:整数约束(is_discrete = True);
  • sphere:任意维欧氏球面,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 等推理算法使用。

九、实战建议:如何在模型中使用

将上述体系落到实际建模中,几个高频组合如下:

  1. 定义随机变量:在模型函数内用pyro.sample("name", dist.Normal(loc, scale).to_event(1))声明多维变量,to_event用于把 batch 维转为事件维,使plate语义与评分正确;
  2. 利用mask与枚举:dist.MaskedDistribution支撑pyro.mask实现缺失数据/观测掩码;Categorical的enumerate_support快速路径与has_enumerate_support配合pyro.infer.enum做离散穷举,参见 tests/infer/test_enum.py;
  3. 构造流式变分后验:base_dist = dist.Normal(...).to_event(1)+ 一组 Transform Module(或工厂函数,如dist.transforms.spline_autoregressive(input_dim))+dist.TransformedDistribution,即可作为AutoGuide之外的自定义变分族;
  4. 验证模式:通过 pyro.distributions.util 导出的validation_enabled()/enable_validation()开关分布参数校验,调试形状错误时尤其有用(如Empirical.log_prob会拒绝带 sample_shape 的向量化输入);
  5. 新增分布:优先继承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
  • 基类与 Mixin:pyro/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 中游刃有余地选择与组合分布。

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

【免费下载链接】pyro

Deep universal probabilistic programming with Python and PyTorch

项目地址:https://gitcode.com/gh_mirrors/py/pyro
点击查看免费下载

相关推荐

上一篇:Carota核心功能解析:探索HTML Canvas上的富文本渲染技术
下一篇:Burn 的 LibTorch 后端(burn-tch)详解:LibTorch 安装、CUDA/MPS 配置与张量实现原理

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/25 3:40:48

ESP32驱动墨水屏实战:GxEPD2库入门与避坑指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/25 3:38:59

SpringBoot+Vue二手交易系统:从业务建模到部署上线全解析

最近我把一套基于 SpringBoot Vue 的二手物品交易管理系统重新翻了出来&#xff0c;项目代号 bootpf&#xff0c;代码包名统一叫 com.bootpf。这套系统从用户注册、商品发布、浏览搜索、购物车、下订单&#xff0c;到后台的商品审核、用户管理和数据统计&#xff0c;基本把二手…

作者头像 李华
网站建设 2026/9/25 3:37:33

STM32入门第0集:从芯片认知到环境搭建,新手避坑指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/25 3:37:08

华为昇腾推理引擎开源:边缘AI部署与性能优化实战

1. 昇腾推理引擎开源这件事&#xff0c;到底在解决什么问题第一次在昇腾社区看到推理引擎开源的消息时&#xff0c;我正蹲在一个边缘计算项目上折腾模型部署。当时手里的活儿是把一个视觉检测模型塞进一台功耗受限的工控机里&#xff0c;芯片用的是昇腾310P3。那会儿最头疼的不…

作者头像 李华