27namespace nda::tensor::device {
30 thread_local bool synchronize =
true;
31 void set_synchronization(
bool do_sync)
noexcept { synchronize = do_sync; }
32 bool get_synchronization() noexcept {
return synchronize; }
35 cutensorHandle_t &get_handle() {
36 struct handle_storage_t {
37 handle_storage_t() { cutensorCreate(&handle); }
38 ~handle_storage_t() { cutensorDestroy(handle); }
39 cutensorHandle_t handle = {};
41 static auto sto = handle_storage_t{};
49 void cutensor_error_check(cutensorStatus_t status, std::string_view func) {
50 if (status != CUTENSOR_STATUS_SUCCESS) {
51 NDA_RUNTIME_ERROR <<
"cuTENSOR runtime error in function " << func <<
"\n"
52 <<
" cutensorStatus_t: " << status <<
"\n"
53 <<
" cutensorGetErrorString: " << cutensorGetErrorString(status) <<
"\n";
58 template <
typename T,
typename U = std::remove_const_t<T>>
59 constexpr auto cuda_data_type() {
60 if constexpr (std::is_same_v<U, float>) {
61 return CUTENSOR_R_32F;
62 }
else if constexpr (std::is_same_v<U, double>) {
63 return CUTENSOR_R_64F;
64 }
else if constexpr (std::is_same_v<U, std::complex<float>>) {
65 return CUTENSOR_C_32F;
66 }
else if constexpr (std::is_same_v<U, std::complex<double>>) {
67 return CUTENSOR_C_64F;
72 template <
typename T,
typename U = std::remove_const_t<T>>
73 constexpr auto cutensor_compute_type() {
74 if constexpr (AnyOf<U, float, std::complex<float>>) {
75 return CUTENSOR_COMPUTE_DESC_32F;
76 }
else if constexpr (AnyOf<U, double, std::complex<double>>) {
77 return CUTENSOR_COMPUTE_DESC_64F;
84 std::uint32_t find_alignment(T *p) {
85 auto const x =
reinterpret_cast<std::uintptr_t
>(p);
86 return std::uint32_t{1} << std::min(std::countr_zero(x), 8);
90 auto to_modes(std::string_view idx) {
return std::vector<std::int32_t>(idx.begin(), idx.end()); }
94 cutensorOperator_t to_cutensor_unary_op(
unary_op op) {
96 case unary_op::IDENTITY:
return CUTENSOR_OP_IDENTITY;
97 case unary_op::SQRT:
return CUTENSOR_OP_SQRT;
98 case unary_op::RELU:
return CUTENSOR_OP_RELU;
99 case unary_op::CONJ:
return CUTENSOR_OP_CONJ;
100 case unary_op::RCP:
return CUTENSOR_OP_RCP;
101 case unary_op::SIGMOID:
return CUTENSOR_OP_SIGMOID;
102 case unary_op::TANH:
return CUTENSOR_OP_TANH;
103 case unary_op::EXP:
return CUTENSOR_OP_EXP;
104 case unary_op::LOG:
return CUTENSOR_OP_LOG;
105 case unary_op::ABS:
return CUTENSOR_OP_ABS;
106 case unary_op::NEG:
return CUTENSOR_OP_NEG;
107 case unary_op::SIN:
return CUTENSOR_OP_SIN;
108 case unary_op::COS:
return CUTENSOR_OP_COS;
109 case unary_op::TAN:
return CUTENSOR_OP_TAN;
110 case unary_op::SINH:
return CUTENSOR_OP_SINH;
111 case unary_op::COSH:
return CUTENSOR_OP_COSH;
112 case unary_op::ASIN:
return CUTENSOR_OP_ASIN;
113 case unary_op::ACOS:
return CUTENSOR_OP_ACOS;
114 case unary_op::ATAN:
return CUTENSOR_OP_ATAN;
115 case unary_op::ASINH:
return CUTENSOR_OP_ASINH;
116 case unary_op::ACOSH:
return CUTENSOR_OP_ACOSH;
117 case unary_op::ATANH:
return CUTENSOR_OP_ATANH;
118 case unary_op::CEIL:
return CUTENSOR_OP_CEIL;
119 case unary_op::FLOOR:
return CUTENSOR_OP_FLOOR;
120 case unary_op::MISH:
return CUTENSOR_OP_MISH;
121 case unary_op::SWISH:
return CUTENSOR_OP_SWISH;
122 case unary_op::SOFT_PLUS:
return CUTENSOR_OP_SOFT_PLUS;
123 case unary_op::SOFT_SIGN:
return CUTENSOR_OP_SOFT_SIGN;
124 default: NDA_RUNTIME_ERROR <<
"nda::tensor::cutensor: unary_op has no cuTENSOR equivalent";
130 cutensorOperator_t to_cutensor_binary_op(
binary_op op) {
132 case binary_op::SUM:
return CUTENSOR_OP_ADD;
133 case binary_op::PROD:
return CUTENSOR_OP_MUL;
134 case binary_op::MAX:
return CUTENSOR_OP_MAX;
135 case binary_op::MIN:
return CUTENSOR_OP_MIN;
136 default: NDA_RUNTIME_ERROR <<
"nda::tensor::cutensor: binary_op has no cuTENSOR equivalent";
141 template <
typename T>
142 auto create_tensor_desc(tensor_view<T> tv) {
143 cutensorTensorDescriptor_t desc{};
144 auto status = cutensorCreateTensorDescriptor(get_handle(), &desc,
static_cast<std::uint32_t
>(tv.ndim), tv.extents, tv.strides,
145 cuda_data_type<T>(), find_alignment(tv.data));
146 cutensor_error_check(status,
"cutensorCreateTensorDescriptor");
151 void destroy_tensor_desc(cutensorTensorDescriptor_t desc) {
152 cutensor_error_check(cutensorDestroyTensorDescriptor(desc),
"cutensorDestroyTensorDescriptor");
156 cutensorPlanPreference_t create_plan_pref() {
157 cutensorPlanPreference_t pref{};
158 cutensor_error_check(cutensorCreatePlanPreference(get_handle(), &pref, CUTENSOR_ALGO_DEFAULT, CUTENSOR_JIT_MODE_NONE),
159 "cutensorCreatePlanPreference");
164 void destroy_plan_pref(cutensorPlanPreference_t pref) {
165 cutensor_error_check(cutensorDestroyPlanPreference(pref),
"cutensorDestroyPlanPreference");
169 cutensorPlan_t create_plan(cutensorOperationDescriptor_t op_desc, cutensorPlanPreference_t pref, std::uint64_t workspace_limit = 0) {
170 cutensorPlan_t plan{};
171 cutensor_error_check(cutensorCreatePlan(get_handle(), &plan, op_desc, pref, workspace_limit),
"cutensorCreatePlan");
176 void destroy_plan(cutensorPlan_t plan) { cutensor_error_check(cutensorDestroyPlan(plan),
"cutensorDestroyPlan"); }
179 void destroy_op_desc(cutensorOperationDescriptor_t op_desc) {
180 cutensor_error_check(cutensorDestroyOperationDescriptor(op_desc),
"cutensorDestroyOperationDescriptor");
184 std::uint64_t estimate_workspace(cutensorOperationDescriptor_t op_desc, cutensorPlanPreference_t pref,
185 cutensorWorksizePreference_t ws_pref = CUTENSOR_WORKSPACE_DEFAULT) {
186 std::uint64_t size = 0;
187 cutensor_error_check(cutensorEstimateWorkspaceSize(get_handle(), op_desc, pref, ws_pref, &size),
"cutensorEstimateWorkspaceSize");
192 template <
typename T>
193 void permute_impl(T alpha,
const_tensor_view<T> A, std::string_view idx_A, tensor_view<T> B, std::string_view idx_B) {
194 auto &handle = get_handle();
197 auto desc_A = create_tensor_desc(A);
198 auto desc_B = create_tensor_desc(B);
201 auto modes_A = to_modes(idx_A);
202 auto modes_B = to_modes(idx_B);
205 cutensorOperationDescriptor_t op_desc{};
206 cutensor_error_check(cutensorCreatePermutation(handle, &op_desc, desc_A, modes_A.data(), to_cutensor_unary_op(A.op), desc_B, modes_B.data(),
207 cutensor_compute_type<T>()),
208 "cutensorCreatePermutation");
211 auto pref = create_plan_pref();
212 auto ws_limit = estimate_workspace(op_desc, pref);
213 auto plan = create_plan(op_desc, pref, ws_limit);
216 cutensor_error_check(cutensorPermute(handle, plan, cuscalar(alpha), A.data, B.data,
nullptr ),
"cutensorPermute");
223 destroy_plan_pref(pref);
224 destroy_op_desc(op_desc);
225 destroy_tensor_desc(desc_A);
226 destroy_tensor_desc(desc_B);
231 template <
typename T>
234 auto &handle = get_handle();
237 auto desc_A = create_tensor_desc(A);
238 auto desc_C = create_tensor_desc(C);
239 auto desc_D = create_tensor_desc(D);
242 auto modes_A = to_modes(idx_A);
243 auto modes_C = to_modes(idx_C);
246 cutensorOperationDescriptor_t op_desc{};
247 cutensor_error_check(cutensorCreateElementwiseBinary(handle, &op_desc, desc_A, modes_A.data(), to_cutensor_unary_op(A.op), desc_C,
248 modes_C.data(), to_cutensor_unary_op(C.op), desc_D, modes_C.data(),
249 to_cutensor_binary_op(op_AC), cutensor_compute_type<T>()),
250 "cutensorCreateElementwiseBinary");
253 auto pref = create_plan_pref();
254 auto ws_limit = estimate_workspace(op_desc, pref);
255 auto plan = create_plan(op_desc, pref, ws_limit);
258 cutensor_error_check(
259 cutensorElementwiseBinaryExecute(handle, plan, cuscalar(alpha), A.data, cuscalar(gamma), C.data, D.data,
nullptr ),
260 "cutensorElementwiseBinaryExecute");
267 destroy_plan_pref(pref);
268 destroy_op_desc(op_desc);
269 destroy_tensor_desc(desc_A);
270 destroy_tensor_desc(desc_C);
271 destroy_tensor_desc(desc_D);
276 template <
typename T>
279 auto &handle = get_handle();
282 auto desc_A = create_tensor_desc(A);
283 auto desc_B = create_tensor_desc(B);
284 auto desc_C = create_tensor_desc(C);
285 auto desc_D = create_tensor_desc(D);
288 auto modes_A = to_modes(idx_A);
289 auto modes_B = to_modes(idx_B);
290 auto modes_C = to_modes(idx_C);
293 cutensorOperationDescriptor_t op_desc{};
294 cutensor_error_check(cutensorCreateElementwiseTrinary(handle, &op_desc, desc_A, modes_A.data(), to_cutensor_unary_op(A.op), desc_B,
295 modes_B.data(), to_cutensor_unary_op(B.op), desc_C, modes_C.data(),
296 to_cutensor_unary_op(C.op), desc_D, modes_C.data(), to_cutensor_binary_op(op_AB),
297 to_cutensor_binary_op(op_ABC), cutensor_compute_type<T>()),
298 "cutensorCreateElementwiseTrinary");
301 auto pref = create_plan_pref();
302 auto ws_limit = estimate_workspace(op_desc, pref);
303 auto plan = create_plan(op_desc, pref, ws_limit);
306 cutensor_error_check(cutensorElementwiseTrinaryExecute(handle, plan, cuscalar(alpha), A.data, cuscalar(beta), B.data, cuscalar(gamma), C.data,
308 "cutensorElementwiseTrinaryExecute");
315 destroy_plan_pref(pref);
316 destroy_op_desc(op_desc);
317 destroy_tensor_desc(desc_A);
318 destroy_tensor_desc(desc_B);
319 destroy_tensor_desc(desc_C);
320 destroy_tensor_desc(desc_D);
326 template <
typename T>
329 auto &handle = get_handle();
332 auto desc_A = create_tensor_desc(A);
333 auto desc_C = create_tensor_desc(C);
334 auto desc_D = create_tensor_desc(D);
337 auto modes_A = to_modes(idx_A);
338 auto modes_C = to_modes(idx_C);
341 cutensorOperationDescriptor_t op_desc{};
342 cutensor_error_check(cutensorCreateReduction(handle, &op_desc, desc_A, modes_A.data(), to_cutensor_unary_op(A.op), desc_C, modes_C.data(),
343 to_cutensor_unary_op(C.op), desc_D, modes_C.data(), to_cutensor_binary_op(op_reduce),
344 cutensor_compute_type<T>()),
345 "cutensorCreateReduction");
348 auto pref = create_plan_pref();
349 auto ws_limit = estimate_workspace(op_desc, pref);
350 auto plan = create_plan(op_desc, pref, ws_limit);
353 std::uint64_t ws_size = 0;
354 cutensor_error_check(cutensorPlanGetAttribute(handle, plan, CUTENSOR_PLAN_REQUIRED_WORKSPACE, &ws_size,
sizeof(ws_size)),
355 "cutensorPlanGetAttribute");
358 void *workspace =
nullptr;
359 if (ws_size > 0) {
device_error_check(cudaMalloc(&workspace, ws_size),
"cudaMalloc"); }
362 cutensor_error_check(
363 cutensorReduce(handle, plan, cuscalar(alpha), A.data, cuscalar(beta), C.data, D.data, workspace, ws_size,
nullptr ),
374 destroy_plan_pref(pref);
375 destroy_op_desc(op_desc);
376 destroy_tensor_desc(desc_A);
377 destroy_tensor_desc(desc_C);
378 destroy_tensor_desc(desc_D);
383 template <
typename T>
386 auto &handle = get_handle();
389 auto desc_A = create_tensor_desc(A);
390 auto desc_B = create_tensor_desc(B);
391 auto desc_C = create_tensor_desc(C);
392 auto desc_D = create_tensor_desc(D);
395 auto modes_A = to_modes(idx_A);
396 auto modes_B = to_modes(idx_B);
397 auto modes_C = to_modes(idx_C);
400 cutensorOperationDescriptor_t op_desc{};
401 cutensor_error_check(cutensorCreateContraction(handle, &op_desc, desc_A, modes_A.data(), to_cutensor_unary_op(A.op), desc_B, modes_B.data(),
402 to_cutensor_unary_op(B.op), desc_C, modes_C.data(), to_cutensor_unary_op(C.op), desc_D,
403 modes_C.data(), cutensor_compute_type<T>()),
404 "cutensorCreateContraction");
407 auto pref = create_plan_pref();
408 auto ws_limit = estimate_workspace(op_desc, pref);
409 auto plan = create_plan(op_desc, pref, ws_limit);
412 std::uint64_t ws_size = 0;
413 cutensor_error_check(cutensorPlanGetAttribute(handle, plan, CUTENSOR_PLAN_REQUIRED_WORKSPACE, &ws_size,
sizeof(ws_size)),
414 "cutensorPlanGetAttribute");
417 void *workspace =
nullptr;
418 if (ws_size > 0) {
device_error_check(cudaMalloc(&workspace, ws_size),
"cudaMalloc"); }
421 cutensor_error_check(
422 cutensorContract(handle, plan, cuscalar(alpha), A.data, B.data, cuscalar(beta), C.data, D.data, workspace, ws_size,
nullptr ),
433 destroy_plan_pref(pref);
434 destroy_op_desc(op_desc);
435 destroy_tensor_desc(desc_A);
436 destroy_tensor_desc(desc_B);
437 destroy_tensor_desc(desc_C);
438 destroy_tensor_desc(desc_D);
444 void permute(
float alpha,
const_tensor_view<float> A, std::string_view idx_A, tensor_view<float> B, std::string_view idx_B) {
445 permute_impl(alpha, A, idx_A, B, idx_B);
447 void permute(
double alpha,
const_tensor_view<double> A, std::string_view idx_A, tensor_view<double> B, std::string_view idx_B) {
448 permute_impl(alpha, A, idx_A, B, idx_B);
450 void permute(std::complex<float> alpha,
const_tensor_view<std::complex<float>> A, std::string_view idx_A, tensor_view<std::complex<float>> B,
451 std::string_view idx_B) {
452 permute_impl(alpha, A, idx_A, B, idx_B);
454 void permute(std::complex<double> alpha,
const_tensor_view<std::complex<double>> A, std::string_view idx_A, tensor_view<std::complex<double>> B,
455 std::string_view idx_B) {
456 permute_impl(alpha, A, idx_A, B, idx_B);
461 std::string_view idx_C, tensor_view<float> D,
binary_op op_AC) {
462 elementwise_binary_impl(alpha, A, idx_A, gamma, C, idx_C, D, op_AC);
465 std::string_view idx_C, tensor_view<double> D,
binary_op op_AC) {
466 elementwise_binary_impl(alpha, A, idx_A, gamma, C, idx_C, D, op_AC);
468 void elementwise_binary(std::complex<float> alpha,
const_tensor_view<std::complex<float>> A, std::string_view idx_A, std::complex<float> gamma,
470 elementwise_binary_impl(alpha, A, idx_A, gamma, C, idx_C, D, op_AC);
472 void elementwise_binary(std::complex<double> alpha,
const_tensor_view<std::complex<double>> A, std::string_view idx_A, std::complex<double> gamma,
474 elementwise_binary_impl(alpha, A, idx_A, gamma, C, idx_C, D, op_AC);
481 elementwise_trinary_impl(alpha, A, idx_A, beta, B, idx_B, gamma, C, idx_C, D, op_AB, op_ABC);
486 elementwise_trinary_impl(alpha, A, idx_A, beta, B, idx_B, gamma, C, idx_C, D, op_AB, op_ABC);
488 void elementwise_trinary(std::complex<float> alpha,
const_tensor_view<std::complex<float>> A, std::string_view idx_A, std::complex<float> beta,
489 const_tensor_view<std::complex<float>> B, std::string_view idx_B, std::complex<float> gamma,
492 elementwise_trinary_impl(alpha, A, idx_A, beta, B, idx_B, gamma, C, idx_C, D, op_AB, op_ABC);
494 void elementwise_trinary(std::complex<double> alpha,
const_tensor_view<std::complex<double>> A, std::string_view idx_A, std::complex<double> beta,
495 const_tensor_view<std::complex<double>> B, std::string_view idx_B, std::complex<double> gamma,
498 elementwise_trinary_impl(alpha, A, idx_A, beta, B, idx_B, gamma, C, idx_C, D, op_AB, op_ABC);
503 tensor_view<float> D,
binary_op op_reduce) {
504 reduce_impl(alpha, A, idx_A, beta, C, idx_C, D, op_reduce);
507 tensor_view<double> D,
binary_op op_reduce) {
508 reduce_impl(alpha, A, idx_A, beta, C, idx_C, D, op_reduce);
510 void reduce(std::complex<float> alpha,
const_tensor_view<std::complex<float>> A, std::string_view idx_A, std::complex<float> beta,
512 reduce_impl(alpha, A, idx_A, beta, C, idx_C, D, op_reduce);
514 void reduce(std::complex<double> alpha,
const_tensor_view<std::complex<double>> A, std::string_view idx_A, std::complex<double> beta,
515 const_tensor_view<std::complex<double>> C, std::string_view idx_C, tensor_view<std::complex<double>> D,
binary_op op_reduce) {
516 reduce_impl(alpha, A, idx_A, beta, C, idx_C, D, op_reduce);
522 contract_impl(alpha, A, idx_A, B, idx_B, beta, C, idx_C, D);
526 contract_impl(alpha, A, idx_A, B, idx_B, beta, C, idx_C, D);
529 std::string_view idx_B, std::complex<float> beta,
const_tensor_view<std::complex<float>> C, std::string_view idx_C,
530 tensor_view<std::complex<float>> D) {
531 contract_impl(alpha, A, idx_A, B, idx_B, beta, C, idx_C, D);
533 void contract(std::complex<double> alpha,
const_tensor_view<std::complex<double>> A, std::string_view idx_A,
534 const_tensor_view<std::complex<double>> B, std::string_view idx_B, std::complex<double> beta,
535 const_tensor_view<std::complex<double>> C, std::string_view idx_C, tensor_view<std::complex<double>> D) {
536 contract_impl(alpha, A, idx_A, B, idx_B, beta, C, idx_C, D);
Provides concepts for the nda library.
Provides a C++ interface for various cuTENSOR routines.
Provides GPU and non-GPU specific functionality.
Provides a custom runtime error class and macros to assert conditions and throw exceptions.
#define device_error_check(ARG1, ARG2)
Trigger a compilation error every time the nda::device_error_check function is called.
void cuda_device_sync(bool do_sync=true, std::string_view func="")
Empty function if CudaSupport is not enabled.
unary_op
Unary element-wise operations for tensor operations.
binary_op
Binary operations for tensor operations.
tensor_view< const T > const_tensor_view
Alias for a tensor_view with const value type.