☰
Pyro 中的摊销隐狄利克雷分配(Amortized LDA):利用离散枚举边缘化词-主题指派变量
2026/9/25 5:07:54 网站建设 项目流程
  • 人工智能
  • 机器学习
  • 深度学习
  • 概率编程

【免费下载链接】pyro

Deep universal probabilistic programming with Python and PyTorch

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

导读

本文基于 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, data

3.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 来逼近后验。示例将变量分成两类分别处理:

  1. 全局变量(topic_weights、topic_words):使用共轭的指数族 guide;
  2. 局部变量(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))

这里有三处值得深入说明:

  1. 输入转换:模型以词 id 向量为观测,而神经网络天然更适合处理定长向量输入。guide 将一个小批量内每篇文档的词 id 通过scatter_add转成形状[num_words, batch_size]的词频直方图,再转置为[batch_size, num_words]喂给 MLP;
  2. plate 的小批量切分:pyro.plate("documents", args.num_docs, batch_size)将第三个位置参数作为subsample_size,在 1000 篇文档中每次随机抽取batch_size(默认 32)篇,并对该上下文内的对数似然按size/batch_size比例缩放,实现无偏的随机小批量训练(见 pyro/primitives.py);
  3. 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))

流程要点:

  1. 合成数据生成:直接调用model(args=args)生成一组"真实"主题参数与文档数据,作为训练与评估基准——模型本身就是数据生成器,这是概率编程的一大便利;
  2. max_plate_nesting=2:声明模型中最深的 plate 嵌套深度为 2(documents内嵌words),这是TraceEnum_ELBO进行维度分配与形状校验的必要信息;
  3. SVI(model, guide, optim, elbo):标准随机变分推断接口,svi.step(...)每次调用完成一次前向 ELBO 计算与参数更新;
  4. functools.partial:把predictor绑定进 guide 的可调用签名,使 guide 与模型共享(data, args, batch_size)的调用约定;
  5. 每 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-topics8int主题数量 K,同时决定topic_weights先验Gamma(1/K, 1)
-w--num-words1024int词表大小,决定topic_words与 MLP 输入维度
-d--num-docs1000int文档总数(合成数据规模)
-wd--num-words-per-doc64int每篇文档的词数,模型假定所有文档等长
-n--num-steps1000intSVI 训练步数
-l--layer-sizes"100-100"strMLP 隐藏层尺寸,用-分隔,可写多个
-lr--learning-rate0.01floatClippedAdam 学习率
-cv--centered-varianceFalsebool是否启用中心化方差(文献 [2] 的梯度方差修正)
-b--batch-size32int每个小批量的文档数(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_ELBOguide 无需建模离散变量,ELBO 精确求和
可重参数化先验PyTorch 的Gamma/Dirichlet获得路径梯度,免除 [1] 中的 Laplace 逼近
全局变量共轭 guidepyro.param+ Gamma/Dirichlet 后验以解析形式逼近全局后验
局部变量摊销 guideMLP 输出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 说明):

  1. Akash Srivastava, Charles Sutton. ICLR 2017. "Autoencoding Variational Inference for Topic Models"——提出了用摊销变分推断训练 LDA 的整体框架,本文示例沿用了其 Adam + 梯度裁剪的训练策略;
  2. 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

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

相关推荐

上一篇:KryoNet错误处理与调试:5个常见问题排查与解决方案指南
下一篇:终极指南:Swoole协程Channel如何彻底改变PHP并发编程 🚀

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

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

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

立即咨询