Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
DeviceScalarOps.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2026 Rasmus Munk Larsen <rmlarsen@gmail.com>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11// Device-resident scalar and element-wise operations built on the NPP signal
12// primitives (CUDA::npps), which keeps the module header-only: no custom CUDA
13// kernels to compile.
14
15#ifndef EIGEN_GPU_DEVICE_SCALAR_OPS_H
16#define EIGEN_GPU_DEVICE_SCALAR_OPS_H
17
18#include <cuda_runtime.h>
19#include <npps_arithmetic_and_logical_operations.h>
20
21#include "./GpuSupport.h"
22
23namespace Eigen {
24namespace gpu {
25namespace internal {
26
27// NPP statuses below NPP_SUCCESS are errors; positive ones are warnings.
28#define EIGEN_NPP_CHECK(expr) \
29 do { \
30 const NppStatus _s = (expr); \
31 if (_s < NPP_SUCCESS) \
32 ::Eigen::gpu::internal::gpu_check_failed_code("NPP", static_cast<int>(_s), #expr, __FILE__, __LINE__); \
33 } while (0)
34
35// NPP requires nCudaDeviceId and the device attributes to match the device that
36// owns `stream`. Re-querying per call is cheap next to the NPP launch itself and
37// keeps multi-device and borrowed-stream callers correct.
38inline NppStreamContext make_npp_stream_ctx(cudaStream_t stream) {
39 NppStreamContext ctx = {};
40 ctx.hStream = stream;
41#if CUDART_VERSION >= 12080
42 // cudaStreamGetDevice reports the stream's owning device irrespective of the
43 // calling thread's current device.
44 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamGetDevice(stream, &ctx.nCudaDeviceId));
45#else
46 // Without cudaStreamGetDevice (pre-CUDA 12.8), a caller borrowing a stream from
47 // another device must cudaSetDevice() first.
48 EIGEN_CUDA_RUNTIME_CHECK(cudaGetDevice(&ctx.nCudaDeviceId));
49#endif
50 EIGEN_CUDA_RUNTIME_CHECK(cudaDeviceGetAttribute(&ctx.nCudaDevAttrComputeCapabilityMajor,
51 cudaDevAttrComputeCapabilityMajor, ctx.nCudaDeviceId));
52 EIGEN_CUDA_RUNTIME_CHECK(cudaDeviceGetAttribute(&ctx.nCudaDevAttrComputeCapabilityMinor,
53 cudaDevAttrComputeCapabilityMinor, ctx.nCudaDeviceId));
54 EIGEN_CUDA_RUNTIME_CHECK(
55 cudaDeviceGetAttribute(&ctx.nMultiProcessorCount, cudaDevAttrMultiProcessorCount, ctx.nCudaDeviceId));
56 EIGEN_CUDA_RUNTIME_CHECK(cudaDeviceGetAttribute(&ctx.nMaxThreadsPerMultiProcessor,
57 cudaDevAttrMaxThreadsPerMultiProcessor, ctx.nCudaDeviceId));
58 EIGEN_CUDA_RUNTIME_CHECK(
59 cudaDeviceGetAttribute(&ctx.nMaxThreadsPerBlock, cudaDevAttrMaxThreadsPerBlock, ctx.nCudaDeviceId));
60 int shared_mem_per_block = 0;
61 EIGEN_CUDA_RUNTIME_CHECK(
62 cudaDeviceGetAttribute(&shared_mem_per_block, cudaDevAttrMaxSharedMemoryPerBlock, ctx.nCudaDeviceId));
63 ctx.nSharedMemPerBlock = static_cast<size_t>(shared_mem_per_block);
64 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamGetFlags(stream, &ctx.nStreamFlags));
65 return ctx;
66}
67
68// c = a / b. The operands are swapped because NPP computes
69// pDst[i] = pSrc2[i] / pSrc1[i].
70inline void device_scalar_div(const float* a, const float* b, float* c, cudaStream_t stream) {
71 NppStreamContext npp_ctx = make_npp_stream_ctx(stream);
72 EIGEN_NPP_CHECK(nppsDiv_32f_Ctx(b, a, c, 1, npp_ctx));
73}
74
75inline void device_scalar_div(const double* a, const double* b, double* c, cudaStream_t stream) {
76 NppStreamContext npp_ctx = make_npp_stream_ctx(stream);
77 EIGEN_NPP_CHECK(nppsDiv_64f_Ctx(b, a, c, 1, npp_ctx));
78}
79
80// a = sqrt(a), in place, with a context filled in before any capture began.
81inline void device_scalar_sqrt(float* a, const NppStreamContext& npp_ctx) {
82 EIGEN_NPP_CHECK(nppsSqrt_32f_I_Ctx(a, 1, npp_ctx));
83}
84
85inline void device_scalar_sqrt(double* a, const NppStreamContext& npp_ctx) {
86 EIGEN_NPP_CHECK(nppsSqrt_64f_I_Ctx(a, 1, npp_ctx));
87}
88
89// c = -a.
90inline void device_scalar_neg(const float* a, float* c, cudaStream_t stream) {
91 NppStreamContext npp_ctx = make_npp_stream_ctx(stream);
92 EIGEN_NPP_CHECK(nppsMulC_32f_Ctx(a, -1.0f, c, 1, npp_ctx));
93}
94
95inline void device_scalar_neg(const double* a, double* c, cudaStream_t stream) {
96 NppStreamContext npp_ctx = make_npp_stream_ctx(stream);
97 EIGEN_NPP_CHECK(nppsMulC_64f_Ctx(a, -1.0, c, 1, npp_ctx));
98}
99
100// c[i] = a[i] * b[i].
101inline void device_cwiseProduct(const float* a, const float* b, float* c, int n, cudaStream_t stream) {
102 NppStreamContext npp_ctx = make_npp_stream_ctx(stream);
103 EIGEN_NPP_CHECK(nppsMul_32f_Ctx(a, b, c, n, npp_ctx));
104}
105
106inline void device_cwiseProduct(const double* a, const double* b, double* c, int n, cudaStream_t stream) {
107 NppStreamContext npp_ctx = make_npp_stream_ctx(stream);
108 EIGEN_NPP_CHECK(nppsMul_64f_Ctx(a, b, c, n, npp_ctx));
109}
110
111// c[i] = a[i] / b[i], with operands swapped as in device_scalar_div.
112inline void device_cwiseQuotient(const float* a, const float* b, float* c, int n, cudaStream_t stream) {
113 NppStreamContext npp_ctx = make_npp_stream_ctx(stream);
114 EIGEN_NPP_CHECK(nppsDiv_32f_Ctx(b, a, c, n, npp_ctx));
115}
116
117inline void device_cwiseQuotient(const double* a, const double* b, double* c, int n, cudaStream_t stream) {
118 NppStreamContext npp_ctx = make_npp_stream_ctx(stream);
119 EIGEN_NPP_CHECK(nppsDiv_64f_Ctx(b, a, c, n, npp_ctx));
120}
121
122// x[i] /= alpha, a true division (NPP divide-by-constant, in place).
123inline void device_divC(float alpha, float* x, int n, cudaStream_t stream) {
124 NppStreamContext npp_ctx = make_npp_stream_ctx(stream);
125 EIGEN_NPP_CHECK(nppsDivC_32f_I_Ctx(alpha, x, n, npp_ctx));
126}
127
128inline void device_divC(double alpha, double* x, int n, cudaStream_t stream) {
129 NppStreamContext npp_ctx = make_npp_stream_ctx(stream);
130 EIGEN_NPP_CHECK(nppsDivC_64f_I_Ctx(alpha, x, n, npp_ctx));
131}
132
133} // namespace internal
134} // namespace gpu
135} // namespace Eigen
136
137#endif // EIGEN_GPU_DEVICE_SCALAR_OPS_H
Namespace containing all symbols from the Eigen library.