REVIEW 4 major objections 5 minor 52 references
JaxSGMC: Modular stochastic gradient MCMC in JAX
T0 review · 4 major / 5 minor · reviewed 2026-08-15 · deepseek-v4-flash
Pith's one-line read JaxSGMC packages stochastic-gradient MCMC samplers as modular JAX building blocks.
desk verdict A solid, genuinely useful JAX library for SG-MCMC with a weak empirical verification section; worth reviewing, but the authors should add quantitative sampler checks. read the letter →
The pith
A machine-rendered reading of the paper's core claim, the machinery that carries it, and where it could break.
The reading
What carries the argument
The central object is the modular sampler architecture: each sampler is built from independent modules---potential.py (log-likelihood/prior potentials), data.py (jit-compatible DataLoaders), adaption.py (RMSProp and covariance preconditioners), integrator.py (Langevin diffusion, leapfrog with friction, OBABO), solver.py (accept/reject or unconditional sample processing), and scheduler.py (step size, burn-in, thinning). The mathematical object that carries the dynamics is the stochastic potential $U(\theta) \approx -\frac{N}{n}\sum_{i} \log p(y_i|x_i,\theta,\mathcal{M}) - \log p(\theta|\mathcal{M})$, whose mini-batch gradient drives each integrator; composing modules around this potential lets the same building blocks express different samplers.
What would settle it
Run one of the library's samplers on a tractable target with a closed-form posterior, such as a Gaussian linear model with known variance, and compare the MCMC sample variance and the posterior predictive coverage to the nominal level. A mismatch beyond Monte Carlo error, especially in the noise injected by the Langevin integrator or in replica-exchange swap acceptance, would show the implementation does not sample the stated distribution.
Extended reading notes
Core claim
On its own terms, JaxSGMC claims that the full variety of modern SG-MCMC algorithms can be organized into a small set of reusable modules---potential evaluation, data batching, preconditioner adaptation, integrators, solvers, and schedulers---and that composing these modules yields both standard samplers like (preconditioned) SGLD, SGHMC, and SGGMC and more recent schemes like replica-exchange SGLD and AMAGOLD under one API. The library is designed so that sampler components can be jit-compiled end-to-end, including the data loading, which keeps Bayesian sampling runtime-competitive with stochastic optimization. This is offered as a practical path for making SG-MCMC a drop-in alternative to optimization in JAX deep-learning workflows.
Load-bearing premise
The whole contribution rests on the implemented samplers being faithful to the algorithms they cite, so that the Markov chains they produce actually converge to the intended posterior; the paper's verification is qualitative or indirect, and a subtle implementation bug would silently bias the uncertainty estimates.
Editorial extensions
If this is right
- A JAX user with a model already written as a function can call a high-level alias to sample the posterior instead of optimizing, with no model rewrite.
- Samplers that require advanced building blocks, such as parallel tempering and amortized Metropolis-Hastings acceptance, become available through a common API rather than as stand-alone code.
- Custom samplers can be assembled from the documented blocks, letting practitioners tailor proposals, preconditioners, and schedules to a problem.
- Because data loading lives inside the jit-compiled loop, SG-MCMC inference can run at a cost comparable to stochastic optimization; the paper reports pSGLD training of 200 epochs on CIFAR-10 at the same order of magnitude as one full-batch HMC proposal.
Reading between the lines
- If the modular decomposition is as clean as presented, SG-MCMC research could shift toward composing and benchmarking blocks rather than reimplementing full samplers, making new samplers easier to compare on identical data-loading and scheduling code.
- A natural stress test the paper leaves implicit is quantitative convergence checking: running the library's samplers on targets with known posterior moments would let users verify the noise scaling and acceptance steps that the qualitative agreement plot does not settle.
- The CIFAR-10 example suggests a direct extension: the same API could be used to ablate which building block (preconditioner, integrator, or swap schedule) most improves posterior coverage, a comparison the paper does not carry out.
- The authors indicate that pSGLD may not fully explore posterior volume; a testable next step is whether AMAGOLD or replica-exchange samplers in this library reduce that gap on the same neural-network potential benchmark.
Signed reviews
Editorial analysis
A structured set of objections, weighed in public.
Referee Report
Summary. The paper presents JaxSGMC, a JAX library for stochastic gradient Markov chain Monte Carlo (SG-MCMC) that provides pre-built samplers (pSGLD, SGHMC, reSGLD, AMAGOLD, SGGMC) and a modular API for composing custom samplers from building blocks such as integrators, schedulers, and data loaders. The software is designed to be domain-independent, with end-to-end jit compilation and host-device data transfer. The paper demonstrates the library with two examples: a linear regression model using a custom-built pSGLD sampler, compared qualitatively to NumPyro HMC, and a CIFAR-10 image classification task using a pre-built pSGLD sampler, reporting classification accuracy and certainty-threshold behavior. The central claims are that JaxSGMC faithfully implements state-of-the-art SG-MCMC samplers and that its modular structure lowers the barrier to using and developing SG-MCMC methods.
Significance. If the implementation is faithful, JaxSGMC fills a practical gap by offering a modular, JAX-native SG-MCMC library that supports recently proposed samplers (reSGLD, AMAGOLD) not available in other JAX-based libraries, and it could accelerate adoption of Bayesian UQ in deep learning and physical modeling. The paper ships a public repository, documentation, and code listings, which are strengths for a software paper. However, the empirical verification of the central sampling-fidelity claim is currently weak, especially for the novel samplers, and one of the code listings appears to contain a substantive error. These issues are load-bearing for a library whose purpose is to draw samples from the correct posterior.
major comments (4)
- [Section 3.1] The validation of pSGLD against NumPyro HMC is qualitative: the text states that the distributions "agree reasonably well" without any numerical diagnostics. To substantiate the claim that the sampler targets the correct posterior, report quantitative measures such as Wasserstein or maximum mean discrepancy, effective sample size, or repeated-seed scatter plots.
- [Section 3.2] The CIFAR-10 example reports classification accuracy and certainty thresholds but provides no convergence diagnostics (e.g., trace plots, ESS, Gelman-Rubin) or posterior-fidelity checks. Ensemble classification accuracy alone does not demonstrate that the sampler draws from the intended posterior; add calibration or predictive-coverage checks, or explicitly label the example as an illustrative runtime/feature demonstration.
- [Section 2.3 / Section 3] The samplers highlighted as novel contributions relative to prior JAX libraries—replica exchange SG-MCMC (reSGLD) and AMAGOLD—are never empirically validated. A subtle implementation error, such as an incorrect swap-acceptance probability in reSGLD or an incorrect noise scaling in the AMAGOLD proposal, would bias the stationary distribution while remaining invisible in the current examples. Add synthetic-data experiments with known target posteriors for these samplers.
- [Listing 2, Section 3.1] The `log_prior` function returns `1 / jnp.exp(sample["log_sigma"])`, which is not the log-density of an exponential prior. For σ = exp(log_sigma) and an exponential(1) prior, the log-density is `-jnp.exp(sample["log_sigma"])` up to an additive constant. As written, the example defines a different potential than described in the text, compromising the illustrative linear regression example.
minor comments (5)
- [Section 4] The phrase "an domain-independent library" should be "a domain-independent library".
- [Figure 2 caption] The caption says "contour plots of Gaussians obtained from the Hamiltonian Monte Carlo (HMC) method"; clarify whether the Gaussians are fitted to HMC samples or derived analytically from the linear regression posterior.
- [Section 3.2] The statement that "the cost of the whole pSGLD training of 200 epochs is the same order of magnitude as the cost of generating a single sample with the full-batch HMC" is vague; specify hardware, the HMC hyperparameters (number of leapfrog steps, trajectory length), and the exact runtime comparison.
- [Section 4] The text refers to "Stochastic Variation Inference"; it should be "Stochastic Variational Inference".
- [Table 2] Table 2 lists the module as `solvers.py` while the text and listings refer to `solver.py`; make the module name consistent.
Circularity Check
No material circularity: the samplers are implementations of cited external algorithms checked against an independent NumPyro benchmark, and the few self-citations are not load-bearing.
full rationale
The paper's central claim is that the JaxSGMC library implements published SG-MCMC samplers and exposes modular building blocks. The derivation chain here is one of software engineering rather than equation-fitting: no parameter is fitted to data and then renamed as a prediction. The linear-regression example compares pSGLD samples against a gold-standard Hamiltonian Monte Carlo implementation in NumPyro, i.e., an independent external reference, and the CIFAR-10 example checks accuracy and certainty-based test-set performance rather than asserting a posterior identity. The cited algorithms (SGLD, pSGLD, SGHMC, reSGLD, AMAGOLD, SGGMC) are external publications; the paper does not derive those samplers from its own outputs. The only self-referential elements appear in the Impact section, where the authors cite their own prior applications of JaxSGMC as anecdotal evidence of utility. Those citations are not load-bearing for the central claim that the library implements the listed samplers; removing them would not weaken the software description or the examples. The skeptical concern that verification is qualitative (e.g., distributions 'agree reasonably well' without quantitative diagnostics, and no distributional check for reSGLD or AMAGOLD) is a correctness or validation-strength issue, not circularity: a weak or missing test does not make the implementation's claim equivalent to its inputs by construction. Accordingly, the circularity burden is low; the appropriate score is 1 for the minor, non-load-bearing self-citations, with no circular steps identified.
Assumptions & free parameters
free parameters (6)
- step_size_first =
0.05 (linear regression), 0.001 (CIFAR)
- step_size_last =
0.001 (linear regression)
- gamma =
0.33
- burn_in =
2000 (linear regression), 35100 (CIFAR)
- batch_size =
256 (CIFAR)
- prior_std =
10 (CIFAR Gaussian prior)
assumptions (4)
- domain assumption Discretized stochastic gradient Langevin dynamics converge to the posterior as the step size tends to zero with appropriately added noise.
- domain assumption Observations in the dataset are conditionally independent given the model parameters, so the log-likelihood factorizes over data points.
- domain assumption JAX's host-callback API correctly and efficiently transfers data from DataLoaders into jit-compiled computations.
- domain assumption The implementations of reSGLD, AMAGOLD, and SGGMC follow the original papers' algorithms exactly, including swap acceptance and amortized Metropolis-Hastings steps.
Cite this review
Pith. "Pith review of JaxSGMC: Modular stochastic gradient MCMC in JAX." pith.science (2026). https://pith.science/paper/GUNIWVJ3
@misc{pith2026250511190,
author = {Pith},
title = {Pith review of: JaxSGMC: Modular stochastic gradient MCMC in JAX},
year = {2026},
howpublished = {\url{https://pith.science/paper/GUNIWVJ3}},
note = {Machine review of arXiv:2505.11190}
}
read the original abstract
We present JaxSGMC, an application-agnostic library for stochastic gradient Markov chain Monte Carlo (SG-MCMC) in JAX. SG-MCMC schemes are uncertainty quantification (UQ) methods that scale to large datasets and high-dimensional models, enabling trustworthy neural network predictions via Bayesian deep learning. JaxSGMC implements several state-of-the-art SG-MCMC samplers to promote UQ in deep learning by reducing the barriers of entry for switching from stochastic optimization to SG-MCMC sampling. Additionally, JaxSGMC allows users to build custom samplers from standard SG-MCMC building blocks. Due to this modular structure, we anticipate that JaxSGMC will accelerate research into novel SG-MCMC schemes and facilitate their application across a broad range of domains.
Figures
Reference graph
Works this paper leans on
-
[1]
J. Devlin, M.-W. Chang, K. Lee, K. Toutanova, BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding, arXiv preprint arXiv:1810.04805 (2018)
arXiv 2018
-
[2]
S. Grigorescu, B. Trasnea, T. Cocias, G. Macesanu, A survey of deep learning techniques for autonomous driving, J. Field Robot. 37 (3) (2020) 362–386
work page 2020
- [3]
-
[4]
Raissi, P
M. Raissi, P. Perdikaris, G. E. Karniadakis, Physics-informed neural networks: A deep learning framework for solving forward and inverse problems involving nonlinear partial differential equations, J. Comput. Phys. 378 (2019) 686–707
2019
- [5]
- [6]
- [7]
-
[8]
S. Arakelyan, R. J. Das, Y. Mao, X. Ren, Exploring distributional shifts in large language models for code analysis, arXiv preprint arXiv:2303.09128 (2023)
arXiv 2023
Show all 52 references
-
[9]
Efron, R
B. Efron, R. J. Tibshirani, An introduction to the bootstrap, CRC press, 1994
1994
-
[10]
J. Lei, M. G’Sell, A. Rinaldo, R. J. Tibshirani, L. Wasserman, Distribution-free predictive inference for regression, J. Am. Stat. Assoc. 113 (523) (2018) 1094–1111
2018
-
[11]
Lakshminarayanan, A
B. Lakshminarayanan, A. Pritzel, C. Blundell, Simple and scalable pre- dictive uncertainty estimation using deep ensembles, in: Advances in Neural Information Processing Systems, Vol. 30, Long Beach, CA, USA, Dec. 4–9, 2017, pp. 6405–6416
2017
-
[12]
R. M. Neal, Handbook of Markov Chain Monte Carlo, 1st Edition, Chap- man and Hall/CRC, New York, USA, 2011, Ch. MCMC using Hamilto- nian Dynamics, pp. 139–188
2011
-
[13]
Welling, Y
M. Welling, Y. W. Teh, Bayesian learning via stochastic gradient Langevin dynamics, in: Proceedings of the 28th International Confer- ence on Machine Learning, Bellevue, WA, USA, Jun. 28 – Jul. 2, 2011, pp. 681–688
2011
-
[14]
Graves, Practical variational inference for neural networks, in: Ad- vances in neural information processing systems, Vol
A. Graves, Practical variational inference for neural networks, in: Ad- vances in neural information processing systems, Vol. 24, 2011
2011
-
[15]
M. D. Hoffman, A. Gelman, The No-U-Turn Sampler: Adaptively Set- ting Path Lengths in Hamiltonian Monte Carlo, J. Mach. Learn. Res. 15 (2014) 1593–1623. 17
2014
-
[16]
T. Chen, E. Fox, C. Guestrin, Stochastic gradient Hamiltonian Monte Carlo, in: Proceedings of the 31st International Conference on Machine Learning, Beijing, China, Jun. 21–26, 2014, pp. 1683–1691
2014
-
[17]
C. Li, C. Chen, D. E. Carlson, L. Carin, Preconditioned Stochastic Gradient Langevin Dynamics for Deep Neural Networks, in: Proceedings of the Thirtieth AAAI Conference on Artificial Intelligence, Phoenix, AZ, USA, February 12–17, 2016, pp. 1788–1794
2016
-
[18]
Nemeth, P
C. Nemeth, P. Fearnhead, Stochastic gradient Markov chain Monte Carlo, J. Am. Stat. Assoc. 116 (533) (2021) 433–450
2021
-
[19]
G. Lamb, B. Paige, Bayesian Graph Neural Networks for Molecular Property Prediction, in: Machine Learning for Molecules Workshop at NeurIPS, MIT Press, Online, Dec. 12, 2020
2020
-
[20]
Z. Zou, X. Meng, A. F. Psaros, G. E. Karniadakis, NeuralUQ: A com- prehensive library for uncertainty quantification in neural differential equations and operators, arXiv preprint arXiv:2208.11866 (2022)
2022 arXiv
-
[21]
J. V. Dillon, I. Langmore, D. Tran, E. Brevdo, S. Vasudevan, D. Moore, B. Patton, A. Alemi, M. Hoffman, R. A. Saurous, Tensorflow distribu- tions, arXiv preprint arXiv:1711.10604 (2017)
2017 arXiv
-
[22]
Bingham, J
E. Bingham, J. P. Chen, M. Jankowiak, F. Obermeyer, N. Pradhan, T. Karaletsos, R. Singh, P. A. Szerlip, P. Horsfall, N. D. Goodman, Pyro: Deep Universal Probabilistic Programming, J. Mach. Learn. Res. 20 (2019) 1–6
2019
-
[23]
M. D. Hoffman, D. M. Blei, C. Wang, J. Paisley, Stochastic Variational Inference, J. Mach. Learn. Res. 14 (2013) 1303–1347
2013
-
[24]
Baker, P
J. Baker, P. Fearnhead, E. B. Fox, C. Nemeth, sgmcmc: An R Pack- age for Stochastic Gradient Markov Chain Monte Carlo, J. Stat. Softw. 91 (3) (2019) 1–27
2019
-
[25]
A. K. Gupta, SG-MCMC (2016). URL https://github.com/akshaykgupta/SG_MCMC
2016
-
[26]
Coullon, C
J. Coullon, C. Nemeth, SGMCMCJax: a lightweight JAX library for stochastic gradient Markov chain Monte Carlo algorithms, J. Open Source Softw. 7 (72) (2022) 4113. 18
2022
-
[27]
W. Deng, Q. Feng, L. Gao, F. Liang, G. Lin, Non-convex Learning via Replica Exchange Stochastic Gradient MCMC, in: Proceedings of the 37th International Conference on Machine Learning, PMLR, Online, Jul. 13–18, 2020, pp. 2474–2483
2020
-
[28]
Zhang, A
R. Zhang, A. F. Cooper, C. De Sa, AMAGOLD: Amortized Metropolis adjustment for efficient stochastic gradient MCMC, in: International Conference on Artificial Intelligence and Statistics, PMLR, Online, Aug. 26–28, 2020, pp. 2142–2152
2020
-
[29]
Garriga-Alonso, V
A. Garriga-Alonso, V. Fortuin, Exact Langevin Dynamics with Stochas- tic Gradients, in: 3rd Symposium on Advances in Approximate Bayesian Inference, Online, Jan. – Feb., 2021
2021
-
[30]
Gallego, D
V. Gallego, D. R. Insua, Stochastic Gradient MCMC with Repulsive Forces, in: Bayesian Deep Learning Workshop at NeurIPS, MIT Press, Montreal, Canada, Dec. 7, 2018
2018
-
[31]
K. He, X. Zhang, S. Ren, J. Sun, Deep residual learning for image recognition, in: Proceedings of the IEEE conference on computer vision and pattern recognition, Las Vegas, NV, USA, Jun. 27–30, 2016, pp. 770–778
2016
-
[32]
Krizhevsky, G
A. Krizhevsky, G. Hinton, Learning multiple layers of features from tiny images, Tech. rep., University of Toronto (2009)
2009
-
[33]
W. K. Hastings, Monte Carlo sampling methods using Markov chains and their applications, Biometrika 57 (1) (1970) 97–109
1970
-
[34]
Y.-A. Ma, T. Chen, E. B. Fox, A Complete Recipe for Stochastic Gra- dient MCMC, in: Advances in Neural Information Processing Systems, Vol. 28, MIT Press, Montreal, Canada, 2015, p. 2917–2925
2015
-
[35]
S. Kim, Q. Song, F. Liang, Stochastic gradient Langevin dynamics with adaptive drifts, J. Stat. Comput. Simul. 92 (2) (2022) 318–336
2022
-
[36]
Zhang, C
R. Zhang, C. Li, J. Zhang, C. Chen, A. G. Wilson, Cyclical Stochas- tic Gradient MCMC for Bayesian Deep Learning, in: 7th International Conference on Learning Representations, New Orleans, LA, USA, May 6–9, 2019. 19
2019
-
[37]
Babuschkin, K
I. Babuschkin, K. Baumli, A. Bell, S. Bhupatiraju, J. Bruce, P. Buchlovsky, D. Budden, T. Cai, A. Clark, I. Danihelka, C. Fan- tacci, J. Godwin, C. Jones, R. Hemsley, T. Hennigan, M. Hessel, S. Hou, S. Kapturowski, T. Keck, I. Kemaev, M. King, M. Kunesch, L. Martens, H. Merzic...
2020
-
[38]
Tieleman, G
T. Tieleman, G. Hinton, Lecture 6.5-rmsprop: Divide the gradient by a running average of its recent magnitude, COURSERA: Neural networks for machine learning 4 (2) (2012) 26–31
2012
-
[39]
S. Ahn, A. Korattikara, M. Welling, Bayesian Posterior Sampling via Stochastic Gradient Fisher Scoring, in: Proceedings of the 29th Inter- national Conference on Machine Learning, Omnipress, Madison, WI, USA, Jun. 26 – Jul. 1, 2012, pp. 1771–1778
2012
-
[40]
Y. W. Teh, A. H. Thiery, S. J. Vollmer, Consistency and Fluctuations For Stochastic Gradient Langevin Dynamics, J. Mach. Learn. Res. 17 (2016) 1–33
2016
-
[41]
D. Phan, N. Pradhan, M. Jankowiak, Composable Effects for Flexible and Accelerated Probabilistic Programming in NumPyro, in: Program Transformations for ML at NeurIPS, MIT Press, Vancouver, Canada, Dec. 14, 2019
2019
-
[42]
Hennigan, T
T. Hennigan, T. Cai, T. Norman, I. Babuschkin, Haiku: Sonnet for JAX (2020). URL http://github.com/deepmind/dm-haiku
2020
-
[43]
A. G. Howard, M. Zhu, B. Chen, D. Kalenichenko, W. Wang, T. Weyand, M. Andreetto, H. Adam, MobileNets: Efficient Convolu- tional Neural Networks for Mobile Vision Applications, arXiv preprint arXiv:1704.04861 (2017)
2017 arXiv
-
[44]
J. Kim, S. Choi, Automated machine learning for soft voting in an en- semble of tree-based classifiers, in: International Workshop on Auto- matic Machine Learning at ICML, Stockholm, Sweden, Jul. 14, 2018. 20
2018
-
[45]
Thaler, G
S. Thaler, G. Doehner, J. Zavadlav, Scalable Bayesian Uncertainty Quantification for Neural Network Potentials: Promise and Pitfalls, J. Chem. Theory Comput. (2023)
2023
-
[46]
Wang, D.-Y
H. Wang, D.-Y. Yeung, A survey on bayesian deep learning, ACM Com- put. Surv. 53 (5) (2020) 1–37
2020
-
[47]
P. Ren, Y. Xiao, X. Chang, P.-Y. Huang, Z. Li, B. B. Gupta, X. Chen, X. Wang, A survey of deep active learning, ACM Comput. Surv. 54 (9) (2021) 1–40
2021
-
[48]
A. G. Wilson, P. Izmailov, Bayesian Deep Learning and a Probabilis- tic Perspective of Generalization, in: Advances in Neural Information Processing Systems, Vol. 33, Online, Dec. 6–12, 2020, pp. 4697–4708
2020
-
[49]
Y. Gal, Z. Ghahramani, Dropout as a bayesian approximation: Repre- senting model uncertainty in deep learning, in: International Conference on Machine Learning, PMLR, 2016, pp. 1050–1059
2016
-
[50]
Hansen, P
L. Hansen, P. Salamon, Neural Network Ensembles, IEEE Trans. Pat- tern Anal. Machine Intell. 12 (10) (1990) 993–1001
1990
-
[51]
Thaler, M
S. Thaler, M. Stupp, J. Zavadlav, Deep coarse-grained potentials via relative entropy minimization, J. Chem. Phys. 157 (2022) 244103
2022
-
[52]
Thaler, J
S. Thaler, J. Zavadlav, Uncertainty Quantification for Molecular Models via Stochastic Gradient MCMC, in: 10th Vienna Conference on Math- ematical Modelling, Vienna, Austria, Jul. 27–29, 2022, pp. 19–20. 21
2022
Reviewed August 15, 2026 · model on record in the stance chip above.
Discussion (0). Continue with ORCID to comment.