Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
DeviceScalar.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 for deferred host synchronization.
12//
13// Reductions (dot, nrm2) write their result straight to device memory under
14// CUBLAS_POINTER_MODE_DEVICE, so no host sync happens until the value is read.
15// Conversion to Scalar is that read. Because the first conversion flushes the
16// stream, later conversions in the same expression only download: a CG iteration
17// costs one sync rather than three.
18
19#ifndef EIGEN_GPU_DEVICE_SCALAR_H
20#define EIGEN_GPU_DEVICE_SCALAR_H
21
22// IWYU pragma: private
23#include "./InternalHeaderCheck.h"
24
25#include "./GpuSupport.h"
26#include "./DeviceScalarOps.h"
27
28namespace Eigen {
29namespace gpu {
30
32template <typename Scalar_>
34 public:
35 using Scalar = Scalar_;
36
39 explicit DeviceScalar(cudaStream_t stream = nullptr) : d_val_(sizeof(Scalar)), stream_(stream) {}
40
41 DeviceScalar(Scalar host_val, cudaStream_t stream) : d_val_(sizeof(Scalar)), stream_(stream) {
42 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(d_val_.get(), &host_val, sizeof(Scalar), cudaMemcpyHostToDevice, stream_));
43 }
44
45 DeviceScalar(DeviceScalar&& o) noexcept : d_val_(std::move(o.d_val_)), stream_(o.stream_) { o.stream_ = nullptr; }
46
47 DeviceScalar& operator=(DeviceScalar&& o) noexcept {
48 if (this != &o) {
49 d_val_ = std::move(o.d_val_);
50 stream_ = o.stream_;
51 o.stream_ = nullptr;
52 }
53 return *this;
54 }
55
60 DeviceScalar(const DeviceScalar& o) : d_val_(sizeof(Scalar)), stream_(o.stream_) {
61 EIGEN_CUDA_RUNTIME_CHECK(
62 cudaMemcpyAsync(d_val_.get(), o.d_val_.get(), sizeof(Scalar), cudaMemcpyDeviceToDevice, stream_));
63 }
64
69 if (this != &o) *this = DeviceScalar(o);
70 return *this;
71 }
72
74 Scalar get() const {
75 Scalar result;
76 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(&result, d_val_.get(), sizeof(Scalar), cudaMemcpyDeviceToHost, stream_));
77 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(stream_));
78 return result;
79 }
80
83 operator Scalar() const { return get(); }
84
85 Scalar* devicePtr() { return static_cast<Scalar*>(d_val_.get()); }
86 const Scalar* devicePtr() const { return static_cast<const Scalar*>(d_val_.get()); }
87 cudaStream_t stream() const { return stream_; }
88
89 // The arithmetic below keeps results on device via the NPP helpers in
90 // DeviceScalarOps.h, which cover real types only: dividing or negating a
91 // complex DeviceScalar does not compile, so convert it to Scalar (a host sync)
92 // first. Unlike DeviceMatrix, DeviceScalar tracks no cross-stream readiness,
93 // so all operands must share one stream.
94
95 friend DeviceScalar operator/(const DeviceScalar& a, const DeviceScalar& b) {
96 eigen_assert(a.stream_ == b.stream_ && "DeviceScalar operator/: operands must share the same stream");
97 DeviceScalar result(a.stream_);
98 gpu::internal::device_scalar_div(a.devicePtr(), b.devicePtr(), result.devicePtr(), a.stream_);
99 return result;
100 }
101
102 friend DeviceScalar operator/(Scalar a, const DeviceScalar& b) {
103 DeviceScalar d_a(a, b.stream_);
104 return d_a / b;
105 }
106
107 friend DeviceScalar operator/(const DeviceScalar& a, Scalar b) {
108 DeviceScalar d_b(b, a.stream_);
109 return a / d_b;
110 }
111
112 DeviceScalar operator-() const {
113 DeviceScalar result(stream_);
114 gpu::internal::device_scalar_neg(devicePtr(), result.devicePtr(), stream_);
115 return result;
116 }
117
118 private:
119 internal::DeviceBuffer d_val_;
120 cudaStream_t stream_ = nullptr;
121};
122
123} // namespace gpu
124} // namespace Eigen
125
126#endif // EIGEN_GPU_DEVICE_SCALAR_H
DeviceScalar(const DeviceScalar &o)
Definition DeviceScalar.h:60
DeviceScalar(cudaStream_t stream=nullptr)
Definition DeviceScalar.h:39
DeviceScalar & operator=(const DeviceScalar &o)
Definition DeviceScalar.h:68
Scalar get() const
Definition DeviceScalar.h:74
Namespace containing all symbols from the Eigen library.