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 開銷。
相關
- 專案
- 專案
- 專案
- 專案
- 專案