Skip to content

More stable incomplete beta gradient roots - #3409

Open
avehtari wants to merge 4 commits into
developfrom
stable-inc-beta
Open

avehtari wants to merge 4 commits into
developfrom
stable-inc-beta

Conversation

@avehtari

Copy link
Copy Markdown
Member

Replace the incomplete beta gradient roots with a continued-fraction implementation from TensorFlow Probability (this PR was Claude assisted)

Summary

grad_reg_inc_beta, inc_beta_dda and inc_beta_ddb compute the derivatives of the regularized incomplete beta function I_z(a, b) with respect to its two shape parameters. They are two separate implementations of the same quantity. Path A (inc_beta_dda, inc_beta_ddb) serves beta_cdf, neg_binomial_cdf and neg_binomial_2_cdf; path B (grad_reg_inc_beta, through grad_2F1) serves every lcdf and lccdf of beta, beta_proportion, neg_binomial, student_t, plus student_t_cdf and the autodiff of inc_beta itself. Both have defects that make the gradient NaN, or wrong by O(1), in regions that ordinary models reach. This PR replaces both with one implementation: the Boik and Robinson-Cox (1998) algorithm in the form implemented by TensorFlow Probability, a power series near the endpoints and a modified-Lentz continued fraction elsewhere, both differentiated term by term, with the prefactor in log space. inc_beta_dda and inc_beta_ddb become wrappers; beta_cdf makes one root call instead of two. Signatures do not change.

The values of the distributions do not change. Only the gradients change.

Comparison, both trees compiled and run over the same 18 000 points in nine regions (details below):

develop branch
non-finite, path B, large shapes near the mode 1325 of 2000 0
non-finite, path A (throws), all regions 100 0
gradient of log I wrong by > 1e-2, path A, b ≫ a 408 of 2000 0
gradient of log I wrong by > 1e-2, path B, negbin n > 200 1203 of 1999 0
worst error of the gradient of log I, branch — 1e-11
worst relative error where the derivative is not negligible, branch — 3.7e-9
spread between double, var, fvar, fvar<fvar> 0 0
speed, grad_reg_inc_beta in double 1x 30x
speed, beta_lcdf per observation, α = β = 5 1x 40x
speed, beta_cdf per observation, α = β = 50 1x 2.4x

The defects in the current code on develop

Everything in this section describes develop at 5252d51d47. All numbers were measured against an mpmath reference at 80 digits built through two independent routes (density-form quadrature and numerical differentiation at 110 digits) that agree at 60 digits at 54 check points, and confirmed by compiling and running the real develop headers.

1. grad_reg_inc_beta returns NaN for every parameter when beta(a, b) underflows. Its callers pass beta(a, b) as betaAB and it divides by it. beta(a, b) underflows for a + b above about 1100 with a near b. beta_lcdf(y | 600, 600) has a NaN gradient with respect to both shapes today, and so does inc_beta(a, b, z) with autodiff shapes. In the sweep region a, b in [100, 12500] near the mode, the region where I is of order 1, 1325 of 2000 points are NaN.

2. grad_reg_inc_beta drops the hypergeometric term when z^a (1-z)^b underflows. When C = z^a (1-z)^b / a is 0 it skips grad_2F1 and returns I (log z - 1/a - ψ(a) + ψ(a+b)). The dropped term is the whole answer when I is near 1, because 2F1(a+b, 1; a+1; z) grows like (1-z)^(-b). beta_lcdf(0.758 | 352, 590) returns d/dα = 0.705; the true value is -3.4e-138. neg_binomial_lcdf(132 | 313, 0.068) returns d/dα = 0 for a true -2.3.

3. inc_beta_dda and inc_beta_ddb stop on an absolute threshold that can be met before the first term. Both stop the series when fabs(summand) < 1e-10, and the first summand carries ((a+1)/(a+b))^3. For b ≫ a that is below 1e-10 before the loop starts, and the returned ratio is its first term. beta_cdf(1.8e-4 | 1.145, 6786) returns d/dα = -0.203 for a true -0.386, at I = 0.65. In the sweep region b ≫ a, small z, 1301 of 2000 points have an error above 1e-6 in the gradient of log I and 408 above 1e-2.

4. inc_beta_dda and inc_beta_ddb reflect on fixed z thresholds, not on which side is small. For z > 0.75 (and other fixed bands) they evaluate 1 - I_{1-z}(b, a) regardless of where the mass is. In the deep lower tail the result is then I' × (a bracket that must cancel to 1e-47) with I' ≈ 1, and the bracket comes out as 1e-15 noise. beta_cdf(0.76 | 446, 5) returns d/dα = 8.9e-16 for a true -1.06e-47; log(beta_cdf(...)) then has a gradient of 2.5e+31. Some large-shape points also reach the k > 1e5 iteration guard and throw, inside the range the docstring calls tested.

5. The existing unit tests assert the wrong values. inc_beta_ddb_test.cpp passes digamma(a) where the fourth parameter is digamma(b), and its expected constants were generated from that call. Both inc_beta_dda_test.cpp and inc_beta_ddb_test.cpp also assert develop's own output at points where it is wrong by orders of magnitude: 1.8226241 for a true -3.94e-5, 9.3959293 for a true 1.48e-8. grad_reg_inc_beta_test.cpp expects d/db = -inf at z = 1, where I ≡ 1 and the derivative is 0.

What this PR does

stan/math/prim/fun/grad_reg_inc_beta.hpp gets a new body, a C++ port of _betainc_partials from TensorFlow Probability (tensorflow_probability/python/math/special.py, Apache-2.0; see the licence note below). The algorithm is Boik and Robinson-Cox (1998), Derivatives of the Incomplete Beta Function, JSS 3(1): the DLMF 8.17.22 continued fraction evaluated by the modified Lentz method with the derivatives of the partial numerators carried through the recurrence, plus a power-series region near the endpoints from Cephes, also differentiated term by term. The prefactor z^a (1-z)^b / (a B(a, b)) is evaluated in log space, which removes defects 1 and 2. The symmetry relation I_z(a, b) = 1 - I_{1-z}(b, a) is applied on the mode for the continued fraction and on the mean for the series, which removes defect 4. There is no absolute threshold; the series stops at eps / a and the fraction at 3 eps, which removes defect 3.

Two changes from the TFP source, both marked in the code. TFP's stop tests watch the value only. When b is a positive integer, the series factor (n - b) and the continued-fraction numerator d_{2m} are exactly 0 at n = m = b; the value is then converged but the derivatives still need the tail of the series. TFP's own d/db is off by 4.2e-4 at (2, 3, 0.25) and by 6.7 % at (10, 1, 0.5), and its d/da by 5.6 % at (1, 100, 0.01), where the symmetry relation makes b' = 1. The port stops only when the derivative increments have also converged. Integer shapes are common (initial values, fixed shapes). The defect will be reported upstream.

inc_beta_dda and inc_beta_ddb become wrappers over grad_reg_inc_beta that return one of the two derivatives. beta_cdf calls grad_reg_inc_beta once instead of inc_beta_dda and inc_beta_ddb separately. prim/fun.hpp includes inc_beta_dda.hpp, inc_beta_ddb.hpp and inc_beta_ddz.hpp, which it previously reached only through prob/beta_cdf.hpp. The betaAB argument of grad_reg_inc_beta is kept for interface compatibility and is no longer used.

Net change: about 180 lines of algorithm replacing 123.

Accuracy after the change

Same 18 000 points, same reference, both trees compiled with the same probe. The metric is the absolute error of the gradient of log I, which is what beta_lcdf returns and what log(beta_cdf(...)) differentiates to; a relative error on a derivative that is 1e-300 of I is not a useful number.

region (2 000 points each) develop: non-finite / error > 1e-2 branch: worst error
a, b < 1 0 / 0 1e-13
1 ≤ a, b ≤ 20, z < 0.75 0 / 0 3e-14
z > 0.75, a < 500 path A 3 / 81, path B 3 / 125 7e-15
b > a, b > 500, z ≤ 0.75 path A 10 / 86, path B 550 / 933 3e-13 (1)
b ≫ a, z small path A 2 / 408, path B 2 / 12 3.5e-9
large a, b, z uniform path A 46 / 36, path B 932 / 235 4e-12 (1)
large a, b, z near the mode path A 0 / 0, path B 1325 / 9 1e-13
negbin pattern, n ≤ 200 path A 0 / 4, path B 0 / 26 4e-13
negbin pattern, 200 < n ≤ 5000 path A 39 / 22, path B 292 / 1203 4e-13

(1) Excluding points where I itself is subnormal in double.

The relative error where the derivative is not negligible (|dI| / I ≥ 1e-3) is at or below 6e-13 in every region except b ≫ a, small z, where one continued-fraction point at the symmetry switch reaches 3.7e-9. The four scalar configurations double, var, fvar<double>, fvar<fvar<double>> give bit-identical values on both trees, and the branch's autodiff columns have the same error profile as double.

Speed

Xeon E5-2680 v3 at 2.50 GHz, exclusive node, performance governor, two rounds agreeing within 2 %. Grid a, b in [0.5, 60], z in [0.05, 0.95]. The loop adds an accumulator-dependent term to every input so it is not loop-invariant; a baseline is subtracted; the cycle counts of the baseline (6) and of 3 digamma + beta (310) show nothing was removed by the compiler.

Function level, net ns per call:

kernel develop branch ratio
grad_reg_inc_beta double, incl. 3 digamma + beta 28 022 922 30x
grad_reg_inc_beta var, value + grad 157 825 9 670 16x
grad_reg_inc_beta fvar 88 108 1 776 50x
grad_reg_inc_beta fvar<fvar> 255 280 3 730 68x
inc_beta_dda double, incl. 2 digamma 1 880 866 2.2x

Distribution level, N = 1000 observations, shared var shapes, value plus reverse sweep, ns per observation:

function shapes develop branch ratio
beta_cdf 5, 5 1 682 1 107 1.5x
beta_lcdf 5, 5 25 822 642 40x
beta_lccdf 5, 5 25 810 642 40x
beta_cdf 50, 50 6 385 2 653 2.4x
beta_lcdf 50, 50 21 295 1 608 13x
beta_lccdf 50, 50 21 297 1 607 13x
neg_binomial_cdf 5, 0.3 1 895 1 724 1.1x
neg_binomial_lcdf 5, 0.3 7 677 983 7.8x
student_t_lcdf ν = 3 4 938 1 051 4.7x

develop's path B cost is the 2F1 power series with three exp and three log per term and no reflection; at z = 0.998 it takes 607 607 terms. The branch's continued fraction costs about 600 cycles above the digamma calls.

Testing

Six test files, all with fixed references from the 80-digit mpmath computation, not from finite differences, which cannot see these defects:

  • prim/fun/grad_reg_inc_beta_test.cpp: the two existing cases (with the z = 1 expectation corrected to 0 and the tolerance tightened to 1e-12), plus cases for each defect class: beta(a, b) underflow, prefactor underflow, integer b, and the tails.
  • prim/fun/inc_beta_dda_test.cpp, prim/fun/inc_beta_ddb_test.cpp: rewritten from the reference at the same twelve points, with the digamma_b argument corrected, plus regression points.
  • rev/fun/inc_beta_test.cpp (new): shape gradients of inc_beta with var arguments at the defect points.
  • rev/prob/beta_cdf_test.cpp (new): beta_cdf shape partials at the b ≫ a point, with both shapes autodiff and with each shape alone.
  • rev/prob/beta_cdf_log_test.cpp (new): beta_lcdf shape gradients at the NaN point, the dropped-term point and an integer shape.

Regression swap: with the develop headers installed under these six files, 6 of 6 suites fail, 13 of 15 cases; with this branch's headers, 15 of 15 pass, before and after the swap. The two cases that pass on develop are the points where it was already correct.

Licence note

The new body of grad_reg_inc_beta.hpp is derived from TensorFlow Probability, which is Apache-2.0. The file carries the TFP copyright notice and points at licenses/tensorflow-probability-license.txt, following the form used for the Boost-derived code in prim/fun/log_modified_bessel_first_kind.hpp. Apache-2.0 code is already in the tree as vendored libraries (lib/tbb_2020.3, lib/benchmark_1.5.1); this would be the first inline derived file under stan/.

Known limits, not addressed here

  • beta_lccdf forms 1 - Pn in linear space after inc_beta. At (352, 590, 0.758) the double I is exactly 1, the value is -inf and the gradients are infinite, where the true value is log(1 - I) ≈ -0.34. That is in the value code, on both trees, and is a separate change (the symmetry relation for the value).
  • At I below about 1e-315 (subnormal), the gradient's relative error rises to 1e-2. The reference itself has few digits there.

Not included

  • No change to grad_2F1, which other functions use.
  • No change to the neg_binomial*_cdf call pattern; they need one derivative and already make one root call.

Release notes

Replaces the incomplete beta gradient roots with a more stable continued-fraction implementation

Checklist

  • Copyright holder: Aalto University

    The copyright holder is typically you or your assignee, such as a university or company. By submitting this pull request, the copyright holder is agreeing to the license the submitted work under the following licenses:
    - Code: BSD 3-clause (https://opensource.org/licenses/BSD-3-Clause)
    - Documentation: CC-BY 4.0 (https://creativecommons.org/licenses/by/4.0/)

  • the basic tests are passing

    • unit tests pass (to run, use: ./runTests.py test/unit)
    • header checks pass, (make test-headers)
    • dependencies checks pass, (make test-math-dependencies)
    • docs build, (make doxygen)
    • code passes the built in C++ standards checks (make cpplint)
  • the code is written in idiomatic C++ and changes are documented in the doxygen

  • the new changes are tested

@avehtari

Copy link
Copy Markdown
Member Author

I'll be away for the next two weeks, but made this PR to avoid duplicated work

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant