@@ -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
12181297CUDA_HOST_DEVICE inline ValueAndPushforward<float , float >
@@ -1489,6 +1568,14 @@ using std::beta_pullback;
14891568using 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+
14921579namespace class_functions {
14931580template <typename T, typename U>
14941581void constructor_pullback (ValueAndPushforward<T, U> rhs,
0 commit comments