23#ifndef EIGEN_GPU_DEVICE_EXPR_H
24#define EIGEN_GPU_DEVICE_EXPR_H
27#include "./InternalHeaderCheck.h"
29#include "./CuBlasSupport.h"
30#include "./type_traits.h"
37template <
typename T,
typename S>
38using require_host_scalar_convertible_t =
39 typename std::enable_if<std::is_convertible<T, S>::value && !is_device_scalar<typename std::decay<T>::type>::value,
45template <
typename Scalar_>
48 using Scalar = Scalar_;
57template <
typename Scalar_>
60 using Scalar = Scalar_;
76template <
typename Inner_>
79 using Inner = std::decay_t<Inner_>;
80 using Scalar = internal::scalar_type_t<Inner>;
81 Scaled(Scalar alpha,
const Inner_& inner) : alpha_(alpha), inner_(inner) {}
82 Scalar scalar()
const {
return alpha_; }
83 const Inner_& inner()
const {
return inner_; }
91template <
typename Lhs,
typename Rhs>
94 using Scalar = internal::scalar_type_t<Lhs>;
95 static_assert(std::is_same<Scalar, internal::scalar_type_t<Rhs>>::value,
96 "DeviceMatrix GEMM: LHS and RHS must have the same scalar type");
98 GemmExpr(
const Lhs& lhs,
const Rhs& rhs) : lhs_(lhs), rhs_(rhs) {}
99 const Lhs& lhs()
const {
return lhs_; }
100 const Rhs& rhs()
const {
return rhs_; }
116template <
typename T,
typename S,
internal::require_host_scalar_convertible_t<T, S> = 0>
118 return {
static_cast<S
>(alpha), m};
121template <
typename T,
typename S,
internal::require_host_scalar_convertible_t<T, S> = 0>
122Scaled<DeviceMatrix<S>> operator*(
const DeviceMatrix<S>& m, T alpha) {
123 return {
static_cast<S
>(alpha), m};
126template <
typename T,
typename S,
internal::require_host_scalar_convertible_t<T, S> = 0>
128 return {S(1) /
static_cast<S
>(alpha), m};
136template <
typename T,
typename S,
internal::require_host_scalar_convertible_t<T, S> = 0>
138 return {
static_cast<S
>(alpha), m};
141template <
typename T,
typename S,
internal::require_host_scalar_convertible_t<T, S> = 0>
143 return {
static_cast<S
>(alpha), m};
146template <
typename T,
typename S,
internal::require_host_scalar_convertible_t<T, S> = 0>
148 return {
static_cast<S
>(alpha), m};
151template <
typename T,
typename S,
internal::require_host_scalar_convertible_t<T, S> = 0>
153 return {
static_cast<S
>(alpha), m};
157template <
typename T,
typename Inner,
158 internal::require_host_scalar_convertible_t<T, internal::scalar_type_t<Inner>> = 0>
160 using S = internal::scalar_type_t<Inner>;
161 return {
static_cast<S
>(alpha) * s.scalar(), s.inner()};
164template <
typename Inner>
166 using S = internal::scalar_type_t<Inner>;
167 return {S(-1) * s.scalar(), s.inner()};
174 static constexpr bool is_device_expr =
false;
177template <
typename Scalar>
179 using scalar_type = Scalar;
180 static constexpr GpuOp op = GpuOp::NoTrans;
186template <
typename Scalar>
187struct device_expr_traits<AdjointView<Scalar>> {
188 using scalar_type = Scalar;
189 static constexpr GpuOp op = GpuOp::ConjTrans;
190 static constexpr bool is_device_expr =
true;
191 static const DeviceMatrix<Scalar>& matrix(
const AdjointView<Scalar>& x) {
return x.matrix(); }
192 static Scalar alpha(
const AdjointView<Scalar>&) {
return Scalar(1); }
195template <
typename Scalar>
197 using scalar_type = Scalar;
198 static constexpr GpuOp op = GpuOp::Trans;
199 static constexpr bool is_device_expr =
true;
200 static const DeviceMatrix<Scalar>& matrix(
const TransposeView<Scalar>& x) {
return x.matrix(); }
201 static Scalar alpha(
const TransposeView<Scalar>&) {
return Scalar(1); }
204template <
typename Inner>
206 using scalar_type = scalar_type_t<Inner>;
207 static constexpr GpuOp op = device_expr_traits<Inner>::op;
208 static constexpr bool is_device_expr =
true;
209 static const DeviceMatrix<scalar_type>& matrix(
const Scaled<Inner>& x) {
210 return device_expr_traits<Inner>::matrix(x.inner());
212 static scalar_type alpha(
const Scaled<Inner>& x) {
return x.scalar() * device_expr_traits<Inner>::alpha(x.inner()); }
216template <
typename Lhs,
typename Rhs,
217 std::enable_if_t<internal::device_expr_traits<Lhs>::is_device_expr &&
218 internal::device_expr_traits<Rhs>::is_device_expr,
230template <
typename Scalar_>
231class DeviceScaledDevice {
233 using Scalar = Scalar_;
259template <
typename Scalar_>
262 using Scalar = Scalar_;
264 : alpha_(alpha), A_(A), beta_(beta), B_(B) {}
265 Scalar alpha()
const {
return alpha_; }
266 Scalar beta()
const {
return beta_; }
280 return {S(1), a, S(1), b};
285DeviceAddExpr<S> operator+(
const DeviceMatrix<S>& a,
const Scaled<DeviceMatrix<S>>& b) {
286 return {S(1), a, b.scalar(), b.inner()};
292 return {a.scalar(), a.inner(), S(1), b};
298 return {S(1), a, S(-1), b};
304 return {S(1), a, -b.scalar(), b.inner()};
310 return {a.scalar(), a.inner(), S(-1), b};
316 return {a.scalar(), a.inner(), b.scalar(), b.inner()};
321 return {a.scalar(), a.inner(), -b.scalar(), b.inner()};
View returned by DeviceMatrix::adjoint(); maps to the cuBLAS conjugate-transpose operand flag.
Definition DeviceExpr.h:46
Linear combination of two device matrices.
Definition DeviceExpr.h:260
RAII wrapper for a dense column-major matrix in GPU device memory.
Definition DeviceMatrix.h:122
RAII wrapper for a scalar in GPU device memory.
Definition DeviceScalar.h:33
Expression that scales a device matrix by a DeviceScalar.
Definition DeviceExpr.h:231
Expression returned by operator*(lhs_expr, rhs_expr), dispatched to cuBLAS GEMM.
Definition DeviceExpr.h:92
Expression returned by operator*(Scalar, DeviceMatrix/View), carrying the scalar factor.
Definition DeviceExpr.h:77
View returned by DeviceMatrix::transpose(); maps to the cuBLAS transpose operand flag.
Definition DeviceExpr.h:58
Namespace containing all symbols from the Eigen library.
Describes GPU device expression types.
Definition DeviceExpr.h:173
Definition type_traits.h:873