10#if !defined(MiniTensor_TensorBase_h)
11#define MiniTensor_TensorBase_h
60template<
typename T,
typename ST>
116 template<
class ArrayT>
119 Index const dimension,
132 template<
class ArrayT>
135 Index const dimension,
150 template<
class ArrayT>
153 Index const dimension,
170 template<
class ArrayT>
173 Index const dimension,
192 template<
class ArrayT>
195 Index const dimension,
216 template<
class ArrayT>
219 Index const dimension,
297 template<
class ArrayT>
310 template<
class ArrayT>
325 template<
class ArrayT>
342 template<
class ArrayT>
361 template<
class ArrayT>
382 template<
class ArrayT>
415 template<
typename S,
typename SS>
424 template<
typename S,
typename SS>
494template<
typename T,
typename ST>
497norm_f(TensorBase<T, ST>
const & X);
502template<
typename T,
typename ST>
510template<
typename R,
typename S,
typename T,
typename SR,
typename SS,
515 TensorBase<R, SR>
const & A,
516 TensorBase<S, SS>
const & B,
517 TensorBase<T, ST> & C);
522template<
typename R,
typename S,
typename T,
typename SR,
typename SS,
527 TensorBase<R, SR>
const & A,
528 TensorBase<S, SS>
const & B,
529 TensorBase<T, ST> & C);
534template<
typename T,
typename ST>
537minus(TensorBase<T, ST>
const & A, TensorBase<T, ST> & B);
542template<
typename T,
typename ST>
545equal(TensorBase<T, ST>
const & A, TensorBase<T, ST>
const & B);
550template<
typename T,
typename ST>
553not_equal(TensorBase<T, ST>
const & A, TensorBase<T, ST>
const & B);
558template<
typename R,
typename S,
typename T,
typename SR,
typename ST>
561scale(TensorBase<R, SR>
const & A, S
const & s, TensorBase<T, ST> & B);
566template<
typename R,
typename S,
typename T,
typename SR,
typename ST>
569divide(TensorBase<R, SR>
const & A, S
const & s, TensorBase<T, ST> & B);
574template<
typename R,
typename S,
typename T,
typename SR,
typename ST>
577split(TensorBase<R, SR>
const & A, S
const & s, TensorBase<T, ST> & B);
587template<
typename T,
typename ST>
592 static_size = ST::static_size();
594 set_number_components(static_size);
603template<
typename T,
typename ST>
607 set_dimension(dimension, order);
615template<
typename T,
typename ST>
618 Index const dimension,
622 set_dimension(dimension, order);
630template<
typename T,
typename ST>
633 Index const dimension,
637 set_dimension(dimension, order);
645template<
typename T,
typename ST>
646template<
class ArrayT>
649 Index const dimension,
654 set_dimension(dimension, order);
659template<
typename T,
typename ST>
660template<
class ArrayT>
663 Index const dimension,
669 set_dimension(dimension, order);
670 fill(data, index1, index2);
674template<
typename T,
typename ST>
675template<
class ArrayT>
678 Index const dimension,
685 set_dimension(dimension, order);
686 fill(data, index1, index2, index3);
690template<
typename T,
typename ST>
691template<
class ArrayT>
694 Index const dimension,
702 set_dimension(dimension, order);
703 fill(data, index1, index2, index3, index4);
707template<
typename T,
typename ST>
708template<
class ArrayT>
711 Index const dimension,
720 set_dimension(dimension, order);
721 fill(data, index1, index2, index3, index4, index5);
725template<
typename T,
typename ST>
726template<
class ArrayT>
729 Index const dimension,
739 set_dimension(dimension, order);
740 fill(data, index1, index2, index3, index4, index5, index6);
744template<
typename T,
typename ST>
747 Index const dimension,
751 set_dimension(dimension, order);
759template<
typename T,
typename ST>
762 dimension_(X.dimension_)
769 for (
Index i = 0; i < number_components; ++i) {
779template<
typename T,
typename ST>
784 if (
this == &X)
return *
this;
791 set_number_components(number_components);
793 for (
Index i = 0; i < number_components; ++i) {
803template<
typename T,
typename ST>
814template<
typename T,
typename ST>
819 dimension_ = dimension;
824 set_number_components(number_components);
832template<
typename T,
typename ST>
837 return components_[i];
843template<
typename T,
typename ST>
848 return components_[i];
854template<
typename T,
typename ST>
859 return components_.size();
865template<
typename T,
typename ST>
870 using S =
typename Sacado::ScalarType<T>::type;
873 old_size = get_number_components();
875 if (number_components < old_size) {
876 for (
auto i = number_components; i < old_size; ++i) {
877 auto & entry = (*this)[i];
878 fill_AD<T>(entry, not_a_number<S>());
879 entry = not_a_number<T>();
883 components_.resize(number_components);
890 new_size = get_number_components();
892 if (new_size > old_size) {
893 for (
auto i = old_size; i < new_size; ++i) {
894 auto & entry = (*this)[i];
895 fill_AD<T>(entry, not_a_number<S>());
896 entry = not_a_number<T>();
906template<
typename T,
typename ST>
911 using S =
typename Sacado::ScalarType<T>::type;
914 number_components = get_number_components();
919 for (
Index i = 0; i < number_components; ++i) {
920 auto & entry = (*this)[i];
921 fill_AD<T>(entry, S(0));
927 for (
Index i = 0; i < number_components; ++i) {
928 auto & entry = (*this)[i];
929 fill_AD<T>(entry, S(0));
935 for (
Index i = 0; i < number_components; ++i) {
936 auto & entry = (*this)[i];
937 fill_AD<T>(entry, S(0));
938 entry =
static_cast<S
>(i);
943 for (
Index i = 0; i < number_components; ++i) {
944 auto & entry = (*this)[i];
945 fill_AD<T>(entry, not_a_number<S>());
946 entry = not_a_number<S>();
952 for (
Index i = 0; i < number_components; ++i) {
953 auto & entry = (*this)[i];
954 fill_AD<T>(entry, S(0));
959 KOKKOS_IF_ON_DEVICE((
965 for (
Index i = 0; i < number_components; ++i) {
966 auto & entry = (*this)[i];
967 fill_AD<T>(entry, S(0));
968 entry = random_uniform<S>();
972 KOKKOS_IF_ON_DEVICE((
978 for (
Index i = 0; i < number_components; ++i) {
979 auto & entry = (*this)[i];
980 fill_AD<T>(entry, S(0));
981 entry = random_normal<S>();
985 KOKKOS_IF_ON_DEVICE((
990 MT_ERROR_EXIT(
"Unknown or undefined (in execution space) specification of "
991 "value for filling components.");
1001template<
typename T,
typename ST>
1006 using S =
typename Sacado::ScalarType<T>::type;
1009 number_components = get_number_components();
1011 for (
Index i = 0; i < number_components; ++i) {
1012 auto & entry = (*this)[i];
1013 fill_AD<T>(entry, S(0));
1023template<
typename T,
typename ST>
1024template<
class ArrayT>
1031 assert(index1 == 0);
1034 number_components = get_number_components();
1037 rank = number_components / data.extent(0);
1046 for (
Index i = 0; i < number_components; ++i) {
1047 (*this)[i] = data(i);
1055template<
typename T,
typename ST>
1056template<
class ArrayT>
1064 assert(index2 == 0);
1067 number_components = get_number_components();
1073 sub_dimension = number_components;
1076 dim = data.extent(1);
1078 while (sub_dimension != 1) {
1080 sub_dimension /= dim;
1083 assert(sub_dimension >= 1);
1094 for (
Index j = 0; j < number_components; ++j) {
1095 (*this)[j] = data(index1, j);
1100 for (
Index i = 0; i < dim; ++i) {
1101 for (
Index j = 0; j < dim; ++j) {
1102 (*this)[dim * i + j] = data(i, j);
1111template<
typename T,
typename ST>
1112template<
class ArrayT>
1121 assert(index3 == 0);
1124 number_components = get_number_components();
1127 dim = data.extent(2);
1133 sub_dimension = number_components;
1135 while (sub_dimension != 1) {
1137 sub_dimension /= dim;
1140 assert(sub_dimension >= 1);
1151 for (
Index k = 0; k < number_components; ++k) {
1152 (*this)[k] = data(index1, index2, k);
1157 for (
Index j = 0; j < dim; ++j) {
1158 for (
Index k = 0; k < dim; ++k) {
1159 (*this)[dim * j + k] = data(index1, j, k);
1165 for (
Index i = 0; i < dim; ++i) {
1166 for (
Index j = 0; j < dim; ++j) {
1167 for (
Index k = 0; k < dim; ++k) {
1168 (*this)[dim * (dim * i + j) + k] = data(i, j, k);
1178template<
typename T,
typename ST>
1179template<
class ArrayT>
1189 assert(index4 == 0);
1192 number_components = get_number_components();
1195 dim = data.extent(2);
1201 sub_dimension = number_components;
1203 while (sub_dimension != 1) {
1205 sub_dimension /= dim;
1208 assert(sub_dimension >= 1);
1219 for (
Index l = 0; l < number_components; ++l) {
1220 (*this)[l] = data(index1, index2, index3, l);
1225 for (
Index k = 0; k < dim; ++k) {
1226 for (
Index l = 0; l < dim; ++l) {
1227 (*this)[dim * k + l] = data(index1, index2, k, l);
1233 for (
Index j = 0; j < dim; ++j) {
1234 for (
Index k = 0; k < dim; ++k) {
1235 for (
Index l = 0; l < dim; ++l) {
1236 (*this)[dim * (dim * j + k) + l] = data(index1, j, k, l);
1243 for (
Index i = 0; i < dim; ++i) {
1244 for (
Index j = 0; j < dim; ++j) {
1245 for (
Index k = 0; k < dim; ++k) {
1246 for (
Index l = 0; l < dim; ++l) {
1247 (*this)[dim * (dim * (dim * i + j) + k) + l] =
1259template<
typename T,
typename ST>
1260template<
class ArrayT>
1271 assert(index5 == 0);
1274 number_components = get_number_components();
1277 dim = data.extent(2);
1283 sub_dimension = number_components;
1285 while (sub_dimension != 1) {
1287 sub_dimension /= dim;
1290 assert(sub_dimension >= 1);
1301 for (
Index m = 0; m < number_components; ++m) {
1302 (*this)[m] = data(index1, index2, index3, index4, m);
1307 for (
Index l = 0; l < dim; ++l) {
1308 for (
Index m = 0; m < dim; ++m) {
1309 (*this)[dim * l + m] = data(index1, index2, index3, l, m);
1315 for (
Index k = 0; k < dim; ++k) {
1316 for (
Index l = 0; l < dim; ++l) {
1317 for (
Index m = 0; m < dim; ++m) {
1318 (*this)[dim * (dim * k + l) + m] = data(index1, index2, k, l, m);
1325 for (
Index j = 0; j < dim; ++j) {
1326 for (
Index k = 0; k < dim; ++k) {
1327 for (
Index l = 0; l < dim; ++l) {
1328 for (
Index m = 0; m < dim; ++m) {
1329 (*this)[dim * (dim * (dim * j + k) + l) + m] =
1330 data(index1, j, k, l, m);
1338 for (
Index i = 0; i < dim; ++i) {
1339 for (
Index j = 0; j < dim; ++j) {
1340 for (
Index k = 0; k < dim; ++k) {
1341 for (
Index l = 0; l < dim; ++l) {
1342 for (
Index m = 0; m < dim; ++m) {
1343 (*this)[dim * (dim * (dim * (dim * i + j) + k) + l) + m] =
1344 data(i, j, k, l, m);
1356template<
typename T,
typename ST>
1357template<
class ArrayT>
1369 assert(index6 == 0);
1372 number_components = get_number_components();
1375 dim = data.extent(2);
1381 sub_dimension = number_components;
1383 while (sub_dimension != 1) {
1385 sub_dimension /= dim;
1388 assert(sub_dimension >= 1);
1399 for (
Index n = 0; n < number_components; ++n) {
1400 (*this)[n] = data(index1, index2, index3, index4, index5, n);
1405 for (
Index m = 0; m < dim; ++m) {
1406 for (
Index n = 0; n < dim; ++n) {
1407 (*this)[dim * m + n] = data(index1, index2, index3, index4, m, n);
1413 for (
Index l = 0; l < dim; ++l) {
1414 for (
Index m = 0; m < dim; ++m) {
1415 for (
Index n = 0; n < dim; ++n) {
1416 (*this)[dim * (dim * l + m) + n] =
1417 data(index1, index2, index3, l, m, n);
1424 for (
Index k = 0; k < dim; ++k) {
1425 for (
Index l = 0; l < dim; ++l) {
1426 for (
Index m = 0; m < dim; ++m) {
1427 for (
Index n = 0; n < dim; ++n) {
1428 (*this)[dim * (dim * (dim * k + l) + m) + n] =
1429 data(index1, index2, k, l, m, n);
1437 for (
Index j = 0; j < dim; ++j) {
1438 for (
Index k = 0; k < dim; ++k) {
1439 for (
Index l = 0; l < dim; ++l) {
1440 for (
Index m = 0; m < dim; ++m) {
1441 for (
Index n = 0; n < dim; ++n) {
1442 (*this)[dim * (dim * (dim * (dim * j + k) + l) + m) + n] =
1443 data(index1, j, k, l, m, n);
1452 for (
Index i = 0; i < dim; ++i) {
1453 for (
Index j = 0; j < dim; ++j) {
1454 for (
Index k = 0; k < dim; ++k) {
1455 for (
Index l = 0; l < dim; ++l) {
1456 for (
Index m = 0; m < dim; ++m) {
1457 for (
Index n = 0; n < dim; ++n) {
1458 (*this)[dim * (dim * (dim * (dim * (dim *
1459 i + j) + k) + l) + m) + n] = data(i, j, k, l, m, n);
1471template<
typename T,
typename ST>
1476 assert(data_ptr != NULL);
1479 number_components = get_number_components();
1481 for (
Index i = 0; i < number_components; ++i) {
1482 (*this)[i] = data_ptr[i];
1491template<
typename T,
typename ST>
1498 assert(data_ptr != NULL);
1506 switch (number_components) {
1509 self.
fill(data_ptr);
1514 switch (component_order) {
1517 self.
fill(data_ptr);
1524 self[0] = data_ptr[0];
1525 self[4] = data_ptr[1];
1526 self[8] = data_ptr[2];
1528 self[1] = data_ptr[3];
1529 self[5] = data_ptr[4];
1530 self[6] = data_ptr[5];
1532 self[3] = data_ptr[6];
1533 self[7] = data_ptr[7];
1534 self[2] = data_ptr[8];
1538 self[0] = data_ptr[0];
1539 self[4] = data_ptr[1];
1540 self[8] = data_ptr[2];
1542 self[1] = data_ptr[3];
1543 self[5] = data_ptr[4];
1544 self[6] = data_ptr[5];
1546 self[3] = data_ptr[3];
1547 self[7] = data_ptr[4];
1548 self[2] = data_ptr[5];
1566template<
typename T,
typename ST>
1567template<
typename S,
typename SS>
1573 number_components = get_number_components();
1577 for (
Index i = 0; i < number_components; ++i) {
1587template<
typename T,
typename ST>
1588template<
typename S,
typename SS>
1594 number_components = get_number_components();
1598 for (
Index i = 0; i < number_components; ++i) {
1608template<
typename T,
typename ST>
1615 number_components = get_number_components();
1617 for (
Index i = 0; i < number_components; ++i) {
1626template<
typename T,
typename ST>
1633 number_components = get_number_components();
1635 for (
Index i = 0; i < number_components; ++i) {
1644template<
typename T,
typename ST>
1656template<
typename T,
typename ST>
1674template<
typename T,
typename ST>
1682 if (s > 0.0)
return std::sqrt(s);
1690template<
typename R,
typename S,
typename T,
typename SR,
typename SS,
1706 for (
Index i = 0; i < number_components; ++i) {
1716template<
typename R,
typename S,
typename T,
typename SR,
typename SS,
1731 for (
Index i = 0; i < number_components; ++i) {
1741template<
typename T,
typename ST>
1751 for (
Index i = 0; i < number_components; ++i) {
1761template<
typename T,
typename ST>
1771 for (
Index i = 0; i < number_components; ++i) {
1772 if (A[i] != B[i])
return false;
1781template<
typename T,
typename ST>
1786 return !(
equal(A, B));
1792template<
typename R,
typename S,
typename T,
typename SR,
typename ST>
1802 for (
Index i = 0; i < number_components; ++i) {
1812template<
typename R,
typename S,
typename T,
typename SR,
typename ST>
1822 for (
Index i = 0; i < number_components; ++i) {
1832template<
typename R,
typename S,
typename T,
typename SR,
typename ST>
1842 for (
Index i = 0; i < number_components; ++i) {
#define KOKKOS_INLINE_FUNCTION
#define MT_ERROR_EXIT(...)
KOKKOS_INLINE_FUNCTION void clear()
KOKKOS_INLINE_FUNCTION TensorBase< T, ST > & operator*=(S const &X)
KOKKOS_INLINE_FUNCTION void minus(TensorBase< T, ST > const &A, TensorBase< T, ST > &B)
KOKKOS_INLINE_FUNCTION void fill(ArrayT &data, Index index1, Index index2, Index index3, Index index4, Index index5, Index index6)
KOKKOS_INLINE_FUNCTION void fill(ArrayT &data, Index index1, Index index2, Index index3, Index index4)
KOKKOS_INLINE_FUNCTION void set_number_components(Index const number_components)
KOKKOS_INLINE_FUNCTION void fill(ArrayT &data, Index index1, Index index2, Index index3, Index index4, Index index5)
KOKKOS_INLINE_FUNCTION void add(TensorBase< R, SR > const &A, TensorBase< S, SS > const &B, TensorBase< T, ST > &C)
KOKKOS_INLINE_FUNCTION T const & operator[](Index const i) const
KOKKOS_INLINE_FUNCTION void fill(T const *data_ptr, ComponentOrder const component_order)
KOKKOS_INLINE_FUNCTION TensorBase(Index const dimension, Index const order, ArrayT &data, Index index1, Index index2, Index index3, Index index4, Index index5)
KOKKOS_INLINE_FUNCTION TensorBase(Index const dimension, Index const order, ArrayT &data, Index index1, Index index2)
KOKKOS_INLINE_FUNCTION T norm_f(TensorBase< T, ST > const &X)
KOKKOS_INLINE_FUNCTION TensorBase(Index const dimension, Index const order, ArrayT &data, Index index1, Index index2, Index index3, Index index4)
KOKKOS_INLINE_FUNCTION bool equal(TensorBase< T, ST > const &A, TensorBase< T, ST > const &B)
KOKKOS_INLINE_FUNCTION void fill(T const &s)
KOKKOS_INLINE_FUNCTION TensorBase< T, ST > & operator-=(TensorBase< S, SS > const &X)
KOKKOS_INLINE_FUNCTION void fill(ArrayT &data, Index index1, Index index2)
KOKKOS_INLINE_FUNCTION void fill(ArrayT &data, Index index1, Index index2, Index index3)
KOKKOS_INLINE_FUNCTION Index get_dimension(Index const order) const
KOKKOS_INLINE_FUNCTION TensorBase(Index const dimension, Index const order, ArrayT &data, Index index1)
KOKKOS_INLINE_FUNCTION void fill(T const *data_ptr)
KOKKOS_INLINE_FUNCTION TensorBase(Index const dimension, Index const order, ArrayT &data, Index index1, Index index2, Index index3)
KOKKOS_INLINE_FUNCTION void fill(Filler const value)
KOKKOS_INLINE_FUNCTION TensorBase(Index const dimension, Index const order)
KOKKOS_INLINE_FUNCTION T norm_f_square(TensorBase< T, ST > const &X)
KOKKOS_INLINE_FUNCTION TensorBase()
KOKKOS_INLINE_FUNCTION void scale(TensorBase< R, SR > const &A, S const &s, TensorBase< T, ST > &B)
KOKKOS_INLINE_FUNCTION TensorBase(Index const dimension, Index const order, T const *data_ptr)
KOKKOS_INLINE_FUNCTION TensorBase(Index const dimension, Index const order, T const &s)
KOKKOS_INLINE_FUNCTION TensorBase< T, ST > & operator/=(S const &X)
KOKKOS_INLINE_FUNCTION TensorBase< T, ST > & operator=(TensorBase< T, ST > const &X)
KOKKOS_INLINE_FUNCTION T & operator[](Index const i)
KOKKOS_INLINE_FUNCTION void split(TensorBase< R, SR > const &A, S const &s, TensorBase< T, ST > &B)
KOKKOS_INLINE_FUNCTION TensorBase(TensorBase< T, ST > const &X)
KOKKOS_INLINE_FUNCTION TensorBase< T, ST > & operator+=(TensorBase< S, SS > const &X)
KOKKOS_INLINE_FUNCTION TensorBase(Index const dimension, Index const order, Filler const value)
KOKKOS_INLINE_FUNCTION bool not_equal(TensorBase< T, ST > const &A, TensorBase< T, ST > const &B)
KOKKOS_INLINE_FUNCTION Index get_number_components() const
KOKKOS_INLINE_FUNCTION void set_dimension(Index const dimension, Index const order)
KOKKOS_INLINE_FUNCTION TensorBase(Index const dimension, Index const order, ArrayT &data, Index index1, Index index2, Index index3, Index index4, Index index5, Index index6)
KOKKOS_INLINE_FUNCTION void divide(TensorBase< R, SR > const &A, S const &s, TensorBase< T, ST > &B)
KOKKOS_INLINE_FUNCTION void fill(ArrayT &data, Index index1)
KOKKOS_INLINE_FUNCTION void subtract(TensorBase< R, SR > const &A, TensorBase< S, SS > const &B, TensorBase< T, ST > &C)
uint32_t Index
Indexing type.
KOKKOS_INLINE_FUNCTION T integer_power(T const &X, Index const exponent)