Skip to content

Commit dc2dd69

Browse files
authored
Feature: Add std::comp_ellint derivatives (#1674)
Fixes #1673
1 parent c4d93fd commit dc2dd69

2 files changed

Lines changed: 116 additions & 3 deletions

File tree

include/clad/Differentiator/BuiltinDerivatives.h

Lines changed: 87 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1213,6 +1213,85 @@ CUDA_HOST_DEVICE void beta_pullback(T x, T y, U d_z, T* d_x, T* d_y) {
12131213
}
12141214
#endif
12151215

1216+
#if __cplusplus >= 201703L && (defined(__cpp_lib_math_special_funcs) || \
1217+
defined(__STDCPP_MATH_SPEC_FUNCS__))
1218+
template <typename T, typename dT,
1219+
typename T_out = decltype(::std::comp_ellint_1(T())),
1220+
typename dT_out = typename AdjOutType<T_out, dT>::type>
1221+
CUDA_HOST_DEVICE ValueAndPushforward<T_out, dT_out>
1222+
comp_ellint_1_pushforward(T k, dT d_k) {
1223+
T_out K = ::std::comp_ellint_1(k);
1224+
T_out E = ::std::comp_ellint_2(k);
1225+
T_out k_sq = k * k;
1226+
T_out one = 1.0;
1227+
T_out derivative = 0.0;
1228+
if (k_sq < 1.0 && k != 0.0)
1229+
derivative = (E - (one - k_sq) * K) / (k * (one - k_sq));
1230+
return {K, static_cast<dT_out>(d_k) * derivative};
1231+
}
1232+
1233+
template <typename T, typename dT,
1234+
typename T_out = decltype(::std::comp_ellint_2(T())),
1235+
typename dT_out = typename AdjOutType<T_out, dT>::type>
1236+
CUDA_HOST_DEVICE ValueAndPushforward<T_out, dT_out>
1237+
comp_ellint_2_pushforward(T k, dT d_k) {
1238+
T_out K = ::std::comp_ellint_1(k);
1239+
T_out E = ::std::comp_ellint_2(k);
1240+
T_out derivative = 0.0;
1241+
if (k != 0.0)
1242+
derivative = (E - K) / k;
1243+
return {E, static_cast<dT_out>(d_k) * derivative};
1244+
}
1245+
1246+
template <typename T1, typename T2, typename dT1, typename dT2,
1247+
typename T_out = decltype(::std::comp_ellint_3(T1(), T2())),
1248+
typename dT_out = typename AdjOutType<T_out, dT1>::type>
1249+
CUDA_HOST_DEVICE ValueAndPushforward<T_out, dT_out>
1250+
comp_ellint_3_pushforward(T1 k, T2 nu, dT1 d_k, dT2 d_nu) {
1251+
T_out K = ::std::comp_ellint_1(k);
1252+
T_out E = ::std::comp_ellint_2(k);
1253+
T_out Pi = ::std::comp_ellint_3(k, nu);
1254+
T_out k2 = k * k;
1255+
T_out one = 1.0;
1256+
T_out grad_k = 0.0;
1257+
if (k2 != static_cast<T_out>(nu) && k != 0.0 && k2 < 1.0) {
1258+
T_out term_k = E / (one - k2);
1259+
grad_k = (k / (k2 - static_cast<T_out>(nu))) * (term_k - Pi);
1260+
}
1261+
T_out grad_nu = 0.0;
1262+
if (nu != 0.0 && nu != 1.0 && k2 != static_cast<T_out>(nu)) {
1263+
T_out p2 = ((k2 - static_cast<T_out>(nu)) / static_cast<T_out>(nu)) * K;
1264+
T_out p3 = ((static_cast<T_out>(nu) * static_cast<T_out>(nu) - k2) /
1265+
static_cast<T_out>(nu)) *
1266+
Pi;
1267+
grad_nu = (one / (2.0 * (static_cast<T_out>(nu) - one) *
1268+
(k2 - static_cast<T_out>(nu)))) *
1269+
(E + p2 + p3);
1270+
}
1271+
return {Pi, (static_cast<dT_out>(d_k) * grad_k) +
1272+
(static_cast<dT_out>(d_nu) * grad_nu)};
1273+
}
1274+
1275+
template <typename T1, typename T2, typename T3>
1276+
CUDA_HOST_DEVICE void comp_ellint_3_pullback(T1 k, T2 nu, T3 d_out, T1* d_k,
1277+
T2* d_nu) {
1278+
auto K = ::std::comp_ellint_1(k);
1279+
auto E = ::std::comp_ellint_2(k);
1280+
auto Pi = ::std::comp_ellint_3(k, nu);
1281+
auto k2 = k * k;
1282+
auto one = 1.0;
1283+
if (k2 != nu && k != 0.0 && k2 < 1.0) {
1284+
auto term_k = E / (one - k2);
1285+
*d_k += d_out * ((k / (k2 - nu)) * (term_k - Pi));
1286+
}
1287+
if (nu != 0.0 && nu != one && k2 != nu) {
1288+
auto p2 = ((k2 - nu) / nu) * K;
1289+
auto p3 = ((nu * nu - k2) / nu) * Pi;
1290+
*d_nu += d_out * ((one / (2.0 * (nu - one) * (k2 - nu))) * (E + p2 + p3));
1291+
}
1292+
}
1293+
#endif
1294+
12161295
} // namespace std
12171296

12181297
CUDA_HOST_DEVICE inline ValueAndPushforward<float, float>
@@ -1489,6 +1568,14 @@ using std::beta_pullback;
14891568
using std::beta_pushforward;
14901569
#endif
14911570

1571+
#if __cplusplus >= 201703L && (defined(__cpp_lib_math_special_funcs) || \
1572+
defined(__STDCPP_MATH_SPEC_FUNCS__))
1573+
using std::comp_ellint_1_pushforward;
1574+
using std::comp_ellint_2_pushforward;
1575+
using std::comp_ellint_3_pullback;
1576+
using std::comp_ellint_3_pushforward;
1577+
#endif
1578+
14921579
namespace class_functions {
14931580
template <typename T, typename U>
14941581
void constructor_pullback(ValueAndPushforward<T, U> rhs,

test/Features/stl-cmath.cpp

Lines changed: 29 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -99,9 +99,9 @@
9999
// D assoc_laguerre / f / l (C++17) associated Laguerre polynomials
100100
// D assoc_legendre/ f / l (C++17) associated Legendre polynomials
101101
// DS beta/ betaf/ betal (C++17) beta function
102-
// D comp_ellint_1/ f / l (C++17) complete elliptic integral (1st kind)
103-
// D comp_ellint_2/ f / l (C++17) complete elliptic integral (2nd kind)
104-
// D comp_ellint_3/ f / l (C++17) complete elliptic integral (3rd kind)
102+
// DS comp_ellint_1/ f / l (C++17) complete elliptic integral (1st kind)
103+
// DS comp_ellint_2/ f / l (C++17) complete elliptic integral (2nd kind)
104+
// DS comp_ellint_3/ f / l (C++17) complete elliptic integral (3rd kind)
105105
// D cyl_bessel_i/ f / l (C++17) modified cylindrical Bessel (regular)
106106
// D cyl_bessel_j/ f / l (C++17) cylindrical Bessel functions (1st kind)
107107
// D cyl_bessel_k/ f / l (C++17) modified cylindrical Bessel (irregular)
@@ -313,6 +313,25 @@ template<typename T> T f_beta(T x){ return std::beta(x,(T)2.0); } // x in (0, +i
313313
inline float f_betaf(float x){ return std::beta(x, 2.0f); }
314314
inline long double f_betal(long double x){ return std::beta(x, 2.0L); }
315315

316+
#if __cplusplus >= 201703L && (defined(__cpp_lib_math_special_funcs) || defined(__STDCPP_MATH_SPEC_FUNCS__))
317+
//------------------------ Elliptic integrals -----------------------------
318+
//
319+
// Domain: k in (-1, 1)
320+
template<typename T> T f_comp_ellint_1(T k) { return std::comp_ellint_1(k); }
321+
float f_comp_ellint_1f(float k) { return std::comp_ellint_1(k); }
322+
long double f_comp_ellint_1l(long double k) { return std::comp_ellint_1(k); }
323+
324+
// Domain: k in (-1, 1)
325+
template<typename T> T f_comp_ellint_2(T k) { return std::comp_ellint_2(k); }
326+
float f_comp_ellint_2f(float k) { return std::comp_ellint_2(k); }
327+
long double f_comp_ellint_2l(long double k) { return std::comp_ellint_2(k); }
328+
329+
// Domain: k in (-1, 1). Fixed nu = 0.5 for testing.
330+
template<typename T> T f_comp_ellint_3(T k) { return std::comp_ellint_3(k, (T)0.5); }
331+
float f_comp_ellint_3f(float k) { return std::comp_ellint_3(k, 0.5f); }
332+
long double f_comp_ellint_3l(long double k) { return std::comp_ellint_3(k, 0.5L); }
333+
#endif
334+
316335
int main() {
317336
// Absolute value
318337
CHECK(abs);
@@ -376,5 +395,12 @@ int main() {
376395
CHECK_ALL(erf);
377396
CHECK_ALL_RANGE(beta, {0.1, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0});
378397

398+
#if __cplusplus >= 201703L && (defined(__cpp_lib_math_special_funcs) || defined(__STDCPP_MATH_SPEC_FUNCS__))
399+
// Elliptic Integrals
400+
CHECK_ALL_RANGE(comp_ellint_1, {-0.9, -0.6, -0.3, 0.0, 0.3, 0.6, 0.9});
401+
CHECK_ALL_RANGE(comp_ellint_2, {-0.9, -0.6, -0.3, 0.0, 0.3, 0.6, 0.9});
402+
CHECK_ALL_RANGE(comp_ellint_3, {-0.9, -0.6, -0.3, 0.0, 0.3, 0.6, 0.9});
403+
#endif
404+
379405
return 0;
380406
}

0 commit comments

Comments
 (0)