13#ifndef dealii_differentiation_ad_ad_number_traits_h
14#define dealii_differentiation_ad_ad_number_traits_h
23#include <boost/type_traits.hpp>
45 template <
typename ScalarType,
64 template <
typename ADNumberType,
typename T =
void>
104 template <
typename ScalarType,
139 template <
typename ADNumberType,
typename T =
void>
172 template <
typename ADNumberType,
typename T =
void>
188 template <
typename ADNumberTrait,
typename T =
void>
206 template <
typename T>
214 template <
typename Number>
226 template <
typename NumberType>
236 template <
typename NumberType,
typename =
void>
246 template <
typename NumberType,
typename =
void>
256 template <
typename NumberType,
typename =
void>
266 template <
typename NumberType,
typename =
void>
285 template <
typename ADNumberTrait,
typename>
286 struct HasRequiredADInfo : std::false_type
301 template <
typename ADNumberTrait>
302 struct HasRequiredADInfo<
304 decltype((void)ADNumberTrait::type_code,
305 (void)ADNumberTrait::is_taped,
306 (void)std::declval<typename ADNumberTrait::real_type>(),
307 (void)std::declval<typename ADNumberTrait::derivative_type>(),
309 : std::conditional_t<
310 std::is_floating_point_v<typename ADNumberTrait::real_type>,
320 template <
typename ScalarType>
321 struct ADNumberInfoFromEnum<
324 std::enable_if_t<std::is_floating_point_v<ScalarType>>>
326 static const bool is_taped =
false;
327 using real_type = ScalarType;
328 using derivative_type = ScalarType;
329 static const unsigned int n_supported_derivative_levels = 0;
338 template <
typename ScalarType>
339 struct Marking<ScalarType,
340 std::enable_if_t<std::is_floating_point_v<ScalarType>>>
349 template <
typename ADNumberType>
351 independent_variable(
const ScalarType &in,
362 template <
typename ADNumberType>
364 dependent_variable(ADNumberType &,
const ScalarType &)
369 "Floating point numbers cannot be marked as dependent variables."));
377 template <
typename ADNumberType>
378 struct Marking<ADNumberType,
379 std::enable_if_t<boost::is_complex<ADNumberType>::value>>
384 template <
typename ScalarType>
386 independent_variable(
const ScalarType &in,
394 "Marking for complex numbers has not yet been implemented."));
401 template <
typename ScalarType>
403 dependent_variable(ADNumberType &,
const ScalarType &)
408 "Marking for complex numbers has not yet been implemented."));
415 template <
typename NumberType,
typename>
416 struct is_taped_ad_number : std::false_type
420 template <
typename NumberType,
typename>
421 struct is_tapeless_ad_number : std::false_type
425 template <
typename NumberType,
typename>
426 struct is_real_valued_ad_number : std::false_type
430 template <
typename NumberType,
typename>
431 struct is_complex_valued_ad_number : std::false_type
440 template <
typename NumberType>
442 : internal::HasRequiredADInfo<ADNumberTraits<std::decay_t<NumberType>>>
450 template <
typename NumberType>
451 struct is_taped_ad_number<
453 std::enable_if_t<ADNumberTraits<std::decay_t<NumberType>>::is_taped>>
462 template <
typename NumberType>
463 struct is_tapeless_ad_number<
465 std::enable_if_t<ADNumberTraits<std::decay_t<NumberType>>::is_tapeless>>
475 template <
typename NumberType>
476 struct is_real_valued_ad_number<
479 ADNumberTraits<std::decay_t<NumberType>>::is_real_valued>>
489 template <
typename NumberType>
490 struct is_complex_valued_ad_number<
493 ADNumberTraits<std::decay_t<NumberType>>::is_complex_valued>>
504 template <
typename Number>
505 struct RemoveComplexWrapper
516 template <
typename Number>
517 struct RemoveComplexWrapper<
std::complex<Number>>
519 using type =
typename RemoveComplexWrapper<Number>::type;
528 template <
typename NumberType>
529 struct ExtractData<NumberType,
530 std::enable_if_t<std::is_floating_point_v<NumberType>>>
535 static const NumberType &
536 value(
const NumberType &x)
546 n_directional_derivatives(
const NumberType &)
556 directional_derivative(
const NumberType &,
const unsigned int)
568 template <
typename ADNumberType>
569 struct ExtractData<
std::complex<ADNumberType>>
572 "Expected an auto-differentiable number.");
578 static std::complex<typename ADNumberTraits<ADNumberType>::scalar_type>
579 value(
const std::complex<ADNumberType> &x)
582 typename ADNumberTraits<ADNumberType>::scalar_type>(
583 ExtractData<ADNumberType>::value(x.real()),
584 ExtractData<ADNumberType>::value(x.imag()));
592 n_directional_derivatives(
const std::complex<ADNumberType> &x)
594 return ExtractData<ADNumberType>::n_directional_derivatives(x.real());
602 typename ADNumberTraits<ADNumberType>::derivative_type>
603 directional_derivative(
const std::complex<ADNumberType> &x,
604 const unsigned int direction)
607 typename ADNumberTraits<ADNumberType>::derivative_type>(
608 ExtractData<ADNumberType>::directional_derivative(x.real(),
610 ExtractData<ADNumberType>::directional_derivative(x.imag(),
616 template <
typename T>
622 template <
typename F>
624 value(
const F &f, std::enable_if_t<!is_ad_number<F>::value> * =
nullptr)
629 return ::internal::NumberType<T>::value(f);
638 template <
typename F>
641 std::enable_if_t<is_ad_number<F>::value &&
642 std::is_floating_point_v<T>> * =
nullptr)
647 return NumberType<T>::value(ExtractData<F>::value(f));
656 template <
typename F>
659 std::enable_if_t<is_ad_number<F>::value && is_ad_number<T>::value>
666 template <
typename T>
667 struct NumberType<
std::complex<T>>
672 template <
typename F>
674 value(
const F &f, std::enable_if_t<!is_ad_number<F>::value> * =
nullptr)
679 return ::internal::NumberType<std::complex<T>>::value(f);
687 template <
typename F>
688 static std::complex<T>
690 std::enable_if_t<is_ad_number<F>::value &&
691 std::is_floating_point_v<T>> * =
nullptr)
696 return std::complex<T>(
697 NumberType<T>::value(ExtractData<F>::value(f)));
700 template <
typename F>
701 static std::complex<T>
702 value(
const std::complex<F> &f)
706 return std::complex<T>(NumberType<T>::value(f.real()),
707 NumberType<T>::value(f.imag()));
735 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
740 std::is_floating_point_v<ScalarType> ||
741 (boost::is_complex<ScalarType>::value &&
742 std::is_floating_point_v<
743 typename internal::RemoveComplexWrapper<ScalarType>::type>)>>
748 static constexpr enum NumberTypes type_code = ADNumberTypeCode;
762 static const bool is_taped;
769 static const bool is_tapeless;
776 static const bool is_real_valued;
783 static const bool is_complex_valued;
790 static const unsigned int n_supported_derivative_levels;
798 static constexpr bool is_taped = internal::ADNumberInfoFromEnum<
799 typename internal::RemoveComplexWrapper<ScalarType>::type,
800 ADNumberTypeCode>::is_taped;
807 static constexpr bool is_tapeless =
808 !(NumberTraits<ScalarType, ADNumberTypeCode>::is_taped);
815 static constexpr bool is_real_valued =
816 (!boost::is_complex<ScalarType>::value);
823 static constexpr bool is_complex_valued =
824 !(NumberTraits<ScalarType, ADNumberTypeCode>::is_real_valued);
831 static constexpr unsigned int n_supported_derivative_levels =
832 internal::ADNumberInfoFromEnum<
833 typename internal::RemoveComplexWrapper<ScalarType>::type,
834 ADNumberTypeCode>::n_supported_derivative_levels;
843 using scalar_type = ScalarType;
849 using real_type =
typename internal::ADNumberInfoFromEnum<
850 typename internal::RemoveComplexWrapper<ScalarType>::type,
851 ADNumberTypeCode>::real_type;
857 using complex_type = std::complex<real_type>;
864 typename std::conditional_t<is_real_valued, real_type, complex_type>;
869 using derivative_type = std::conditional_t<
871 typename internal::ADNumberInfoFromEnum<
872 typename internal::RemoveComplexWrapper<ScalarType>::type,
873 ADNumberTypeCode>::derivative_type,
874 std::complex<
typename internal::ADNumberInfoFromEnum<
875 typename internal::RemoveComplexWrapper<ScalarType>::type,
876 ADNumberTypeCode>::derivative_type>>;
882 static scalar_type get_scalar_value(
const ad_type &x)
891 internal::ExtractData<ad_type>::value(x));
898 static derivative_type get_directional_derivative(
899 const ad_type &x,
const unsigned int direction)
901 return internal::ExtractData<ad_type>::directional_derivative(
910 static unsigned int n_directional_derivatives(
const ad_type &x)
912 return internal::ExtractData<ad_type>::n_directional_derivatives(x);
916 static_assert((is_real_valued ==
true ?
917 std::is_same_v<ad_type, real_type> :
918 std::is_same_v<ad_type, complex_type>),
919 "Incorrect template type selected for ad_type");
921 static_assert((is_complex_valued ==
true ?
922 boost::is_complex<scalar_type>::value :
924 "Expected a complex float_type");
926 static_assert((is_complex_valued ==
true ?
927 boost::is_complex<ad_type>::value :
929 "Expected a complex ad_type");
934 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
935 const bool NumberTraits<
939 std::is_floating_point_v<ScalarType> ||
940 (boost::is_complex<ScalarType>::value &&
941 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
942 ScalarType>::type>)>>::is_taped =
943 internal::ADNumberInfoFromEnum<
944 typename internal::RemoveComplexWrapper<ScalarType>::type,
945 ADNumberTypeCode>::is_taped;
948 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
949 const bool NumberTraits<
953 std::is_floating_point_v<ScalarType> ||
954 (boost::is_complex<ScalarType>::value &&
955 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
956 ScalarType>::type>)>>::is_tapeless =
957 !(NumberTraits<ScalarType, ADNumberTypeCode>::is_taped);
960 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
961 const bool NumberTraits<
965 std::is_floating_point_v<ScalarType> ||
966 (boost::is_complex<ScalarType>::value &&
967 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
968 ScalarType>::type>)>>::is_real_valued =
969 (!boost::is_complex<ScalarType>::value);
972 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
973 const bool NumberTraits<
977 std::is_floating_point_v<ScalarType> ||
978 (boost::is_complex<ScalarType>::value &&
979 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
980 ScalarType>::type>)>>::is_complex_valued =
981 !(NumberTraits<ScalarType, ADNumberTypeCode>::is_real_valued);
984 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
985 const unsigned int NumberTraits<
989 std::is_floating_point_v<ScalarType> ||
990 (boost::is_complex<ScalarType>::value &&
991 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
992 ScalarType>::type>)>>::n_supported_derivative_levels =
993 internal::ADNumberInfoFromEnum<
994 typename internal::RemoveComplexWrapper<ScalarType>::type,
995 ADNumberTypeCode>::n_supported_derivative_levels;
999 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
1000 constexpr bool NumberTraits<
1004 std::is_floating_point_v<ScalarType> ||
1005 (boost::is_complex<ScalarType>::value &&
1006 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1007 ScalarType>::type>)>>::is_taped;
1010 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
1011 constexpr bool NumberTraits<
1015 std::is_floating_point_v<ScalarType> ||
1016 (boost::is_complex<ScalarType>::value &&
1017 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1018 ScalarType>::type>)>>::is_tapeless;
1021 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
1022 constexpr bool NumberTraits<
1026 std::is_floating_point_v<ScalarType> ||
1027 (boost::is_complex<ScalarType>::value &&
1028 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1029 ScalarType>::type>)>>::is_real_valued;
1032 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
1033 constexpr bool NumberTraits<
1037 std::is_floating_point_v<ScalarType> ||
1038 (boost::is_complex<ScalarType>::value &&
1039 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1040 ScalarType>::type>)>>::is_complex_valued;
1043 template <
typename ScalarType, enum NumberTypes ADNumberTypeCode>
1044 constexpr unsigned int NumberTraits<
1048 std::is_floating_point_v<ScalarType> ||
1049 (boost::is_complex<ScalarType>::value &&
1050 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1051 ScalarType>::type>)>>::n_supported_derivative_levels;
1066 template <
typename ScalarType>
1067 struct NumberTraits<
1071 std::is_floating_point_v<ScalarType> ||
1072 (boost::is_complex<ScalarType>::value &&
1073 std::is_floating_point_v<
1074 typename internal::RemoveComplexWrapper<ScalarType>::type>)>>
1093 static const bool is_taped;
1100 static const bool is_tapeless;
1107 static const bool is_real_valued;
1114 static const bool is_complex_valued;
1121 static const unsigned int n_supported_derivative_levels;
1129 static constexpr bool is_taped =
false;
1136 static constexpr bool is_tapeless =
false;
1143 static constexpr bool is_real_valued =
1144 (!boost::is_complex<ScalarType>::value);
1151 static constexpr bool is_complex_valued = !is_real_valued;
1158 static constexpr unsigned int n_supported_derivative_levels = 0;
1167 using scalar_type = ScalarType;
1174 typename ::numbers::NumberTraits<scalar_type>::real_type;
1180 using complex_type = std::complex<real_type>;
1186 using ad_type = ScalarType;
1191 using derivative_type = ScalarType;
1198 get_scalar_value(
const ad_type &x)
1207 static derivative_type
1208 get_directional_derivative(
const ad_type &,
const unsigned int)
1213 "Floating point/arithmetic numbers have no directional derivatives."));
1214 return derivative_type();
1223 n_directional_derivatives(
const ad_type &)
1228 "Floating point/arithmetic numbers have no directional derivatives."));
1235 template <
typename ScalarType>
1236 const bool NumberTraits<
1240 std::is_floating_point_v<ScalarType> ||
1241 (boost::is_complex<ScalarType>::value &&
1242 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1243 ScalarType>::type>)>>::is_taped =
false;
1246 template <
typename ScalarType>
1247 const bool NumberTraits<
1251 std::is_floating_point_v<ScalarType> ||
1252 (boost::is_complex<ScalarType>::value &&
1253 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1254 ScalarType>::type>)>>::is_tapeless =
false;
1257 template <
typename ScalarType>
1258 const bool NumberTraits<
1262 std::is_floating_point_v<ScalarType> ||
1263 (boost::is_complex<ScalarType>::value &&
1264 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1265 ScalarType>::type>)>>::is_real_valued =
1266 (!boost::is_complex<ScalarType>::value);
1269 template <
typename ScalarType>
1270 const bool NumberTraits<
1274 std::is_floating_point_v<ScalarType> ||
1275 (boost::is_complex<ScalarType>::value &&
1276 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1277 ScalarType>::type>)>>::is_complex_valued =
1278 !(NumberTraits<ScalarType, NumberTypes::none>::is_real_valued);
1281 template <
typename ScalarType>
1282 const unsigned int NumberTraits<
1286 std::is_floating_point_v<ScalarType> ||
1287 (boost::is_complex<ScalarType>::value &&
1288 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1289 ScalarType>::type>)>>::n_supported_derivative_levels = 0;
1293 template <
typename ScalarType>
1294 constexpr bool NumberTraits<
1298 std::is_floating_point_v<ScalarType> ||
1299 (boost::is_complex<ScalarType>::value &&
1300 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1301 ScalarType>::type>)>>::is_taped;
1304 template <
typename ScalarType>
1305 constexpr bool NumberTraits<
1309 std::is_floating_point_v<ScalarType> ||
1310 (boost::is_complex<ScalarType>::value &&
1311 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1312 ScalarType>::type>)>>::is_tapeless;
1315 template <
typename ScalarType>
1316 constexpr bool NumberTraits<
1320 std::is_floating_point_v<ScalarType> ||
1321 (boost::is_complex<ScalarType>::value &&
1322 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1323 ScalarType>::type>)>>::is_real_valued;
1326 template <
typename ScalarType>
1327 constexpr bool NumberTraits<
1331 std::is_floating_point_v<ScalarType> ||
1332 (boost::is_complex<ScalarType>::value &&
1333 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1334 ScalarType>::type>)>>::is_complex_valued;
1337 template <
typename ScalarType>
1338 constexpr unsigned int NumberTraits<
1342 std::is_floating_point_v<ScalarType> ||
1343 (boost::is_complex<ScalarType>::value &&
1344 std::is_floating_point_v<
typename internal::RemoveComplexWrapper<
1345 ScalarType>::type>)>>::n_supported_derivative_levels;
1368 template <
typename ScalarType>
1369 struct ADNumberTraits<
1371 std::enable_if_t<std::is_floating_point_v<ScalarType>>>
1372 : NumberTraits<ScalarType, NumberTypes::none>
1378 template <
typename ComplexScalarType>
1379 struct ADNumberTraits<
1382 boost::is_complex<ComplexScalarType>::value &&
1383 std::is_floating_point_v<typename ComplexScalarType::value_type>>>
1384 : NumberTraits<ComplexScalarType, NumberTypes::none>
#define DEAL_II_NAMESPACE_OPEN
#define DEAL_II_NAMESPACE_CLOSE
#define Assert(cond, exc)
static ::ExceptionBase & ExcMessage(std::string arg1)
#define AssertThrow(cond, exc)
* * * RotationFunction< dim, Number >::RotationFunction Number(dim)
static constexpr const T & value(const T &t)