14#ifndef EIGEN_GPU_EIGENSOLVER_H
15#define EIGEN_GPU_EIGENSOLVER_H
18#include "./InternalHeaderCheck.h"
20#include "./GpuSolverContext.h"
24template <
typename Scalar_>
25class SelfAdjointEigenSolver {
27 using Scalar = Scalar_;
28 using RealScalar =
typename NumTraits<Scalar>::Real;
29 using PlainMatrix = Eigen::Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
30 using RealVector = Eigen::Matrix<RealScalar, Dynamic, 1>;
32 SelfAdjointEigenSolver() =
default;
37 explicit SelfAdjointEigenSolver(Context& ctx) : solver_ctx_(ctx) {}
40 template <
typename InputType>
41 explicit SelfAdjointEigenSolver(
const DenseBase<InputType>& A,
int options =
ComputeEigenvectors) {
45 explicit SelfAdjointEigenSolver(
const DeviceMatrix<Scalar>& d_A,
int options =
ComputeEigenvectors) {
46 compute(d_A, options);
50 explicit SelfAdjointEigenSolver(DeviceMatrix<Scalar>&& d_A,
int options =
ComputeEigenvectors) {
51 compute(std::move(d_A), options);
55 template <
typename InputType>
56 SelfAdjointEigenSolver(Context& ctx,
const DenseBase<InputType>& A,
int options =
ComputeEigenvectors)
62 SelfAdjointEigenSolver(Context& ctx,
const DeviceMatrix<Scalar>& d_A,
int options =
ComputeEigenvectors)
64 compute(d_A, options);
68 SelfAdjointEigenSolver(Context& ctx, DeviceMatrix<Scalar>&& d_A,
int options =
ComputeEigenvectors)
70 compute(std::move(d_A), options);
73 ~SelfAdjointEigenSolver() =
default;
75 SelfAdjointEigenSolver(
const SelfAdjointEigenSolver&) =
delete;
76 SelfAdjointEigenSolver& operator=(
const SelfAdjointEigenSolver&) =
delete;
78 SelfAdjointEigenSolver(SelfAdjointEigenSolver&& o) noexcept
79 : solver_ctx_(std::move(o.solver_ctx_)),
80 d_A_(std::move(o.d_A_)),
81 d_W_(std::move(o.d_W_)),
82 compute_eigenvectors_(o.compute_eigenvectors_),
85 o.compute_eigenvectors_ =
true;
90 SelfAdjointEigenSolver& operator=(SelfAdjointEigenSolver&& o)
noexcept {
92 solver_ctx_ = std::move(o.solver_ctx_);
93 d_A_ = std::move(o.d_A_);
94 d_W_ = std::move(o.d_W_);
95 compute_eigenvectors_ = o.compute_eigenvectors_;
98 o.compute_eigenvectors_ =
true;
105 template <
typename InputType>
106 SelfAdjointEigenSolver& compute(
const DenseBase<InputType>& A,
int options =
ComputeEigenvectors) {
112 SelfAdjointEigenSolver& compute(
const DeviceMatrix<Scalar>& d_A,
int options =
ComputeEigenvectors) {
113 if (!begin_compute(d_A, options))
return *
this;
115 const size_t mat_bytes =
static_cast<size_t>(lda_) *
static_cast<size_t>(n_) *
sizeof(Scalar);
116 internal::ensure_sized(d_A_, mat_bytes);
117 EIGEN_CUDA_RUNTIME_CHECK(
118 cudaMemcpyAsync(d_A_.get(), d_A.data(), mat_bytes, cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
127 SelfAdjointEigenSolver& compute(DeviceMatrix<Scalar>&& d_A,
int options =
ComputeEigenvectors) {
128 if (d_A.isView())
return compute(
static_cast<const DeviceMatrix<Scalar>&
>(d_A), options);
129 if (!begin_compute(d_A, options))
return *
this;
131 d_A_ = internal::DeviceBuffer::adopt(
static_cast<void*
>(d_A.release()),
132 static_cast<size_t>(lda_) *
static_cast<size_t>(n_) *
sizeof(Scalar));
140 Index cols()
const {
return n_; }
141 Index rows()
const {
return n_; }
144 RealVector eigenvalues()
const {
145 eigen_assert(solver_ctx_.info() ==
Success);
148 solver_ctx_.download(W.data(), d_W_.get(),
static_cast<size_t>(n_) *
sizeof(RealScalar));
155 PlainMatrix eigenvectors()
const {
156 eigen_assert(solver_ctx_.info() ==
Success);
157 eigen_assert(compute_eigenvectors_ &&
"eigenvectors() requires ComputeEigenvectors option");
158 PlainMatrix V(n_, n_);
160 solver_ctx_.download(V.data(), d_A_.get(),
static_cast<size_t>(lda_) *
static_cast<size_t>(n_) *
sizeof(Scalar));
173 DeviceMatrix<RealScalar> d_eigenvalues()
const {
174 eigen_assert(solver_ctx_.info() ==
Success);
176 v.recordReady(solver_ctx_.stream());
182 DeviceMatrix<Scalar> d_eigenvectors()
const {
183 eigen_assert(solver_ctx_.info() ==
Success);
184 eigen_assert(compute_eigenvectors_ &&
"d_eigenvectors() requires ComputeEigenvectors option");
186 v.recordReady(solver_ctx_.stream());
190 cudaStream_t stream()
const {
return solver_ctx_.stream(); }
193 mutable internal::GpuSolverContext solver_ctx_;
196 internal::DeviceBuffer d_A_;
197 internal::DeviceBuffer d_W_;
198 bool compute_eigenvectors_ =
true;
204 bool begin_compute(
const DeviceMatrix<Scalar>& d_A,
int options) {
205 eigen_assert(d_A.rows() == d_A.cols() &&
"SelfAdjointEigenSolver requires a square matrix");
207 "options must be ComputeEigenvectors or EigenvaluesOnly");
210 if (!solver_ctx_.begin_compute(n_ != 0)) {
211 d_A_ = internal::DeviceBuffer();
212 d_W_ = internal::DeviceBuffer();
216 d_A.waitReady(solver_ctx_.stream());
221 constexpr cudaDataType_t dtype = internal::cusolver_data_type<Scalar>::value;
222 constexpr cudaDataType_t rtype = internal::cuda_data_type<RealScalar>::value;
224 solver_ctx_.mark_pending();
226 internal::ensure_sized(d_W_,
static_cast<size_t>(n_) *
sizeof(RealScalar));
228 const cusolverEigMode_t jobz = compute_eigenvectors_ ? CUSOLVER_EIG_MODE_VECTOR : CUSOLVER_EIG_MODE_NOVECTOR;
230 constexpr cublasFillMode_t uplo = CUBLAS_FILL_MODE_LOWER;
232 size_t dev_ws = 0, host_ws = 0;
233 EIGEN_CUSOLVER_CHECK(cusolverDnXsyevd_bufferSize(solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, jobz, uplo,
234 n_, dtype, d_A_.get(), lda_, rtype, d_W_.get(), dtype, &dev_ws,
237 solver_ctx_.ensure_scratch(dev_ws);
238 solver_ctx_.h_workspace_.resize(host_ws);
240 EIGEN_CUSOLVER_CHECK(cusolverDnXsyevd(solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, jobz, uplo, n_, dtype,
241 d_A_.get(), lda_, rtype, d_W_.get(), dtype, solver_ctx_.scratch_workspace(),
242 dev_ws, host_ws > 0 ? solver_ctx_.h_workspace_.data() :
nullptr, host_ws,
243 solver_ctx_.scratch_info()));
245 solver_ctx_.enqueue_info_copy();
static DeviceMatrix fromHost(const DenseBase< Derived > &host, cudaStream_t stream=nullptr)
Definition DeviceMatrix.h:226
static DeviceMatrix view(Scalar *device_ptr, Index rows, Index cols)
Definition DeviceMatrix.h:576
Namespace containing all symbols from the Eigen library.