[论文解读] A Universal Marginalizer for Amortized Inference in Generative Models
本文提出一种通用边缘化器(Universal Marginalizer, UM),即一个单一神经网络,可近似贝叶斯网络中任意证据集下的所有条件边缘分布。通过结合祖先采样与可学习掩码机制,并在多样化观测模式上进行训练,UM 实现了摊销推理,并通过混合提议分布显著提升了重要性采样效率,使所需样本数最多减少 8 倍,同时保持高精度。
We consider the problem of inference in a causal generative model where the set of available observations differs between data instances. We show how combining samples drawn from the graphical model with an appropriate masking function makes it possible to train a single neural network to approximate all the corresponding conditional marginal distributions and thus amortize the cost of inference. We further demonstrate that the efficiency of importance sampling may be improved by basing proposals on the output of the neural network. We also outline how the same network can be used to generate samples from an approximate joint posterior via a chain decomposition of the graph.
研究动机与目标
- 解决复杂因果生成模型在不同证据集下推理的高计算成本问题。
- 开发一个单一神经网络,可近似贝叶斯网络中任意观测下所有节点的条件边缘分布。
- 通过将 UM 输出用作提议分布,提升重要性采样的效率。
- 利用图模型的链式分解,通过 UM 实现后验分布的联合采样。
- 证明单一 UM 可替代多个专用模型,从而降低推理开销。
提出的方法
- 训练一个前馈神经网络,以预测给定任意观测节点子集 XO 时,贝叶斯网络中所有节点 Xi 的后验边缘分布 P(Xi|XO= xO)。
- 训练过程中使用概率掩码,其中每个节点以均匀概率独立地被掩码,以模拟任意观测模式。
- 使用 2 位或 33 位表示法对观测和未观测节点进行编码,其中未观测节点包含其先验概率。
- 在多标签分类设置下,使用二元交叉熵损失函数训练网络,以预测节点概率。
- 通过混合参数 β,将祖先采样与 UM 预测的边缘分布相结合,构建混合重要性采样提议。
- 利用 UM 输出指导拓扑有序的链式分解中的顺序采样,从而提升提议质量。
实验结果
研究问题
- RQ1一个单一神经网络能否对任意观测变量子集,近似贝叶斯网络中所有条件边缘分布?
- RQ2将 UM 输出用作重要性采样中的提议分布,是否相比标准方法能提升采样效率?
- RQ3UM 的性能如何随不同网络架构和输入编码方式而变化?
- RQ4UM 是否可用于通过图模型的链式分解生成联合后验样本?
- RQ5在重要性采样中,祖先采样与基于 UM 的提议之间,最优混合策略是什么?
主要发现
- 参数量最大的单层神经网络(2048 个单元)表现最佳,使用 33 位表示法并包含先验信息时,平均绝对误差为 0.0060,最大绝对误差为 0.3223。
- 混合重要性采样中 β = 0.25 时,在 250,000 个样本下,真实与估计边缘分布的相关性达到 95%,优于标准 IS(β = 0)方法,后者需 200 万个样本才能达到 92% 的相关性。
- 混合提议方法在 200 万个样本下实现了 96% 的相关性,证明其在提升准确率的同时显著降低了计算成本。
- 包含未观测节点先验概率的 33 位表示法,在平均误差和最大误差指标上均略优于 2 位表示法。
- 基于 UM 的混合提议方法使所需样本数最多减少 8 倍,同时保持或提升相关性与有效样本量(ESS)。
- 该方法可使用单一模型实现所有可能证据集的摊销推理,无需为每种观测模式单独训练模型。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。