- 人工智能
- 机器学习
- 深度学习
- 概率编程
【免费下载链接】pyro
Deep universal probabilistic programming with Python and PyTorch
导读
本文基于 Pyro 官方教程文档 Example: Amortized Latent Dirichlet Allocation 及其对应示例 examples/lda.py,完整讲解如何在 Pyro 中实现摊销化的 Latent Dirichlet Allocation(LDA)主题模型。核心技术亮点在于:通过 Pyro 的离散变量并行枚举机制(infer={"enumerate": "parallel"}配合TraceEnum_ELBO)将词-主题指派变量word_topics在模型内直接边缘化,从而让 guide 无需建模这一离散隐变量;同时利用 PyTorch 可重参数化的 Gamma/Dirichlet 分布获得路径梯度,配合共轭 guide 与 MLP 摊销 guide 完成高效的变分推断。读完本文,你将掌握 LDA 在 Pyro 中的完整建模、枚举式边缘化推断、摊销推理以及命令行训练流程。
一、问题背景:为什么在 Pyro 中实现 LDA 需要"枚举"
Latent Dirichlet Allocation 是一个经典的主题模型:每篇文档由一组隐藏主题混合(doc_topics)驱动,每个主题由一组词分布(topic_words)刻画,文档中的每个词都隐含地指派给某一个主题(word_topics)。在标准的变分推断中,word_topics这类离散指派变量通常需要被显式地建模为变分分布并求解,推导繁琐且容易引入较高的方差。
Pyro 的示例 examples/lda.py 给出的思路是:直接把这个离散指派变量在模型内部边缘化掉。具体而言,示例在word_topics采样点处标注infer={"enumerate": "parallel"},并使用TraceEnum_ELBO做推断,这样 ELBO 的计算会自动对word_topics的所有可能取值求和(枚举),guide 中完全不需要出现这个变量。这正对应原文档开篇所述:该示例 "demonstrating how to marginalize out discrete assignment variables in a Pyro model",并将文档视为词 id 向量的批量矩阵(而非词频直方图),用词 id 的 Categorical 分布建模观测。
此外,示例还采用了 [1] 中提出的摊销变分推断框架(autoencoding variational inference for topic models),但去掉了其中的 Laplace 逼近,转而使用 PyTorch 提供的可重参数化 Gamma 与 Dirichlet 分布 [2],使得梯度估计完全基于路径梯度(pathwise gradients / reparameterization trick),无需近似展开。
二、完整代码与运行方式
原文档tutorial/source/lda.rst的主体正是通过literalinclude嵌入的完整示例代码,下面将其完整展开并逐步剖析(行号对应 examples/lda.py 实际源码)。
2.1 依赖与全局设置
import argparse import functools import logging import torch from torch import nn from torch.distributions import constraints import pyro import pyro.distributions as dist from pyro.infer import SVI, JitTraceEnum_ELBO, TraceEnum_ELBO from pyro.optim import ClippedAdam logging.basicConfig(format="%(relativeCreated) 9d %(message)s", level=logging.INFO)代码依赖 PyTorch 与 Pyro,并引入本次推断的两大核心组件:
TraceEnum_ELBO/JitTraceEnum_ELBO:来自 pyro/infer/traceenum_elbo.py,是支持离散采样点穷举枚举的 ELBO 实现;ClippedAdam:来自 pyro/optim/clipped_adam.py,是带梯度裁剪、学习率衰减与中心化方差选项的 Adam 变体。
2.2 运行命令
python examples/lda.py示例默认生成 1000 篇合成文档、每篇 64 个词、词表大小 1024、8 个主题,训练 1000 步 SVI。全部命令行参数见文末"六、命令行参数速查表"。
三、模型:全局主题参数与局部文档变量
模型的完整定义如下(examples/lda.py):
# This is a fully generative model of a batch of documents. # data is a [num_words_per_doc, num_documents] shaped array of word ids # (specifically it is not a histogram). We assume in this simple example # that all documents have the same number of words. def model(data=None, args=None, batch_size=None): # Globals. with pyro.plate("topics", args.num_topics): topic_weights = pyro.sample( "topic_weights", dist.Gamma(1.0 / args.num_topics, 1.0) ) topic_words = pyro.sample( "topic_words", dist.Dirichlet(torch.ones(args.num_words) / args.num_words) ) # Locals. with pyro.plate("documents", args.num_docs) as ind: if data is not None: with pyro.util.ignore_jit_warnings(): assert data.shape == (args.num_words_per_doc, args.num_docs) data = data[:, ind] doc_topics = pyro.sample("doc_topics", dist.Dirichlet(topic_weights)) with pyro.plate("words", args.num_words_per_doc): # The word_topics variable is marginalized out during inference, # achieved by specifying infer={"enumerate": "parallel"} and using # TraceEnum_ELBO for inference. Thus we can ignore this variable in # the guide. word_topics = pyro.sample( "word_topics", dist.Categorical(doc_topics), infer={"enumerate": "parallel"}, ) data = pyro.sample( "doc_words", dist.Categorical(topic_words[word_topics]), obs=data ) return topic_weights, topic_words, data3.1 数据约定:词 id 向量而非词频直方图
注释明确强调:data是形状为[num_words_per_doc, num_documents]的词 id 矩阵("specifically it is not a histogram"),即每个元素是一个词的整数编号,而不是文档的词频统计向量。示例假设所有文档具有相同的词数args.num_words_per_doc。这一点与后面 guide 中将数据转换为直方图的操作形成对照:模型以词 id 为观测,guide 的神经网络则以词频直方图为输入(见第四节)。
3.2 全局变量:主题权重与主题-词分布
with pyro.plate("topics", args.num_topics): topic_weights = pyro.sample( "topic_weights", dist.Gamma(1.0 / args.num_topics, 1.0) ) topic_words = pyro.sample( "topic_words", dist.Dirichlet(torch.ones(args.num_words) / args.num_words) )topic_weights:长度num_topics的向量,先验取Gamma(1/num_topics, 1),作为后续每篇文档doc_topics ~ Dirichlet(topic_weights)的浓度参数;topic_words:形状为[num_topics, num_words]的矩阵,每行是一个主题上的词分布,先验取以均匀分布为均值(torch.ones(...) / args.num_words)的Dirichlet。
两者都被放在pyro.plate("topics", args.num_topics)上下文内,声明主题维上的条件独立性。pyro.plate是 Pyro 中表达"条件独立变量序列"的原语,可顺序使用也可作为上下文管理器并行(向量化)使用,并支持通过subsample_size做小批量切分(见 pyro/primitives.py)。
3.3 局部变量:文档主题与词观测
with pyro.plate("documents", args.num_docs) as ind: if data is not None: with pyro.util.ignore_jit_warnings(): assert data.shape == (args.num_words_per_doc, args.num_docs) data = data[:, ind] doc_topics = pyro.sample("doc_topics", dist.Dirichlet(topic_weights)) with pyro.plate("words", args.num_words_per_doc): word_topics = pyro.sample( "word_topics", dist.Categorical(doc_topics), infer={"enumerate": "parallel"}, ) data = pyro.sample( "doc_words", dist.Categorical(topic_words[word_topics]), obs=data )- 外层
pyro.plate("documents", args.num_docs) as ind遍历文档;当传入真实data时,通过data[:, ind]按当前索引(含小批量)切取对应文档列; doc_topics ~ Dirichlet(topic_weights)是每篇文档的主题混合;- 内层
pyro.plate("words", args.num_words_per_doc)遍历每篇文档中的每个词位; word_topics:每个词位的主题指派,服从Categorical(doc_topics),其infer={"enumerate": "parallel"}声明该离散变量在推断时做并行枚举;doc_words ~ Categorical(topic_words[word_topics])以obs=data的形式接收观测,是唯一的观测采样点。
3.4 枚举边缘化的关键一行
word_topics采样点上的infer={"enumerate": "parallel"}是本示例的核心手法。它告诉推断算法:对该离散变量做穷举枚举而非采样。结合TraceEnum_ELBO,模型对word_topics的所有可能取值求和(边缘化),doc_words的对数似然项因此成为关于doc_topics的充分统计量。这样 guide 只需推断连续的doc_topics,离散指派变量完全不必出现在 guide 中——这正是代码注释 "Thus we can ignore this variable in the guide" 的含义。
从 pyro/infer/traceenum_elbo.py 的类注释可以看到该 ELBO 的能力边界:支持对离散采样点穷举枚举、支持 guide 内的局部并行采样;对模型采样点标注infer={'enumerate': 'parallel'}且该采样点不出现在 guide 中即可完成模型侧枚举。其内部通过EnumMessenger(pyro/infer/traceenum_elbo.py)为模型与 guide 分配独立的枚举维度,并借助contract_tensor_tree与SampleRing(pyro/infer/traceenum_elbo.py)以 tensor message passing 方式对因子做收缩,最终由Dice(guide_trace, ordering).compute_expectation(costs)(pyro/infer/traceenum_elbo.py)计算期望,实现对枚举变量的精确边缘化。
若想对 guide 中的离散点批量配置枚举,可使用pyro.infer.enum.config_enumerate(默认策略"parallel",另有"sequential"、"flat"、None可选,见 pyro/infer/enum.py);本示例只需要模型侧单点枚举,因此直接在采样点标注即可。
四、Guide:共轭 guide + 摊销 MLP guide
变分推断需要定义一个 guide 来逼近后验。示例将变量分成两类分别处理:
- 全局变量(
topic_weights、topic_words):使用共轭的指数族 guide; - 局部变量(
doc_topics):使用由神经网络参数化的摊销 guide(amortized inference)。
4.1 共轭 guide:全局变量的指数族后验
def parametrized_guide(predictor, data, args, batch_size=None): # Use a conjugate guide for global variables. topic_weights_posterior = pyro.param( "topic_weights_posterior", lambda: torch.ones(args.num_topics), constraint=constraints.positive, ) topic_words_posterior = pyro.param( "topic_words_posterior", lambda: torch.ones(args.num_topics, args.num_words), constraint=constraints.greater_than(0.5), ) with pyro.plate("topics", args.num_topics): pyro.sample("topic_weights", dist.Gamma(topic_weights_posterior, 1.0)) pyro.sample("topic_words", dist.Dirichlet(topic_words_posterior))topic_weights_posterior:形状[num_topics]的可学习参数,经constraints.positive约束为正数,作为 Gamma 后验的浓度参数(rate 固定为 1);topic_words_posterior:形状[num_topics, num_words]的可学习参数,经constraints.greater_than(0.5)约束(Dirichlet 浓度参数需大于 0,此处进一步限定大于 0.5),作为 Dirichlet 后验的浓度参数;- 两处
pyro.sample的采样点名称与模型中的全局变量一一对应,Pyro 通过名称匹配自动计算每个采样点的 score 项。
pyro.param声明可训练参数并存入全局参数仓库(param store);pyro.module("predictor", predictor)则把神经网络的参数注册进同一个仓库(见 examples/lda.py)。
4.2 摊销 guide:MLP 输出 doc_topics 的 Delta 分布
def make_predictor(args): layer_sizes = ( [args.num_words] + [int(s) for s in args.layer_sizes.split("-")] + [args.num_topics] ) logging.info("Creating MLP with sizes {}".format(layer_sizes)) layers = [] for in_size, out_size in zip(layer_sizes, layer_sizes[1:]): layer = nn.Linear(in_size, out_size) layer.weight.data.normal_(0, 0.001) layer.bias.data.normal_(0, 0.001) layers.append(layer) layers.append(nn.Sigmoid()) layers.append(nn.Softmax(dim=-1)) return nn.Sequential(*layers)make_predictor构造一个多层感知机:输入维度为词表大小num_words,隐藏层由--layer-sizes(默认"100-100")指定,输出维度为num_topics;隐层激活函数为 Sigmoid,最后一层 Softmax 保证输出是合法的概率单纯形(主题混合)。权重与偏置均以均值 0、标准差 0.001 的正态分布初始化,保证起始时接近均匀输出。
# Use an amortized guide for local variables. pyro.module("predictor", predictor) with pyro.plate("documents", args.num_docs, batch_size) as ind: data = data[:, ind] # The neural network will operate on histograms rather than word # index vectors, so we'll convert the raw data to a histogram. counts = torch.zeros(args.num_words, ind.size(0)).scatter_add( 0, data, torch.ones(data.shape) ) doc_topics = predictor(counts.transpose(0, 1)) pyro.sample("doc_topics", dist.Delta(doc_topics, event_dim=1))这里有三处值得深入说明:
- 输入转换:模型以词 id 向量为观测,而神经网络天然更适合处理定长向量输入。guide 将一个小批量内每篇文档的词 id 通过
scatter_add转成形状[num_words, batch_size]的词频直方图,再转置为[batch_size, num_words]喂给 MLP; - plate 的小批量切分:
pyro.plate("documents", args.num_docs, batch_size)将第三个位置参数作为subsample_size,在 1000 篇文档中每次随机抽取batch_size(默认 32)篇,并对该上下文内的对数似然按size/batch_size比例缩放,实现无偏的随机小批量训练(见 pyro/primitives.py); - Delta 作为变分分布:
dist.Delta(doc_topics, event_dim=1)将 MLP 的确定性输出封装为退化分布。event_dim=1表明每个"样本"是一个向量(主题混合),Delta 类本身是可重参数化的(has_rsample = True,见 pyro/distributions/delta.py),因此doc_topics的梯度可以通过路径梯度直接回传到 MLP。这正是"摊销"的含义:局部变量doc_topics的后验由神经网络参数化,对任意新文档都能即时推理,无需为每篇文档单独优化变分参数。
五、训练:SVI、ClippedAdam 与 JIT 加速
5.1 主流程
def main(args): logging.info("Generating data") pyro.set_rng_seed(0) pyro.clear_param_store() # We can generate synthetic data directly by calling the model. true_topic_weights, true_topic_words, data = model(args=args) # We'll train using SVI. logging.info("-" * 40) logging.info("Training on {} documents".format(args.num_docs)) predictor = make_predictor(args) guide = functools.partial(parametrized_guide, predictor) Elbo = JitTraceEnum_ELBO if args.jit else TraceEnum_ELBO elbo = Elbo(max_plate_nesting=2) optim = ClippedAdam( {"lr": args.learning_rate, "centered_variance": args.centered_variance} ) svi = SVI(model, guide, optim, elbo) logging.info("Step\tLoss") for step in range(args.num_steps): loss = svi.step(data, args=args, batch_size=args.batch_size) if step % 10 == 0: logging.info("{: >5d}\t{}".format(step, loss)) loss = elbo.loss(model, guide, data, args=args) logging.info("final loss = {}".format(loss))流程要点:
- 合成数据生成:直接调用
model(args=args)生成一组"真实"主题参数与文档数据,作为训练与评估基准——模型本身就是数据生成器,这是概率编程的一大便利; max_plate_nesting=2:声明模型中最深的 plate 嵌套深度为 2(documents内嵌words),这是TraceEnum_ELBO进行维度分配与形状校验的必要信息;SVI(model, guide, optim, elbo):标准随机变分推断接口,svi.step(...)每次调用完成一次前向 ELBO 计算与参数更新;functools.partial:把predictor绑定进 guide 的可调用签名,使 guide 与模型共享(data, args, batch_size)的调用约定;- 每 10 步打印一次损失,结束后用完整数据集再算一次
elbo.loss(...)给出最终损失。
5.2 ClippedAdam:为什么 LDA 需要梯度裁剪
示例按 [1] 的建议使用 Adam 优化器并对梯度做裁剪。Pyro 提供现成的ClippedAdam(pyro/optim/clipped_adam.py),它在torch.optim.Adam基础上增加了三个特性:
- 梯度裁剪:每步更新前执行
grad.clamp_(-clip_norm, clip_norm),默认clip_norm=10.0,可抑制异常大梯度对参数更新的破坏; - 学习率衰减:每步执行
group["lr"] *= group["lrd"],lrd默认 1.0(不衰减); - 中心化方差(centered variance):当
centered_variance=True时,二阶矩估计改用grad - exp_avg(即去中心化的梯度),对应文献 [2] 的式 (2) 的路径梯度方差修正方案——这正是示例引入 [2] 的目的:在重参数化梯度框架下,中心化方差可以降低梯度方差、稳定训练。
命令行参数-cv/--centered-variance默认False,读者可开启后对比训练曲线。
5.3 JIT 编译:JitTraceEnum_ELBO 的适用前提
传--jit时改用JitTraceEnum_ELBO。从 pyro/infer/traceenum_elbo.py 的注释可知,该实现通过pyro.ops.jit.compile编译loss_and_grads,能显著降低 Python 开销,但只适用于有限的一类模型:
- 模型必须具有静态结构(每次执行的计算图不随数据变化);
- 模型不能依赖任何全局数据(参数仓库除外);
- 所有张量输入必须通过位置参数传入,非张量输入通过关键字参数传入(编译按
**kwargs缓存)。
本示例的model与parametrized_guide满足这些约束(数据、args、batch_size均为入参),因此可以安全开启--jit。训练结束后可调用pyro.infer.util中的校验例程验证形状与枚举配置(TraceEnum_ELBO在开启验证时会自动检查模型/guide 枚举约束,见 pyro/infer/traceenum_elbo.py)。
六、命令行参数速查表
参数解析位于 examples/lda.py,全部选项如下:
| 短选项 | 长选项 | 默认值 | 类型 | 含义 |
|---|---|---|---|---|
-t | --num-topics | 8 | int | 主题数量 K,同时决定topic_weights先验Gamma(1/K, 1) |
-w | --num-words | 1024 | int | 词表大小,决定topic_words与 MLP 输入维度 |
-d | --num-docs | 1000 | int | 文档总数(合成数据规模) |
-wd | --num-words-per-doc | 64 | int | 每篇文档的词数,模型假定所有文档等长 |
-n | --num-steps | 1000 | int | SVI 训练步数 |
-l | --layer-sizes | "100-100" | str | MLP 隐藏层尺寸,用-分隔,可写多个 |
-lr | --learning-rate | 0.01 | float | ClippedAdam 学习率 |
-cv | --centered-variance | False | bool | 是否启用中心化方差(文献 [2] 的梯度方差修正) |
-b | --batch-size | 32 | int | 每个小批量的文档数(plate 的subsample_size) |
--jit | 关闭 | flag | 启用JitTraceEnum_ELBO编译加速 |
例如,用 16 个主题、更大词表、双隐藏层 MLP 训练 2000 步并开启 JIT:
python examples/lda.py -t 16 -w 4096 -n 2000 -l "200-200" --jit七、算法与实现要点回顾
| 设计选择 | 实现方式 | 效果 |
|---|---|---|
| 边缘化离散指派变量 | word_topics标注infer={"enumerate": "parallel"}+TraceEnum_ELBO | guide 无需建模离散变量,ELBO 精确求和 |
| 可重参数化先验 | PyTorch 的Gamma/Dirichlet | 获得路径梯度,免除 [1] 中的 Laplace 逼近 |
| 全局变量共轭 guide | pyro.param+ Gamma/Dirichlet 后验 | 以解析形式逼近全局后验 |
| 局部变量摊销 guide | MLP 输出Delta(doc_topics, event_dim=1) | 对新文档即时推理,参数全局共享 |
| 梯度裁剪与方差修正 | ClippedAdam(clip_norm=10, centered_variance=...) | 稳定训练、降低梯度方差 |
| 小批量训练 | pyro.plate(..., batch_size)自动缩放对数似然 | 支持大规模文档集合的无偏随机优化 |
文中所有结论均可在仓库对应文件中验证:模型与 guide 的完整实现见 examples/lda.py,枚举式 ELBO 的机制见 pyro/infer/traceenum_elbo.py,枚举配置 API 见 pyro/infer/enum.py,优化器实现见 pyro/optim/clipped_adam.py。官方教程入口为 tutorial/source/lda.rst,读者可结合 tutorial/source/enumeration.ipynb 进一步学习 Pyro 离散枚举的一般用法。
八、参考资料
示例的实现依据以下两篇文献(对应 examples/lda.py 文件头部的 References 说明):
- Akash Srivastava, Charles Sutton. ICLR 2017. "Autoencoding Variational Inference for Topic Models"——提出了用摊销变分推断训练 LDA 的整体框架,本文示例沿用了其 Adam + 梯度裁剪的训练策略;
- Martin Jankowiak, Fritz Obermeyer. ICML 2018. "Pathwise gradients beyond the reparametrization trick"——给出了重参数化梯度之外的路径梯度理论,并支持本文示例使用 PyTorch 可重参数化 Gamma/Dirichlet 分布、以
centered_variance选项降低梯度方差的做法。
两个文献的完整引用信息已包含在 examples/lda.py 的文档字符串中,读者可按需查阅原文。
- 人工智能
- 机器学习
- 深度学习
- 概率编程
【免费下载链接】pyro
Deep universal probabilistic programming with Python and PyTorch
相关推荐
解决常见问题:TagLib使用中的5个实用技巧
解决常见问题:TagLib 使用中的5个实用技巧 TagLib 是一款强大的媒体文件元数据读写库,能够帮助开发者轻松处理各种音频、视频和图像文件的元数据信息。本
numpy-ml 潜狄利克雷分配(LDA)实现深度解析:基于变分 EM 训练的非平滑主题模型
numpy ml 潜狄利克雷分配(LDA)实现深度解析:基于变分 EM 训练的非平滑主题模型 导读 :本文以 docs/numpy_ml.lda.lda.rst
机器学习人工智能在 DGL 上用消息传递实现潜在狄利克雷分配(LDA):以二部多重图驱动变分推断的完整实战指南
在 DGL 上用消息传递实现潜在狄利克雷分配(LDA):以二部多重图驱动变分推断的完整实战指南 本指南深入讲解 DGL(Deep Graph Library)官
人工智能机器学习深度学习图计算
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考