Cytnx v1.0.0
Loading...
Searching...
No Matches
Type.hpp
Go to the documentation of this file.
1#ifndef CYTNX_TYPE_H_
2#define CYTNX_TYPE_H_
3
4#include <complex>
5#include <cstdint>
6#include <string>
7#include <type_traits>
8#include <tuple>
9#include <array>
10#include <utility>
11#include <vector>
12#include <variant>
13
14#include "cytnx_error.hpp" // also brings in cuComplex.h
15
16#define MKL_Complex8 std::complex<float>
17#define MKL_Complex16 std::complex<double>
18
19#ifdef BACKEND_TORCH
20typedef int32_t blas_int;
21#else
22
23 #ifdef UNI_MKL
24 #include <mkl.h>
25typedef MKL_INT blas_int;
26 #else
27typedef int32_t blas_int;
28 #endif
29
30#endif
31
32// @cond
33namespace cytnx {
34
35 template <class T>
36 using vec3d = std::vector<std::vector<std::vector<T>>>;
37
38 template <class T>
39 using vec2d = std::vector<std::vector<T>>;
40
41 typedef double cytnx_double;
42 typedef float cytnx_float;
43 typedef uint64_t cytnx_uint64;
44 typedef uint32_t cytnx_uint32;
45 typedef uint16_t cytnx_uint16;
46 typedef int64_t cytnx_int64;
47 typedef int32_t cytnx_int32;
48 typedef int16_t cytnx_int16;
49 typedef std::size_t cytnx_size_t;
50 typedef std::complex<float> cytnx_complex64;
51 typedef std::complex<double> cytnx_complex128;
52 typedef bool cytnx_bool;
53
54 namespace internal {
55 template <class>
56 struct is_complex_impl : std::false_type {};
57
58 template <class T>
59 struct is_complex_impl<std::complex<T>> : std::true_type {};
60
61 template <typename>
62 struct is_complex_floating_point_impl : std::false_type {};
63
64 template <typename T>
65 struct is_complex_floating_point_impl<std::complex<T>> : std::is_floating_point<T> {};
66
67 template <std::size_t I, typename T, typename Tuple>
68 constexpr std::size_t index_in_tuple_helper() {
69 static_assert(I < std::tuple_size_v<Tuple>, "Type not found!");
70 if constexpr (std::is_same_v<T, std::tuple_element_t<I, Tuple>>) {
71 return I;
72 } else {
73 return index_in_tuple_helper<I + 1, T, Tuple>();
74 }
75 }
76
77 } // namespace internal
78
79 // helper metafunction to transform a variant into another variant via a
80 // transform template alias
81 template <typename V, template <typename> class Transform>
82 struct make_variant_from_transform;
83
84 template <template <typename> class Transform, typename... Args>
85 struct make_variant_from_transform<std::variant<Args...>, Transform> {
86 using type = std::variant<typename Transform<Args>::type...>;
87 };
88
89 // helper type alias for make_variant_from_transform
90 template <typename V, template <typename> class Transform>
91 using make_variant_from_transform_t = typename make_variant_from_transform<V, Transform>::type;
92
93 template <typename T>
94 using is_complex = internal::is_complex_impl<std::remove_cv_t<T>>;
95
96 template <typename T>
97 using is_complex_floating_point = internal::is_complex_floating_point_impl<std::remove_cv_t<T>>;
98
99 // is_complex_v checks if a data type is of type std::complex
100 // usage: is_complex_v<T> returns true or false for a data type T
101 template <typename T>
102 constexpr bool is_complex_v = is_complex<T>::value;
103
104 // is_complex_floating_point_v<T> is a template constant that is true if T is of type complex<U>
105 // where U is a floating point type, and false otherwise.
106 template <typename T>
107 constexpr bool is_complex_floating_point_v = is_complex_floating_point<T>::value;
108
109 // variant_index<T, Variant> returns the index of type T in the Variant, or compile error if not
110 // found
111 template <typename T, typename Variant>
112 struct variant_index;
113
114 template <typename T, typename... Types>
115 struct variant_index<T, std::variant<Types...>> {
116 static constexpr std::size_t value = std::variant_size_v<std::variant<Types...>>;
117 };
118
119 template <typename T, typename... Types>
120 struct variant_index<T, std::variant<T, Types...>> {
121 static constexpr std::size_t value = 0;
122 };
123
124 template <typename T, typename U, typename... Types>
125 struct variant_index<T, std::variant<U, Types...>> {
126 static constexpr std::size_t value = 1 + variant_index<T, std::variant<Types...>>::value;
127 };
128
129 // helper template variable
130 template <typename T, typename Variant>
131 static constexpr std::size_t variant_index_v = variant_index<T, Variant>::value;
132
133 namespace internal {
134 // type_size returns the sizeof(T) for the supported types. This is the same as
135 // sizeof(T), except that size_type<void> is 0.
136 template <typename T>
137 inline constexpr int type_size = sizeof(T);
138 template <>
139 inline constexpr int type_size<void> = 0;
140 } // namespace internal
141
142 // the list of supported types. The dtype() of an object is an index into this list.
143 // std::variant works better than std::tuple here since a variant is constrained to only
144 // hold each type once, and we have std::variant_alternative_t<n> to get the n'th type,
145 // as well as the variant_index_v helper to get the index of a given type
146 using Type_list =
147 std::variant<void, cytnx_complex128, cytnx_complex64, cytnx_double, cytnx_float, cytnx_int64,
148 cytnx_uint64, cytnx_int32, cytnx_uint32, cytnx_int16, cytnx_uint16, cytnx_bool>;
149
150 // For GPU storage, the types are slightly different because CUDA uses their own complex type
151#ifdef UNI_GPU
152 using Type_list_gpu =
153 std::variant<void, cuDoubleComplex, cuComplex, cytnx_double, cytnx_float, cytnx_int64,
154 cytnx_uint64, cytnx_int32, cytnx_uint32, cytnx_int16, cytnx_uint16, cytnx_bool>;
155#endif
156
157 // The number of supported types
158 constexpr int N_Type = std::variant_size_v<Type_list>;
159 constexpr int N_fType = 5;
160
161 // The friendly name of each type
162 template <typename T>
163 inline constexpr char* Type_names = nullptr;
164 template <>
165 inline constexpr const char* Type_names<void> = "Void";
166 template <>
167 inline constexpr const char* Type_names<cytnx_complex128> = "Complex Double (Complex Float64)";
168 template <>
169 inline constexpr const char* Type_names<cytnx_complex64> = "Complex Float (Complex Float32)";
170 template <>
171 inline constexpr const char* Type_names<cytnx_double> = "Double (Float64)";
172 template <>
173 inline constexpr const char* Type_names<cytnx_float> = "Float (Float32)";
174 template <>
175 inline constexpr const char* Type_names<cytnx_int64> = "Int64";
176 template <>
177 inline constexpr const char* Type_names<cytnx_uint64> = "Uint64";
178 template <>
179 inline constexpr const char* Type_names<cytnx_int32> = "Int32";
180 template <>
181 inline constexpr const char* Type_names<cytnx_uint32> = "Uint32";
182 template <>
183 inline constexpr const char* Type_names<cytnx_int16> = "Int16";
184 template <>
185 inline constexpr const char* Type_names<cytnx_uint16> = "Uint16";
186 template <>
187 inline constexpr const char* Type_names<cytnx_bool> = "Bool";
188
189 // The corresponding Python enumeration name
190 template <typename T>
191 inline constexpr char* Type_enum_name = nullptr;
192 template <>
193 inline constexpr const char* Type_enum_name<void> = "Void";
194 template <>
195 inline constexpr const char* Type_enum_name<cytnx_complex128> = "ComplexDouble";
196 template <>
197 inline constexpr const char* Type_enum_name<cytnx_complex64> = "ComplexFloat";
198 template <>
199 inline constexpr const char* Type_enum_name<cytnx_double> = "Double";
200 template <>
201 inline constexpr const char* Type_enum_name<cytnx_float> = "Float";
202 template <>
203 inline constexpr const char* Type_enum_name<cytnx_int64> = "Int64";
204 template <>
205 inline constexpr const char* Type_enum_name<cytnx_uint64> = "Uint64";
206 template <>
207 inline constexpr const char* Type_enum_name<cytnx_int32> = "Int32";
208 template <>
209 inline constexpr const char* Type_enum_name<cytnx_uint32> = "Uint32";
210 template <>
211 inline constexpr const char* Type_enum_name<cytnx_int16> = "Int16";
212 template <>
213 inline constexpr const char* Type_enum_name<cytnx_uint16> = "Uint16";
214 template <>
215 inline constexpr const char* Type_enum_name<cytnx_bool> = "Bool";
216
217 struct Type_struct {
218 const char* name; // char* is OK here, it is only ever initialized from a string literal
219 const char* enum_name;
220 bool is_unsigned;
221 bool is_complex;
222 bool is_float;
223 bool is_int;
224 unsigned int typeSize;
225 };
226
227 template <typename T>
228 struct Type_struct_t {
229 static constexpr unsigned int cy_typeid = variant_index_v<T, Type_list>;
230#ifdef UNI_GPU
231 static constexpr unsigned int cy_typeid_gpu = variant_index_v<T, Type_list_gpu>;
232#endif
233 static constexpr const char* name = Type_names<T>;
234 static constexpr const char* enum_name = Type_enum_name<T>;
235 static constexpr bool is_complex = is_complex_v<T>;
236 static constexpr bool is_unsigned = std::is_unsigned_v<T>;
237 static constexpr bool is_float = std::is_floating_point_v<T> || is_complex_floating_point_v<T>;
238 static constexpr bool is_int = std::is_integral_v<T> && !std::is_same_v<T, bool>;
239 static constexpr std::size_t typeSize = internal::type_size<T>;
240
241 static constexpr Type_struct construct() {
242 return {name, enum_name, is_unsigned, is_complex, is_float, is_int, typeSize};
243 }
244 };
245
246 namespace internal {
247 template <typename Variant, std::size_t... Indices>
248 constexpr auto make_type_array_helper(std::index_sequence<Indices...>) {
249 return std::array<Type_struct, sizeof...(Indices)>{
250 Type_struct_t<std::variant_alternative_t<Indices, Variant>>::construct()...};
251 }
252 template <typename Variant>
253 constexpr auto make_type_array() {
254 return make_type_array_helper<Variant>(
255 std::make_index_sequence<std::variant_size_v<Variant>>());
256 }
257 } // namespace internal
258
259 class Type_class {
260 private:
261 public:
262 // Typeinfos is a std::array<Type_struct> for each type in Type_list
263 static constexpr auto Typeinfos = internal::make_type_array<Type_list>();
264
265 template <typename T>
266 static constexpr unsigned int cy_typeid_v = variant_index_v<T, Type_list>;
267
268#ifdef UNI_GPU
269 template <typename T>
270 static constexpr unsigned int cy_typeid_gpu_v = variant_index_v<T, Type_list_gpu>;
271#endif
272
273 enum Type : unsigned int {
274 Void = cy_typeid_v<void>,
275 ComplexDouble = cy_typeid_v<cytnx_complex128>,
276 ComplexFloat = cy_typeid_v<cytnx_complex64>,
277 Double = cy_typeid_v<cytnx_double>,
278 Float = cy_typeid_v<cytnx_float>,
279 Int64 = cy_typeid_v<cytnx_int64>,
280 Uint64 = cy_typeid_v<cytnx_uint64>,
281 Int32 = cy_typeid_v<cytnx_int32>,
282 Uint32 = cy_typeid_v<cytnx_uint32>,
283 Int16 = cy_typeid_v<cytnx_int16>,
284 Uint16 = cy_typeid_v<cytnx_uint16>,
285 Bool = cy_typeid_v<cytnx_bool>
286 };
287
288 static constexpr void check_type(unsigned int type_id) {
289 cytnx_error_msg(type_id >= N_Type, "[ERROR] invalid type_id: %s", type_id);
290 }
291
292 // This could be constexpr returning constexpr char*, but there is lots of code that
293 // assumes that it returns a std::string and calls getname(n).c_str()
294 static std::string getname(unsigned int type_id) {
295 check_type(type_id);
296 return Typeinfos[type_id].name;
297 }
298 // This cannot be constexpr as we define it in a .cpp file,
299 // and typeid(T).name() is not constexpr
300 static unsigned int c_typename_to_id(const std::string& c_name);
301 static char const* enum_name(unsigned int type_id) {
302 check_type(type_id);
303 return Typeinfos[type_id].enum_name;
304 }
305 static constexpr unsigned int typeSize(unsigned int type_id) {
306 check_type(type_id);
307 return Typeinfos[type_id].typeSize;
308 }
309 static constexpr bool is_unsigned(unsigned int type_id) {
310 check_type(type_id);
311 return Typeinfos[type_id].is_unsigned;
312 }
313 static constexpr bool is_complex(unsigned int type_id) {
314 check_type(type_id);
315 return Typeinfos[type_id].is_complex;
316 }
317 static constexpr bool is_float(unsigned int type_id) {
318 check_type(type_id);
319 return Typeinfos[type_id].is_float;
320 }
321 static constexpr bool is_int(unsigned int type_id) {
322 check_type(type_id);
323 return Typeinfos[type_id].is_int;
324 }
325
326 template <class T>
327 static constexpr unsigned int cy_typeid(const T& rc) {
328 return cy_typeid_v<T>;
329 }
330
331 // Find a common type for typeL and typeR
332 static constexpr unsigned int type_promote(unsigned int typeL, unsigned int typeR) {
333 if (typeL < typeR) {
334 if (typeL == 0) return 0;
335
336 if (!is_unsigned(typeR) && is_unsigned(typeL)) {
337 return typeL - 1;
338 } else {
339 return typeL;
340 }
341 } else {
342 if (typeR == 0) return 0;
343 if (!is_unsigned(typeL) && is_unsigned(typeR)) {
344 return typeR - 1;
345 } else {
346 return typeR;
347 }
348 }
349 }
350
351 // type metafunction for type promotion
352 template <typename TL, typename TR>
353 using type_promote_t =
354 std::variant_alternative_t<Type_class::type_promote(variant_index_v<TL, Type_list>,
355 variant_index_v<TR, Type_list>),
356 Type_list>;
357
358 // Helper to promote two pointer types (note does _not_ return another pointer type)
359 template <typename TL, typename TR>
360 struct type_promote_from_pointer {
361 using type = void;
362 };
363
364 template <typename TL, typename TR>
365 struct type_promote_from_pointer<TL*, TR*> {
366 using type = type_promote_t<std::decay_t<TL>, std::decay_t<TR>>;
367 };
368
369 // helper typedef
370 template <typename TL, typename TR>
371 using type_promote_from_pointer_t = typename type_promote_from_pointer<TL, TR>::type;
372
373#ifdef UNI_GPU
374 // .. and we need a version where TL and TR are GPU device pointers
375 template <typename TL, typename TR>
376 using type_promote_gpu_t =
377 std::variant_alternative_t<Type_class::type_promote(variant_index_v<TL, Type_list_gpu>,
378 variant_index_v<TR, Type_list_gpu>),
379 Type_list_gpu>;
380
381 template <typename TL, typename TR>
382 struct type_promote_from_gpu_pointer {
383 using type = void;
384 };
385
386 template <typename TL, typename TR>
387 struct type_promote_from_gpu_pointer<TL*, TR*> {
388 using type = type_promote_gpu_t<std::decay_t<TL>, std::decay_t<TR>>;
389 };
390
391 // helper typedef
392 template <typename TL, typename TR>
393 using type_promote_from_gpu_pointer_t = typename type_promote_from_gpu_pointer<TL, TR>::type;
394#endif
395
396 }; // Type_class
398
426 constexpr Type_class Type;
427
428 extern int __blasINTsize__;
429
430 extern bool User_debug;
431
432} // namespace cytnx
433
434#endif // CYTNX_TYPE_H_
int __blasINTsize__
int32_t blas_int
Definition Type.hpp:27
bool User_debug
constexpr Type_class Type
data type
Definition Type.hpp:426
#define cytnx_error_msg(is_true, format,...)
Definition cytnx_error.hpp:27
Definition Accessor.hpp:12
@ U
Definition Symmetry.hpp:32