15#ifndef EIGEN_GPU_DEVICE_SCALAR_OPS_H
16#define EIGEN_GPU_DEVICE_SCALAR_OPS_H
18#include <cuda_runtime.h>
19#include <npps_arithmetic_and_logical_operations.h>
21#include "./GpuSupport.h"
28#define EIGEN_NPP_CHECK(expr) \
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__); \
38inline NppStreamContext make_npp_stream_ctx(cudaStream_t stream) {
39 NppStreamContext ctx = {};
41#if CUDART_VERSION >= 12080
44 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamGetDevice(stream, &ctx.nCudaDeviceId));
48 EIGEN_CUDA_RUNTIME_CHECK(cudaGetDevice(&ctx.nCudaDeviceId));
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));
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));
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));
81inline void device_scalar_sqrt(
float* a,
const NppStreamContext& npp_ctx) {
82 EIGEN_NPP_CHECK(nppsSqrt_32f_I_Ctx(a, 1, npp_ctx));
85inline void device_scalar_sqrt(
double* a,
const NppStreamContext& npp_ctx) {
86 EIGEN_NPP_CHECK(nppsSqrt_64f_I_Ctx(a, 1, npp_ctx));
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));
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));
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));
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));
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));
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));
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));
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));
Namespace containing all symbols from the Eigen library.