MiniTensor Version of the Day
Loading...
Searching...
No Matches
MiniTensor_TensorBase.h
Go to the documentation of this file.
1// @HEADER
2// *****************************************************************************
3// MiniTensor Package
4//
5// Copyright 2016 NTESS and the MiniTensor contributors.
6// SPDX-License-Identifier: BSD-3-Clause
7// *****************************************************************************
8// @HEADER
9
10#if !defined(MiniTensor_TensorBase_h)
11#define MiniTensor_TensorBase_h
12
13#include <algorithm>
14#include <cassert>
15#include <iostream>
16#include <vector>
17
18#include "MiniTensor_Storage.h"
19#include "MiniTensor_Scalar.h"
20
21namespace minitensor {
22
25
29enum class Filler {
30 ZEROS,
31 ONES,
33 RANDOM,
36 NANS
37};
38
47
51enum class Source {
52 ARRAY
53};
54
60template<typename T, typename ST>
62{
63public:
64
68 using value_type = T;
69
73 using storage_type = ST;
74
80
86 explicit
88 TensorBase(Index const dimension, Index const order);
89
97 TensorBase(Index const dimension, Index const order, Filler const value);
98
106 TensorBase(Index const dimension, Index const order, T const & s);
107
115 // TensorBase for Kokkos Data Types (we can 't use pointers with Kokkos::View)
116 template<class ArrayT>
119 Index const dimension,
120 Index const order,
121 ArrayT & data,
122 Index index1);
123
132 template<class ArrayT>
135 Index const dimension,
136 Index const order,
137 ArrayT & data,
138 Index index1,
139 Index index2);
140
150 template<class ArrayT>
153 Index const dimension,
154 Index const order,
155 ArrayT & data,
156 Index index1,
157 Index index2,
158 Index index3);
159
170 template<class ArrayT>
173 Index const dimension,
174 Index const order,
175 ArrayT & data,
176 Index index1,
177 Index index2,
178 Index index3,
179 Index index4);
180
192 template<class ArrayT>
195 Index const dimension,
196 Index const order,
197 ArrayT & data,
198 Index index1,
199 Index index2,
200 Index index3,
201 Index index4,
202 Index index5);
203
216 template<class ArrayT>
219 Index const dimension,
220 Index const order,
221 ArrayT & data,
222 Index index1,
223 Index index2,
224 Index index3,
225 Index index4,
226 Index index5,
227 Index index6);
228
229 //TensorBase for Shards and other data Types
237 TensorBase(Index const dimension, Index const order, T const * data_ptr);
244
252
258 T const &
259 operator[](Index const i) const;
260
266 T &
268
273 Index
275
281 void
282 fill(Filler const value);
283
289 void
290 fill(T const & s);
291
297 template<class ArrayT>
299 void
301 ArrayT & data,
302 Index index1);
303
310 template<class ArrayT>
312 void
314 ArrayT & data,
315 Index index1,
316 Index index2);
317
325 template<class ArrayT>
327 void
329 ArrayT & data,
330 Index index1,
331 Index index2,
332 Index index3);
333
342 template<class ArrayT>
344 void
346 ArrayT & data,
347 Index index1,
348 Index index2,
349 Index index3,
350 Index index4);
351
361 template<class ArrayT>
363 void
365 ArrayT & data,
366 Index index1,
367 Index index2,
368 Index index3,
369 Index index4,
370 Index index5);
371
382 template<class ArrayT>
384 void
386 ArrayT & data,
387 Index index1,
388 Index index2,
389 Index index3,
390 Index index4,
391 Index index5,
392 Index index6);
393
399 void fill(T const * data_ptr);
400
401
408 void
409 fill(T const * data_ptr, ComponentOrder const component_order);
410
415 template<typename S, typename SS>
419
424 template<typename S, typename SS>
428
433 template<typename S>
436 operator*=(S const & X);
437
442 template<typename S>
445 operator/=(S const & X);
446
451 void
453
454protected:
455
460 void
461 set_number_components(Index const number_components);
462
467 Index
468 get_dimension(Index const order) const;
469
475 void
476 set_dimension(Index const dimension, Index const order);
477
481 ST
483
487 Index
489};
490
494template<typename T, typename ST>
496T
497norm_f(TensorBase<T, ST> const & X);
498
502template<typename T, typename ST>
504T
505norm_f_square(TensorBase<T, ST> const & X);
506
510template<typename R, typename S, typename T, typename SR, typename SS,
511 typename ST>
513void
514add(
515 TensorBase<R, SR> const & A,
516 TensorBase<S, SS> const & B,
517 TensorBase<T, ST> & C);
518
522template<typename R, typename S, typename T, typename SR, typename SS,
523 typename ST>
525void
527 TensorBase<R, SR> const & A,
528 TensorBase<S, SS> const & B,
529 TensorBase<T, ST> & C);
530
534template<typename T, typename ST>
536void
537minus(TensorBase<T, ST> const & A, TensorBase<T, ST> & B);
538
542template<typename T, typename ST>
544bool
545equal(TensorBase<T, ST> const & A, TensorBase<T, ST> const & B);
546
550template<typename T, typename ST>
552bool
553not_equal(TensorBase<T, ST> const & A, TensorBase<T, ST> const & B);
554
558template<typename R, typename S, typename T, typename SR, typename ST>
560void
561scale(TensorBase<R, SR> const & A, S const & s, TensorBase<T, ST> & B);
562
566template<typename R, typename S, typename T, typename SR, typename ST>
568void
569divide(TensorBase<R, SR> const & A, S const & s, TensorBase<T, ST> & B);
570
574template<typename R, typename S, typename T, typename SR, typename ST>
576void
577split(TensorBase<R, SR> const & A, S const & s, TensorBase<T, ST> & B);
578
579} // namespace minitensor
580
581namespace minitensor
582{
583
584//
585// Default constructor.
586//
587template<typename T, typename ST>
590{
591 Index const
592 static_size = ST::static_size();
593
594 set_number_components(static_size);
595 fill(Filler::NANS);
596
597 return;
598}
599
600//
601// Construction that initializes to NaNs
602//
603template<typename T, typename ST>
605TensorBase<T, ST>::TensorBase(Index const dimension, Index const order)
606{
607 set_dimension(dimension, order);
608 fill(Filler::NANS);
609 return;
610}
611
612//
613// Create with specified value
614//
615template<typename T, typename ST>
618 Index const dimension,
619 Index const order,
620 Filler const value)
621{
622 set_dimension(dimension, order);
623 fill(value);
624 return;
625}
626
627//
628// Construction from a scalar
629//
630template<typename T, typename ST>
633 Index const dimension,
634 Index const order,
635 T const & s)
636{
637 set_dimension(dimension, order);
638 fill(s);
639 return;
640}
641
642//
643// Construction from array
644//Kokkos data Types:
645template<typename T, typename ST>
646template<class ArrayT>
649 Index const dimension,
650 Index const order,
651 ArrayT & data,
652 Index index1)
653{
654 set_dimension(dimension, order);
655 fill(data, index1);
656 return;
657}
658
659template<typename T, typename ST>
660template<class ArrayT>
663 Index const dimension,
664 Index const order,
665 ArrayT & data,
666 Index index1,
667 Index index2)
668{
669 set_dimension(dimension, order);
670 fill(data, index1, index2);
671 return;
672}
673
674template<typename T, typename ST>
675template<class ArrayT>
678 Index const dimension,
679 Index const order,
680 ArrayT & data,
681 Index index1,
682 Index index2,
683 Index index3)
684{
685 set_dimension(dimension, order);
686 fill(data, index1, index2, index3);
687 return;
688}
689
690template<typename T, typename ST>
691template<class ArrayT>
694 Index const dimension,
695 Index const order,
696 ArrayT & data,
697 Index index1,
698 Index index2,
699 Index index3,
700 Index index4)
701{
702 set_dimension(dimension, order);
703 fill(data, index1, index2, index3, index4);
704 return;
705}
706
707template<typename T, typename ST>
708template<class ArrayT>
711 Index const dimension,
712 Index const order,
713 ArrayT & data,
714 Index index1,
715 Index index2,
716 Index index3,
717 Index index4,
718 Index index5)
719{
720 set_dimension(dimension, order);
721 fill(data, index1, index2, index3, index4, index5);
722 return;
723}
724
725template<typename T, typename ST>
726template<class ArrayT>
729 Index const dimension,
730 Index const order,
731 ArrayT & data,
732 Index index1,
733 Index index2,
734 Index index3,
735 Index index4,
736 Index index5,
737 Index index6)
738{
739 set_dimension(dimension, order);
740 fill(data, index1, index2, index3, index4, index5, index6);
741 return;
742}
743
744template<typename T, typename ST>
747 Index const dimension,
748 Index const order,
749 T const * data_ptr)
750{
751 set_dimension(dimension, order);
752 fill(data_ptr);
753 return;
754}
755
756//
757// Copy constructor
758//
759template<typename T, typename ST>
762 dimension_(X.dimension_)
763{
764 Index const
765 number_components = X.get_number_components();
766
767 set_number_components(number_components);
768
769 for (Index i = 0; i < number_components; ++i) {
770 (*this)[i] = X[i];
771 }
772
773 return;
774}
775
776//
777// Copy assignment
778//
779template<typename T, typename ST>
783{
784 if (this == &X) return *this;
785
786 dimension_ = X.dimension_;
787
788 Index const
789 number_components = X.get_number_components();
790
791 set_number_components(number_components);
792
793 for (Index i = 0; i < number_components; ++i) {
794 (*this)[i] = X[i];
795 }
796
797 return *this;
798}
799
800//
801// Get dimension
802//
803template<typename T, typename ST>
805Index
807{
808 return dimension_;
809}
810
811//
812// Set dimension
813//
814template<typename T, typename ST>
816void
817TensorBase<T, ST>::set_dimension(Index const dimension, Index const order)
818{
819 dimension_ = dimension;
820
821 Index const
822 number_components = integer_power(dimension, order);
823
824 set_number_components(number_components);
825
826 return;
827}
828
829//
830// Linear access to components
831//
832template<typename T, typename ST>
834T const &
836{
837 return components_[i];
838}
839
840//
841// Linear access to components
842//
843template<typename T, typename ST>
845T &
847{
848 return components_[i];
849}
850
851//
852// Get total number of components
853//
854template<typename T, typename ST>
856Index
858{
859 return components_.size();
860}
861
862//
863// Allocate space for components
864//
865template<typename T, typename ST>
867void
869{
870 using S = typename Sacado::ScalarType<T>::type;
871
872 Index const
873 old_size = get_number_components();
874
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>();
880 }
881 }
882
883 components_.resize(number_components);
884
885 // Bound the growth loop by the size the storage reports after the resize,
886 // which static storage clamps to its capacity, rather than by the requested
887 // count, which the optimizer cannot bound. The two are equal whenever the
888 // resize was valid.
889 Index const
890 new_size = get_number_components();
891
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>();
897 }
898 }
899
900 return;
901}
902
903//
904// Fill components with value.
905//
906template<typename T, typename ST>
908void
910{
911 using S = typename Sacado::ScalarType<T>::type;
912
913 Index const
914 number_components = get_number_components();
915
916 switch (value) {
917
918 case Filler::ZEROS:
919 for (Index i = 0; i < number_components; ++i) {
920 auto & entry = (*this)[i];
921 fill_AD<T>(entry, S(0));
922 entry = S(0);
923 }
924 break;
925
926 case Filler::ONES:
927 for (Index i = 0; i < number_components; ++i) {
928 auto & entry = (*this)[i];
929 fill_AD<T>(entry, S(0));
930 entry = S(1);
931 }
932 break;
933
934 case Filler::SEQUENCE:
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);
939 }
940 break;
941
942 case Filler::NANS:
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>();
947 }
948 break;
949
950 case Filler::RANDOM:
951 KOKKOS_IF_ON_HOST((
952 for (Index i = 0; i < number_components; ++i) {
953 auto & entry = (*this)[i];
954 fill_AD<T>(entry, S(0));
955 entry = random<S>();
956 }
957 break;
958 ))
959 KOKKOS_IF_ON_DEVICE((
960 [[fallthrough]];
961 ))
962
964 KOKKOS_IF_ON_HOST((
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>();
969 }
970 break;
971 ))
972 KOKKOS_IF_ON_DEVICE((
973 [[fallthrough]];
974 ))
975
977 KOKKOS_IF_ON_HOST((
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>();
982 }
983 break;
984 ))
985 KOKKOS_IF_ON_DEVICE((
986 [[fallthrough]];
987 ))
988
989 default:
990 MT_ERROR_EXIT("Unknown or undefined (in execution space) specification of "
991 "value for filling components.");
992 break;
993 }
994
995 return;
996}
997
998//
999// Fill components from argument
1000//
1001template<typename T, typename ST>
1003void
1005{
1006 using S = typename Sacado::ScalarType<T>::type;
1007
1008 Index const
1009 number_components = get_number_components();
1010
1011 for (Index i = 0; i < number_components; ++i) {
1012 auto & entry = (*this)[i];
1013 fill_AD<T>(entry, S(0));
1014 entry = s;
1015 }
1016
1017 return;
1018}
1019
1020//
1021// Fill components from array defined by pointer.
1022//
1023template<typename T, typename ST>
1024template<class ArrayT>
1026void
1028 ArrayT & data,
1029 Index index1)
1030{
1031 assert(index1 == 0);
1032
1033 Index const
1034 number_components = get_number_components();
1035
1036 Index const
1037 rank = number_components / data.extent(0);
1038
1039 switch (rank) {
1040
1041 default:
1042 MT_ERROR_EXIT("Invalid rank.");
1043 break;
1044
1045 case 1:
1046 for (Index i = 0; i < number_components; ++i) {
1047 (*this)[i] = data(i);
1048 }
1049 break;
1050 }
1051
1052 return;
1053}
1054
1055template<typename T, typename ST>
1056template<class ArrayT>
1058void
1060 ArrayT & data,
1061 Index index1,
1062 Index index2)
1063{
1064 assert(index2 == 0);
1065
1066 Index const
1067 number_components = get_number_components();
1068
1069 Index
1070 rank = 0;
1071
1072 Index
1073 sub_dimension = number_components;
1074
1075 Index const
1076 dim = data.extent(1);
1077
1078 while (sub_dimension != 1) {
1079
1080 sub_dimension /= dim;
1081 ++rank;
1082
1083 assert(sub_dimension >= 1);
1084
1085 }
1086
1087 switch (rank) {
1088
1089 default:
1090 MT_ERROR_EXIT("Invalid rank.");
1091 break;
1092
1093 case 1:
1094 for (Index j = 0; j < number_components; ++j) {
1095 (*this)[j] = data(index1, j);
1096 }
1097 break;
1098
1099 case 2:
1100 for (Index i = 0; i < dim; ++i) {
1101 for (Index j = 0; j < dim; ++j) {
1102 (*this)[dim * i + j] = data(i, j);
1103 }
1104 }
1105 break;
1106 }
1107
1108 return;
1109}
1110
1111template<typename T, typename ST>
1112template<class ArrayT>
1114void
1116 ArrayT & data,
1117 Index index1,
1118 Index index2,
1119 Index index3)
1120{
1121 assert(index3 == 0);
1122
1123 Index const
1124 number_components = get_number_components();
1125
1126 Index const
1127 dim = data.extent(2);
1128
1129 Index
1130 rank = 0;
1131
1132 Index
1133 sub_dimension = number_components;
1134
1135 while (sub_dimension != 1) {
1136
1137 sub_dimension /= dim;
1138 ++rank;
1139
1140 assert(sub_dimension >= 1);
1141
1142 }
1143
1144 switch (rank) {
1145
1146 default:
1147 MT_ERROR_EXIT("Invalid rank.");
1148 break;
1149
1150 case 1:
1151 for (Index k = 0; k < number_components; ++k) {
1152 (*this)[k] = data(index1, index2, k);
1153 }
1154 break;
1155
1156 case 2:
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);
1160 }
1161 }
1162 break;
1163
1164 case 3:
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);
1169 }
1170 }
1171 }
1172 break;
1173 }
1174
1175 return;
1176}
1177
1178template<typename T, typename ST>
1179template<class ArrayT>
1181void
1183 ArrayT & data,
1184 Index index1,
1185 Index index2,
1186 Index index3,
1187 Index index4)
1188{
1189 assert(index4 == 0);
1190
1191 Index const
1192 number_components = get_number_components();
1193
1194 Index const
1195 dim = data.extent(2);
1196
1197 Index
1198 rank = 0;
1199
1200 Index
1201 sub_dimension = number_components;
1202
1203 while (sub_dimension != 1) {
1204
1205 sub_dimension /= dim;
1206 ++rank;
1207
1208 assert(sub_dimension >= 1);
1209
1210 }
1211
1212 switch (rank) {
1213
1214 default:
1215 MT_ERROR_EXIT("Invalid rank.");
1216 break;
1217
1218 case 1:
1219 for (Index l = 0; l < number_components; ++l) {
1220 (*this)[l] = data(index1, index2, index3, l);
1221 }
1222 break;
1223
1224 case 2:
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);
1228 }
1229 }
1230 break;
1231
1232 case 3:
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);
1237 }
1238 }
1239 }
1240 break;
1241
1242 case 4:
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] =
1248 data(i, j, k, l);
1249 }
1250 }
1251 }
1252 }
1253 break;
1254 }
1255
1256 return;
1257}
1258
1259template<typename T, typename ST>
1260template<class ArrayT>
1262void
1264 ArrayT & data,
1265 Index index1,
1266 Index index2,
1267 Index index3,
1268 Index index4,
1269 Index index5)
1270{
1271 assert(index5 == 0);
1272
1273 Index const
1274 number_components = get_number_components();
1275
1276 Index const
1277 dim = data.extent(2);
1278
1279 Index
1280 rank = 0;
1281
1282 Index
1283 sub_dimension = number_components;
1284
1285 while (sub_dimension != 1) {
1286
1287 sub_dimension /= dim;
1288 ++rank;
1289
1290 assert(sub_dimension >= 1);
1291
1292 }
1293
1294 switch (rank) {
1295
1296 default:
1297 MT_ERROR_EXIT("Invalid rank.");
1298 break;
1299
1300 case 1:
1301 for (Index m = 0; m < number_components; ++m) {
1302 (*this)[m] = data(index1, index2, index3, index4, m);
1303 }
1304 break;
1305
1306 case 2:
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);
1310 }
1311 }
1312 break;
1313
1314 case 3:
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);
1319 }
1320 }
1321 }
1322 break;
1323
1324 case 4:
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);
1331 }
1332 }
1333 }
1334 }
1335 break;
1336
1337 case 5:
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);
1345 }
1346 }
1347 }
1348 }
1349 }
1350 break;
1351 }
1352
1353 return;
1354}
1355
1356template<typename T, typename ST>
1357template<class ArrayT>
1359void
1361 ArrayT & data,
1362 Index index1,
1363 Index index2,
1364 Index index3,
1365 Index index4,
1366 Index index5,
1367 Index index6)
1368{
1369 assert(index6 == 0);
1370
1371 Index const
1372 number_components = get_number_components();
1373
1374 Index const
1375 dim = data.extent(2);
1376
1377 Index
1378 rank = 0;
1379
1380 Index
1381 sub_dimension = number_components;
1382
1383 while (sub_dimension != 1) {
1384
1385 sub_dimension /= dim;
1386 ++rank;
1387
1388 assert(sub_dimension >= 1);
1389
1390 }
1391
1392 switch (rank) {
1393
1394 default:
1395 MT_ERROR_EXIT("Invalid rank.");
1396 break;
1397
1398 case 1:
1399 for (Index n = 0; n < number_components; ++n) {
1400 (*this)[n] = data(index1, index2, index3, index4, index5, n);
1401 }
1402 break;
1403
1404 case 2:
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);
1408 }
1409 }
1410 break;
1411
1412 case 3:
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);
1418 }
1419 }
1420 }
1421 break;
1422
1423 case 4:
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);
1430 }
1431 }
1432 }
1433 }
1434 break;
1435
1436 case 5:
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);
1444 }
1445 }
1446 }
1447 }
1448 }
1449 break;
1450
1451 case 6:
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);
1460 }
1461 }
1462 }
1463 }
1464 }
1465 }
1466 break;
1467 }
1468
1469 return;
1470}
1471template<typename T, typename ST>
1473void
1474TensorBase<T, ST>::fill(T const * data_ptr)
1475{
1476 assert(data_ptr != NULL);
1477
1478 Index const
1479 number_components = get_number_components();
1480
1481 for (Index i = 0; i < number_components; ++i) {
1482 (*this)[i] = data_ptr[i];
1483 }
1484
1485 return;
1486}
1487
1488//
1489// Fill components from array defined by pointer.
1490//
1491template<typename T, typename ST>
1493void
1495 T const * data_ptr,
1496 ComponentOrder const component_order)
1497{
1498 assert(data_ptr != NULL);
1499
1501 self = (*this);
1502
1503 Index const
1504 number_components = self.get_number_components();
1505
1506 switch (number_components) {
1507
1508 default:
1509 self.fill(data_ptr);
1510 break;
1511
1512 case 9:
1513
1514 switch (component_order) {
1515
1517 self.fill(data_ptr);
1518 break;
1519
1521 // 0 1 2 3 4 5 6 7 8
1522 // XX YY ZZ XY YZ ZX YX ZY XZ
1523 // 0 4 8 1 5 6 3 7 2
1524 self[0] = data_ptr[0];
1525 self[4] = data_ptr[1];
1526 self[8] = data_ptr[2];
1527
1528 self[1] = data_ptr[3];
1529 self[5] = data_ptr[4];
1530 self[6] = data_ptr[5];
1531
1532 self[3] = data_ptr[6];
1533 self[7] = data_ptr[7];
1534 self[2] = data_ptr[8];
1535 break;
1536
1538 self[0] = data_ptr[0];
1539 self[4] = data_ptr[1];
1540 self[8] = data_ptr[2];
1541
1542 self[1] = data_ptr[3];
1543 self[5] = data_ptr[4];
1544 self[6] = data_ptr[5];
1545
1546 self[3] = data_ptr[3];
1547 self[7] = data_ptr[4];
1548 self[2] = data_ptr[5];
1549 break;
1550
1551 default:
1552 MT_ERROR_EXIT("Unknown component order.");
1553 break;
1554
1555 }
1556
1557 break;
1558 }
1559
1560 return;
1561}
1562
1563//
1564// Component increment
1565//
1566template<typename T, typename ST>
1567template<typename S, typename SS>
1571{
1572 Index const
1573 number_components = get_number_components();
1574
1575 assert(number_components == X.get_number_components());
1576
1577 for (Index i = 0; i < number_components; ++i) {
1578 (*this)[i] += X[i];
1579 }
1580
1581 return *this;
1582}
1583
1584//
1585// Component decrement
1586//
1587template<typename T, typename ST>
1588template<typename S, typename SS>
1592{
1593 Index const
1594 number_components = get_number_components();
1595
1596 assert(number_components == X.get_number_components());
1597
1598 for (Index i = 0; i < number_components; ++i) {
1599 (*this)[i] -= X[i];
1600 }
1601
1602 return *this;
1603}
1604
1605//
1606// Component scale
1607//
1608template<typename T, typename ST>
1609template<typename S>
1613{
1614 Index const
1615 number_components = get_number_components();
1616
1617 for (Index i = 0; i < number_components; ++i) {
1618 (*this)[i] *= X;
1619 }
1620 return *this;
1621}
1622
1623//
1624// Component divide
1625//
1626template<typename T, typename ST>
1627template<typename S>
1631{
1632 Index const
1633 number_components = get_number_components();
1634
1635 for (Index i = 0; i < number_components; ++i) {
1636 (*this)[i] /= X;
1637 }
1638 return *this;
1639}
1640
1641//
1642// Fill with zeros
1643//
1644template<typename T, typename ST>
1646void
1648{
1649 fill(Filler::ZEROS);
1650 return;
1651}
1652
1653//
1654// Square of Frobenius norm
1655//
1656template<typename T, typename ST>
1658T
1660{
1661 T
1662 s = 0.0;
1663
1664 for (Index i = 0; i < X.get_number_components(); ++i) {
1665 s += X[i] * X[i];
1666 }
1667
1668 return s;
1669}
1670
1671//
1672// Frobenius norm
1673//
1674template<typename T, typename ST>
1676T
1678{
1679 T const
1680 s = norm_f_square(X);
1681
1682 if (s > 0.0) return std::sqrt(s);
1683
1684 return 0.0;
1685}
1686
1687//
1688// Base addition
1689//
1690template<typename R, typename S, typename T, typename SR, typename SS,
1691 typename ST>
1693void
1695 TensorBase<R, SR> const & A,
1696 TensorBase<S, SS> const & B,
1698 )
1699{
1700 Index const
1701 number_components = A.get_number_components();
1702
1703 assert(B.get_number_components() == number_components);
1704 assert(C.get_number_components() == number_components);
1705
1706 for (Index i = 0; i < number_components; ++i) {
1707 C[i] = A[i] + B[i];
1708 }
1709
1710 return;
1711}
1712
1713//
1714// Base subtraction
1715//
1716template<typename R, typename S, typename T, typename SR, typename SS,
1717 typename ST>
1719void
1721 TensorBase<R, SR> const & A,
1722 TensorBase<S, SS> const & B,
1724{
1725 Index const
1726 number_components = A.get_number_components();
1727
1728 assert(B.get_number_components() == number_components);
1729 assert(C.get_number_components() == number_components);
1730
1731 for (Index i = 0; i < number_components; ++i) {
1732 C[i] = A[i] - B[i];
1733 }
1734
1735 return;
1736}
1737
1738//
1739// Base minus
1740//
1741template<typename T, typename ST>
1743void
1745{
1746 Index const
1747 number_components = A.get_number_components();
1748
1749 assert(B.get_number_components() == number_components);
1750
1751 for (Index i = 0; i < number_components; ++i) {
1752 B[i] = -A[i];
1753 }
1754
1755 return;
1756}
1757
1758//
1759// Base equality
1760//
1761template<typename T, typename ST>
1763bool
1765{
1766 Index const
1767 number_components = A.get_number_components();
1768
1769 assert(B.get_number_components() == number_components);
1770
1771 for (Index i = 0; i < number_components; ++i) {
1772 if (A[i] != B[i]) return false;
1773 }
1774
1775 return true;
1776}
1777
1778//
1779// Base not equality
1780//
1781template<typename T, typename ST>
1783bool
1785{
1786 return !(equal(A, B));
1787}
1788
1789//
1790// Base scaling
1791//
1792template<typename R, typename S, typename T, typename SR, typename ST>
1794void
1795scale(TensorBase<R, SR> const & A, S const & s, TensorBase<T, ST> & B)
1796{
1797 Index const
1798 number_components = A.get_number_components();
1799
1800 assert(B.get_number_components() == number_components);
1801
1802 for (Index i = 0; i < number_components; ++i) {
1803 B[i] = s * A[i];
1804 }
1805
1806 return;
1807}
1808
1809//
1810// Base division
1811//
1812template<typename R, typename S, typename T, typename SR, typename ST>
1814void
1815divide(TensorBase<R, SR> const & A, S const & s, TensorBase<T, ST> & B)
1816{
1817 Index const
1818 number_components = A.get_number_components();
1819
1820 assert(B.get_number_components() == number_components);
1821
1822 for (Index i = 0; i < number_components; ++i) {
1823 B[i] = A[i] / s;
1824 }
1825
1826 return;
1827}
1828
1829//
1830// Base split (scalar divided by tensor)
1831//
1832template<typename R, typename S, typename T, typename SR, typename ST>
1834void
1835split(TensorBase<R, SR> const & A, S const & s, TensorBase<T, ST> & B)
1836{
1837 Index const
1838 number_components = A.get_number_components();
1839
1840 assert(B.get_number_components() == number_components);
1841
1842 for (Index i = 0; i < number_components; ++i) {
1843 B[i] = s / A[i];
1844 }
1845
1846 return;
1847}
1848
1849} // namespace minitensor
1850namespace minitensor {
1851
1852// Placeholder for now.
1853
1855} // namespace minitensor
1856
1857#endif //MiniTensor_TensorBase_h
#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)