Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
28 commits
Select commit Hold shift + click to select a range
056dd03
Add erfcx scaled complementary error function
avehtari Sep 15, 2026
e0707b1
Flatten erfcx Horner evaluation to satisfy cpplint line length
avehtari Sep 16, 2026
266fb5c
Use fma instead of Dekker split in OpenCL erfcx device function
avehtari Sep 16, 2026
8c16f29
Use fma instead of Dekker split in erfcx lower branch
avehtari Sep 16, 2026
ea151be
Use Cody second-interval rational for erfcx on 0.46875 <= x < 4
avehtari Sep 17, 2026
f55de06
Replace erfcx erfc call below 0.46875 with a Chebyshev expansion and …
avehtari Sep 17, 2026
7dfed73
Mirror the new erfcx branches in the OpenCL device function and tests
avehtari Sep 17, 2026
bf74584
Merge commit '5252d51d47c1d5e78005fc043ad996fad6dd8da8' into HEAD
yashikno Sep 17, 2026
37d8262
[Jenkins] auto-formatting by clang-format version 10.0.0-4ubuntu1
stan-buildbot Sep 17, 2026
6ff4541
Add reverse-mode OpenCL support for erfcx
avehtari Sep 17, 2026
6630449
Add OpenCL tests for erfcx
avehtari Sep 17, 2026
b9c353f
Add fixed-reference tests for the erfcx derivative in the tail
avehtari Sep 18, 2026
baebb7e
Take the erfcx derivative from the tail rational to avoid cancellation
avehtari Sep 18, 2026
ce4dea1
Split the erfcx small-branch polynomial and merge the reverse-mode ov…
avehtari Sep 18, 2026
b9be613
[Jenkins] auto-formatting by clang-format version 10.0.0-4ubuntu1
stan-buildbot Sep 18, 2026
dadf537
Add OpenCL tail tests for erfcx including fixed-reference derivatives
avehtari Sep 18, 2026
3453e8a
Take the OpenCL erfcx derivative from the tail rational to avoid canc…
avehtari Sep 18, 2026
d997ed8
Merge branch 'stable-erfcx' of github.com:stan-dev/math into stable-e…
avehtari Sep 18, 2026
3b1ded9
[Jenkins] auto-formatting by clang-format version 10.0.0-4ubuntu1
stan-buildbot Sep 18, 2026
8663ecc
Apply clang-format 10 to the OpenCL erfcx files
avehtari Sep 18, 2026
2933f28
Apply batched suggestions from code review
avehtari Sep 18, 2026
56f3061
Merge branch 'stable-erfcx' of github.com:stan-dev/math into stable-e…
avehtari Sep 18, 2026
0b1094f
[Jenkins] auto-formatting by clang-format version 10.0.0-4ubuntu1
stan-buildbot Sep 18, 2026
9c8ddfd
Merge remote-tracking branch 'origin/stable-erfcx' into stable-erfcx
avehtari Sep 18, 2026
cd7227a
Fix the out-of-bounds erfcx_small loop and the C++ constructs in the …
avehtari Sep 18, 2026
6aff1cd
Keep the previous OpenCL erfcx device function, which is 8 to 13 perc…
avehtari Sep 18, 2026
17fcb9d
Apply the new review suggestions; pair the Cody middle coefficients w…
avehtari Sep 18, 2026
3612a0b
Apply the 2026-09-19 review; keep the cheap derivative and share the …
avehtari Sep 19, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions stan/math/fwd/fun.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
#include <stan/math/fwd/fun/digamma.hpp>
#include <stan/math/fwd/fun/erf.hpp>
#include <stan/math/fwd/fun/erfc.hpp>
#include <stan/math/fwd/fun/erfcx.hpp>
#include <stan/math/fwd/fun/exp.hpp>
#include <stan/math/fwd/fun/exp2.hpp>
#include <stan/math/fwd/fun/expm1.hpp>
Expand Down
33 changes: 33 additions & 0 deletions stan/math/fwd/fun/erfcx.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
#ifndef STAN_MATH_FWD_FUN_ERFCX_HPP
#define STAN_MATH_FWD_FUN_ERFCX_HPP

#include <stan/math/fwd/meta.hpp>
#include <stan/math/fwd/core.hpp>
#include <stan/math/prim/fun/constants.hpp>
#include <stan/math/prim/fun/erfcx.hpp>
#include <cmath>

namespace stan {
namespace math {

/**
* Return the scaled complementary error function of the argument.
*
* The derivative comes from `internal::erfcx_derivative`. Below 4 it is
* `2 * x * erfcx(x) - 2 / sqrt(pi)`, which reuses the value. At 4 and above
* that difference cancels, so the tail rational supplies the derivative
* directly.
*
* @tparam T inner type of the fvar
* @param x argument
* @return scaled complementary error function of the argument
*/
template <typename T>
inline fvar<T> erfcx(const fvar<T>& x) {
T v = erfcx(x.val_);
return fvar<T>(v, x.d_ * internal::erfcx_derivative(x.val_, v));
}

} // namespace math
} // namespace stan
#endif
4 changes: 4 additions & 0 deletions stan/math/opencl/kernel_generator/elt_function_cl.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
#include <stan/math/opencl/kernels/device_functions/binomial_coefficient_log.hpp>
#include <stan/math/opencl/kernels/device_functions/beta.hpp>
#include <stan/math/opencl/kernels/device_functions/digamma.hpp>
#include <stan/math/opencl/kernels/device_functions/erfcx.hpp>
#include <stan/math/opencl/kernels/device_functions/inv_logit.hpp>
#include <stan/math/opencl/kernels/device_functions/inv_Phi.hpp>
#include <stan/math/opencl/kernels/device_functions/inv_square.hpp>
Expand Down Expand Up @@ -297,6 +298,7 @@ ADD_UNARY_FUNCTION_PASS_ZERO(trunc)

ADD_UNARY_FUNCTION_WITH_INCLUDES(digamma,
opencl_kernels::digamma_device_function)
ADD_UNARY_FUNCTION_WITH_INCLUDES(erfcx, opencl_kernels::erfcx_device_function)
ADD_UNARY_FUNCTION_WITH_INCLUDES(log1m, opencl_kernels::log1m_device_function)
ADD_UNARY_FUNCTION_WITH_INCLUDES(log_inv_logit,
opencl_kernels::log1p_exp_device_function,
Expand Down Expand Up @@ -351,6 +353,8 @@ ADD_BINARY_FUNCTION_WITH_INCLUDES(ldexp)
ADD_BINARY_FUNCTION_WITH_INCLUDES(pow)
ADD_BINARY_FUNCTION_WITH_INCLUDES(copysign)

ADD_BINARY_FUNCTION_WITH_INCLUDES(
erfcx_derivative, stan::math::opencl_kernels::erfcx_device_function)
ADD_BINARY_FUNCTION_WITH_INCLUDES(
beta, stan::math::opencl_kernels::beta_device_function)
ADD_BINARY_FUNCTION_WITH_INCLUDES(
Expand Down
205 changes: 205 additions & 0 deletions stan/math/opencl/kernels/device_functions/erfcx.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,205 @@
#ifndef STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_ERFCX_HPP
#define STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_ERFCX_HPP
#ifdef STAN_OPENCL

#include <stan/math/opencl/stringify.hpp>
#include <string>

namespace stan {
namespace math {
namespace opencl_kernels {

// \cond
static constexpr const char* erfcx_device_function
= "\n"
"#ifndef STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_ERFCX\n"
"#define STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_ERFCX\n" STRINGIFY(
// \endcond
/** \ingroup opencl_kernels
*
* Correction factor of the Cody (1969) third-interval rational:
* erfcx(x) = (INV_SQRT_PI + u * correction(u)) / x, u = 1 / x^2.
*
* Split out so the derivative can reuse it. See
* erfcx_tail_derivative.
*
* @param u inverse square of the argument, 0 <= u <= 1/16
* @return P(u) / Q(u)
*/
double erfcx_tail_correction(double u) {
double p = 0.0163153871373020978498;
p = 0.305326634961232344035 + u * p;
p = 0.360344899949804439429 + u * p;
p = 0.125781726111229246204 + u * p;
p = 0.0160837851487422766278 + u * p;
p = 0.000658749161529837803157 + u * p;
double q = -1.0;
q = -2.56852019228982242072 + u * q;
q = -1.87295284992346047209 + u * q;
q = -0.527905102951428412248 + u * q;
q = -0.0605183413124413191178 + u * q;
q = -0.00233520497626869185443 + u * q;
return p / q;
}

/** \ingroup opencl_kernels
*
* Cody (1969) third-interval rational, for x >= 4.
*
* @param x argument
* @return scaled complementary error function
*/
double erfcx_cody_tail(double x) {
double u = 1.0 / (x * x);
return (M_2_SQRTPI * 0.5 + u * erfcx_tail_correction(u)) / x;
}

/** \ingroup opencl_kernels
*
* Derivative of erfcx, 2 * x * erfcx(x) - 2 / sqrt(pi).
*
* That difference cancels for large x: both terms approach
* 2 / sqrt(pi) while the result decays like 1 / (sqrt(pi) * x^2).
* For x >= 4 the constant cancels analytically against the
* leading term of the tail rational, leaving 2 * u * C(u).
* Above 30 it comes from the asymptotic expansion of
* -sqrt(pi) * x^2 * erfcx'(x) in u, whose coefficients satisfy
* c_{n+1} = -(n + 3/2) * c_n.
*
* Takes the value as an argument so the reverse pass does not
* evaluate erfcx twice.
*
* @param x argument
* @param value erfcx(x)
* @return derivative of erfcx at x
*/
double erfcx_derivative(double x, double value) {
if (x < 4.0) {
return 2.0 * x * value - M_2_SQRTPI;
}
double u = 1.0 / (x * x);
if (x < 30.0) {
return 2.0 * u * erfcx_tail_correction(u);
}
// -sqrt(pi) * x^2 * erfcx'(x), asymptotic, ascending in u
const double s[8]
= {1.0, -1.5, 3.75, -13.125,
59.0625, -324.84375, 2111.484375, -15836.1328125};
double series = s[7];
for (int i = 6; i >= 0; --i) {
series = series * u + s[i];
}
return -(0.5 * M_2_SQRTPI) * u * series;
}

/** \ingroup opencl_kernels
*
* Cody (1969) second-interval rational, for 0.46875 <= x <= 4.
* Yields erfcx directly: the exponential is cancelled
* analytically, so there is no exp and no erfc call.
*
* @param y argument
* @return scaled complementary error function
*/
double erfcx_cody_middle(double y) {
double p = 2.15311535474403846e-8 * y;
p = (p + 5.64188496988670089e-1) * y;
p = (p + 8.88314979438837594) * y;
p = (p + 66.1191906371416295) * y;
p = (p + 298.635138197400131) * y;
p = (p + 881.952221241769090) * y;
p = (p + 1712.04761263407058) * y;
p = (p + 2051.07837782607147) * y;
double q = y;
q = (q + 15.7449261107098347) * y;
q = (q + 117.693950891312499) * y;
q = (q + 537.181101862009858) * y;
q = (q + 1621.38957456669019) * y;
q = (q + 3290.79923573345963) * y;
q = (q + 4362.61909014324716) * y;
q = (q + 3439.36767414372164) * y;
return (p + 1230.33935479799725) / (q + 1230.33935480374942);
}

/** \ingroup opencl_kernels
*
* Degree-18 Chebyshev-economized expansion of erfcx about zero,
* for |x| < 0.46875. Covers both signs with no branch and no
* library call.
*
* @param x argument
* @return scaled complementary error function
*/
double erfcx_small(double x) {
double p = 3.05977060678449757e-06;
p = -9.35890030086883823e-06 + x * p;
p = 2.46655529768908249e-05 + x * p;
p = -7.08163358203131886e-05 + x * p;
p = 1.98445679338826757e-04 + x * p;
p = -5.34506929034156810e-04 + x * p;
p = 1.38888415444527033e-03 + x * p;
p = -3.47359067853470795e-03 + x * p;
p = 8.33333374332981443e-03 + x * p;
p = -1.91048337772546720e-02 + x * p;
p = 4.16666666458337179e-02 + x * p;
p = -8.59717459974174147e-02 + x * p;
p = 1.66666666667239644e-01 + x * p;
p = -3.00901111227312890e-01 + x * p;
p = 4.99999999999992839e-01 + x * p;
p = -7.52252778063651983e-01 + x * p;
p = 1.0 + x * p;
p = -1.12837916709551256 + x * p;
return 1.0 + x * p;
}

/** \ingroup opencl_kernels
*
* Return the scaled complementary error function
* exp(x * x) * erfc(x) of the kernel generator expression.
*
* Mirrors stan/math/prim/fun/erfcx.hpp branch for branch and
* formula for formula; the two must be kept in step. The whole
* positive axis is covered without a library call. On the
* negative side exp is unavoidable, since erfcx grows like
* 2*exp(x*x); its argument is corrected with fma, because exp
* amplifies the rounding of x * x into roughly x * x * eps.
*
* The correction must be written with fma and not with a Dekker
* split. The split relies on t - (t - x) not being simplified
* to x, which holds in floating point but not over the reals,
* and the OpenCL compiler does simplify it: the split form
* silently degrades to the uncorrected product on device while
* still being correct on the host. Do not reintroduce it.
*
* @param x argument
* @return scaled complementary error function of the argument
*/
double erfcx(double x) {
if (x >= 4.0) {
return erfcx_cody_tail(x);
} else if (x >= 0.46875) {
return erfcx_cody_middle(x);
} else if (x > -0.46875) {
return erfcx_small(x);
} else if (x < -27.0) {
return INFINITY;
}
Comment thread
avehtari marked this conversation as resolved.
double h = x * x;
double two_exp_x2 = 2.0 * exp(h) * (1.0 + fma(x, x, -h));
if (x < -6.1) {
return two_exp_x2;
}
double y = -x;
return two_exp_x2
- (y >= 4.0 ? erfcx_cody_tail(y) : erfcx_cody_middle(y));
}
// \cond
) "\n#endif\n"; // NOLINT
// \endcond

} // namespace opencl_kernels
} // namespace math
} // namespace stan

#endif
#endif
1 change: 1 addition & 0 deletions stan/math/opencl/rev.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
#include <stan/math/opencl/rev/elt_multiply.hpp>
#include <stan/math/opencl/rev/erf.hpp>
#include <stan/math/opencl/rev/erfc.hpp>
#include <stan/math/opencl/rev/erfcx.hpp>
#include <stan/math/opencl/rev/exp.hpp>
#include <stan/math/opencl/rev/exp2.hpp>
#include <stan/math/opencl/rev/expm1.hpp>
Expand Down
40 changes: 40 additions & 0 deletions stan/math/opencl/rev/erfcx.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
#ifndef STAN_MATH_OPENCL_REV_ERFCX_HPP
#define STAN_MATH_OPENCL_REV_ERFCX_HPP
#ifdef STAN_OPENCL

#include <stan/math/opencl/kernel_generator.hpp>
#include <stan/math/rev/core.hpp>
#include <stan/math/rev/fun/value_of.hpp>

namespace stan {
namespace math {

/**
* Returns the elementwise `erfcx()` of a var_value<matrix_cl<double>>.
*
* The derivative `2 * x * erfcx(x) - 2 / sqrt(pi)` reuses the function
* value, so no second `exp` or `erfc` evaluation is needed. That difference
* cancels for `x >= 4`, where both terms approach `2 / sqrt(pi)` while the
* result decays like `1 / (sqrt(pi) * x^2)`; the device function
* `erfcx_derivative` takes the derivative from the tail rational there
* instead. Without that the error reaches 2.55e+11 ulp at `x = 1e6`. The
* CPU implementation branches at the same point.
*
* @param A argument
* @return Elementwise `erfcx()` of the input.
*/
template <typename T,
require_all_kernel_expressions_and_none_scalar_t<T>* = nullptr>
inline var_value<matrix_cl<double>> erfcx(const var_value<T>& A) {
return make_callback_var(
erfcx(A.val()), [A](vari_value<matrix_cl<double>>& res) mutable {
A.adj()
+= elt_multiply(res.adj(), erfcx_derivative(A.val(), res.val()));
});
}

} // namespace math
} // namespace stan

#endif
#endif
1 change: 1 addition & 0 deletions stan/math/prim/fun.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,7 @@
#include <stan/math/prim/fun/elt_multiply.hpp>
#include <stan/math/prim/fun/erf.hpp>
#include <stan/math/prim/fun/erfc.hpp>
#include <stan/math/prim/fun/erfcx.hpp>
#include <stan/math/prim/fun/eval.hpp>
#include <stan/math/prim/fun/exp.hpp>
#include <stan/math/prim/fun/exp2.hpp>
Expand Down
Loading
Loading