pyro-ppl/numpyro
Probabilistic programming with NumPy powered by JAX for autograd and JIT compilation to GPU/TPU/CPU.
解决的问题
NumPyro 是一个旨在对模型进行贝叶斯推论的概率编程库。它通过利用 JAX 进行自动微分和即时 (JIT) 编译,解决了概率编程中推论速度慢的问题,使模型能够在 CPU、GPU 和 TPU 上高效运行。
工作原理
NumPyro 为 Pyro 概率编程语言提供类 NumPy 的后端。它使用 JAX 将推论算法的整个集成步骤编译为 XLA 优化的内核,从而减少 Python 开销。该库包含一套完整的分布类、约束和双射变换,并支持用于实现自定义推论算法的效果处理器。
适用对象
面向需要进行复杂概率建模和贝叶斯推论的研究人员和开发人员,特别是那些已经熟悉 Pyro 或 PyTorch 分布 API 的人员。
亮点
- JAX 驱动的性能:利用 JIT 编译和 autograd 在硬件加速器上加速 MCMC 和变分推论。
- 核心 MCMC 算法:实现了 No-U-Turn Sampler (NUTS)、Hamiltonian Monte Carlo (HMC)、用于离散变量的 MixedHMC 以及用于大型数据集的 HMCECS。
- 变分推论:支持自动微分变分推论 (ADVI),并为包括具有离散隐变量的模型在内的模型提供灵活的指南。
- 灵活的分布:提供广泛的分布类,并支持来自 TensorFlow Probability (TFP) 的分布。
- Python 中的 C++ 级速度:通过 Iterative NUTS 消除了 NUTS 树构建阶段的 Python 开销。
相关
- 项目
- 项目
- 项目
- 项目
- 项目