Eigen  5.0.1
 
Loading...
Searching...
No Matches
TriangularMatrixVector_BLAS.h
1/*
2 Copyright (c) 2011, Intel Corporation. All rights reserved.
3
4 Redistribution and use in source and binary forms, with or without modification,
5 are permitted provided that the following conditions are met:
6
7 * Redistributions of source code must retain the above copyright notice, this
8 list of conditions and the following disclaimer.
9 * Redistributions in binary form must reproduce the above copyright notice,
10 this list of conditions and the following disclaimer in the documentation
11 and/or other materials provided with the distribution.
12 * Neither the name of Intel Corporation nor the names of its contributors may
13 be used to endorse or promote products derived from this software without
14 specific prior written permission.
15
16 THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND
17 ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED
18 WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
19 DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT OWNER OR CONTRIBUTORS BE LIABLE FOR
20 ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES
21 (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES;
22 LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON
23 ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT
24 (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS
25 SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
26
27 ********************************************************************************
28 * Content : Eigen bindings to BLAS F77
29 * Triangular matrix-vector product functionality based on ?TRMV.
30 ********************************************************************************
31*/
32// SPDX-License-Identifier: BSD-3-Clause
33
34#ifndef EIGEN_TRIANGULAR_MATRIX_VECTOR_BLAS_H
35#define EIGEN_TRIANGULAR_MATRIX_VECTOR_BLAS_H
36
37// IWYU pragma: private
38#include "../InternalHeaderCheck.h"
39
40namespace Eigen {
41
42namespace internal {
43
44/**********************************************************************
45 * This file implements triangular matrix-vector multiplication using BLAS
46 **********************************************************************/
47
48// trmv/hemv specialization
49
50template <typename Index, int Mode, typename LhsScalar, bool ConjLhs, typename RhsScalar, bool ConjRhs,
51 int StorageOrder>
52struct triangular_matrix_vector_product_trmv
53 : triangular_matrix_vector_product<Index, Mode, LhsScalar, ConjLhs, RhsScalar, ConjRhs, StorageOrder, BuiltIn> {};
54
55#define EIGEN_BLAS_TRMV_SPECIALIZE(Scalar) \
56 template <typename Index, int Mode, bool ConjLhs, bool ConjRhs> \
57 struct triangular_matrix_vector_product<Index, Mode, Scalar, ConjLhs, Scalar, ConjRhs, ColMajor, Specialized> { \
58 static void run(Index rows_, Index cols_, const Scalar* lhs_, Index lhsStride, const Scalar* rhs_, Index rhsIncr, \
59 Scalar* res_, Index resIncr, Scalar alpha) { \
60 triangular_matrix_vector_product_trmv<Index, Mode, Scalar, ConjLhs, Scalar, ConjRhs, ColMajor>::run( \
61 rows_, cols_, lhs_, lhsStride, rhs_, rhsIncr, res_, resIncr, alpha); \
62 } \
63 }; \
64 template <typename Index, int Mode, bool ConjLhs, bool ConjRhs> \
65 struct triangular_matrix_vector_product<Index, Mode, Scalar, ConjLhs, Scalar, ConjRhs, RowMajor, Specialized> { \
66 static void run(Index rows_, Index cols_, const Scalar* lhs_, Index lhsStride, const Scalar* rhs_, Index rhsIncr, \
67 Scalar* res_, Index resIncr, Scalar alpha) { \
68 triangular_matrix_vector_product_trmv<Index, Mode, Scalar, ConjLhs, Scalar, ConjRhs, RowMajor>::run( \
69 rows_, cols_, lhs_, lhsStride, rhs_, rhsIncr, res_, resIncr, alpha); \
70 } \
71 };
72
73EIGEN_BLAS_TRMV_SPECIALIZE(double)
74EIGEN_BLAS_TRMV_SPECIALIZE(float)
75EIGEN_BLAS_TRMV_SPECIALIZE(dcomplex)
76EIGEN_BLAS_TRMV_SPECIALIZE(scomplex)
77
78// implements col-major: res += alpha * op(triangular) * vector
79#define EIGEN_BLAS_TRMV_CM(EIGTYPE, BLASTYPE, EIGPREFIX, BLASPREFIX, BLASPOSTFIX) \
80 template <typename Index, int Mode, bool ConjLhs, bool ConjRhs> \
81 struct triangular_matrix_vector_product_trmv<Index, Mode, EIGTYPE, ConjLhs, EIGTYPE, ConjRhs, ColMajor> { \
82 enum { \
83 IsLower = (Mode & Lower) == Lower, \
84 SetDiag = (Mode & (ZeroDiag | UnitDiag)) ? 0 : 1, \
85 IsUnitDiag = (Mode & UnitDiag) ? 1 : 0, \
86 IsZeroDiag = (Mode & ZeroDiag) ? 1 : 0, \
87 LowUp = IsLower ? Lower : Upper \
88 }; \
89 static void run(Index rows_, Index cols_, const EIGTYPE* lhs_, Index lhsStride, const EIGTYPE* rhs_, \
90 Index rhsIncr, EIGTYPE* res_, Index resIncr, EIGTYPE alpha) { \
91 if (rows_ == 0 || cols_ == 0) return; \
92 EIGEN_IF_CONSTEXPR (ConjLhs || IsZeroDiag) { \
93 triangular_matrix_vector_product<Index, Mode, EIGTYPE, ConjLhs, EIGTYPE, ConjRhs, ColMajor, BuiltIn>::run( \
94 rows_, cols_, lhs_, lhsStride, rhs_, rhsIncr, res_, resIncr, alpha); \
95 return; \
96 } \
97 Index size = (std::min)(rows_, cols_); \
98 Index rows = IsLower ? rows_ : size; \
99 Index cols = IsLower ? size : cols_; \
100 \
101 typedef VectorX##EIGPREFIX VectorRhs; \
102 EIGTYPE *x, *y; \
103 \
104 /* Set x*/ \
105 Map<const VectorRhs, 0, InnerStride<> > rhs(rhs_, cols, InnerStride<>(rhsIncr)); \
106 VectorRhs x_tmp; \
107 EIGEN_IF_CONSTEXPR (ConjRhs) \
108 x_tmp = rhs.conjugate(); \
109 else \
110 x_tmp = rhs; \
111 x = x_tmp.data(); \
112 \
113 /* Square part handling */ \
114 \
115 char trans, uplo, diag; \
116 BlasIndex m, n, lda, incx, incy; \
117 EIGTYPE const* a; \
118 EIGTYPE beta(1); \
119 \
120 /* Set m, n */ \
121 n = convert_index<BlasIndex>(size); \
122 lda = convert_index<BlasIndex>(lhsStride); \
123 incx = 1; \
124 incy = convert_index<BlasIndex>(resIncr); \
125 \
126 /* Set uplo, trans and diag*/ \
127 trans = 'N'; \
128 uplo = IsLower ? 'L' : 'U'; \
129 diag = IsUnitDiag ? 'U' : 'N'; \
130 \
131 /* call ?TRMV*/ \
132 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, trmv), BLASPOSTFIX) \
133 (&uplo, &trans, &diag, &n, (const BLASTYPE*)lhs_, &lda, (BLASTYPE*)x, &incx); \
134 \
135 /* Add op(a_tr)rhs into res*/ \
136 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, axpy), BLASPOSTFIX) \
137 (&n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)x, &incx, (BLASTYPE*)res_, &incy); \
138 /* Non-square case - doesn't fit to BLAS ?TRMV. Fall to default triangular product*/ \
139 if (size < (std::max)(rows, cols)) { \
140 EIGEN_IF_CONSTEXPR (ConjRhs) \
141 x_tmp = rhs.conjugate(); \
142 else \
143 x_tmp = rhs; \
144 x = x_tmp.data(); \
145 if (size < rows) { \
146 y = res_ + size * resIncr; \
147 a = lhs_ + size; \
148 m = convert_index<BlasIndex>(rows - size); \
149 n = convert_index<BlasIndex>(size); \
150 } else { \
151 x += size; \
152 y = res_; \
153 a = lhs_ + size * lda; \
154 m = convert_index<BlasIndex>(size); \
155 n = convert_index<BlasIndex>(cols - size); \
156 } \
157 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, gemv), BLASPOSTFIX) \
158 (&trans, &m, &n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)a, &lda, (const BLASTYPE*)x, \
159 &incx, (const BLASTYPE*)&numext::real_ref(beta), (BLASTYPE*)y, &incy); \
160 } \
161 } \
162 };
163
164#ifdef EIGEN_USE_MKL
165EIGEN_BLAS_TRMV_CM(double, double, d, d, )
166EIGEN_BLAS_TRMV_CM(dcomplex, MKL_Complex16, cd, z, )
167EIGEN_BLAS_TRMV_CM(float, float, f, s, )
168EIGEN_BLAS_TRMV_CM(scomplex, MKL_Complex8, cf, c, )
169#else
170EIGEN_BLAS_TRMV_CM(double, double, d, d, EIGEN_BLAS_POSTFIX)
171EIGEN_BLAS_TRMV_CM(dcomplex, double, cd, z, EIGEN_BLAS_POSTFIX)
172EIGEN_BLAS_TRMV_CM(float, float, f, s, EIGEN_BLAS_POSTFIX)
173EIGEN_BLAS_TRMV_CM(scomplex, float, cf, c, EIGEN_BLAS_POSTFIX)
174#endif
175
176// implements row-major: res += alpha * op(triangular) * vector
177#define EIGEN_BLAS_TRMV_RM(EIGTYPE, BLASTYPE, EIGPREFIX, BLASPREFIX, BLASPOSTFIX) \
178 template <typename Index, int Mode, bool ConjLhs, bool ConjRhs> \
179 struct triangular_matrix_vector_product_trmv<Index, Mode, EIGTYPE, ConjLhs, EIGTYPE, ConjRhs, RowMajor> { \
180 enum { \
181 IsLower = (Mode & Lower) == Lower, \
182 SetDiag = (Mode & (ZeroDiag | UnitDiag)) ? 0 : 1, \
183 IsUnitDiag = (Mode & UnitDiag) ? 1 : 0, \
184 IsZeroDiag = (Mode & ZeroDiag) ? 1 : 0, \
185 LowUp = IsLower ? Lower : Upper \
186 }; \
187 static void run(Index rows_, Index cols_, const EIGTYPE* lhs_, Index lhsStride, const EIGTYPE* rhs_, \
188 Index rhsIncr, EIGTYPE* res_, Index resIncr, EIGTYPE alpha) { \
189 if (rows_ == 0 || cols_ == 0) return; \
190 EIGEN_IF_CONSTEXPR (IsZeroDiag) { \
191 triangular_matrix_vector_product<Index, Mode, EIGTYPE, ConjLhs, EIGTYPE, ConjRhs, RowMajor, BuiltIn>::run( \
192 rows_, cols_, lhs_, lhsStride, rhs_, rhsIncr, res_, resIncr, alpha); \
193 return; \
194 } \
195 Index size = (std::min)(rows_, cols_); \
196 Index rows = IsLower ? rows_ : size; \
197 Index cols = IsLower ? size : cols_; \
198 \
199 typedef VectorX##EIGPREFIX VectorRhs; \
200 EIGTYPE *x, *y; \
201 \
202 /* Set x*/ \
203 Map<const VectorRhs, 0, InnerStride<> > rhs(rhs_, cols, InnerStride<>(rhsIncr)); \
204 VectorRhs x_tmp; \
205 EIGEN_IF_CONSTEXPR (ConjRhs) \
206 x_tmp = rhs.conjugate(); \
207 else \
208 x_tmp = rhs; \
209 x = x_tmp.data(); \
210 \
211 /* Square part handling */ \
212 \
213 char trans, uplo, diag; \
214 BlasIndex m, n, lda, incx, incy; \
215 EIGTYPE const* a; \
216 EIGTYPE beta(1); \
217 \
218 /* Set m, n */ \
219 n = convert_index<BlasIndex>(size); \
220 lda = convert_index<BlasIndex>(lhsStride); \
221 incx = 1; \
222 incy = convert_index<BlasIndex>(resIncr); \
223 \
224 /* Set uplo, trans and diag*/ \
225 trans = ConjLhs ? 'C' : 'T'; \
226 uplo = IsLower ? 'U' : 'L'; \
227 diag = IsUnitDiag ? 'U' : 'N'; \
228 \
229 /* call ?TRMV*/ \
230 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, trmv), BLASPOSTFIX) \
231 (&uplo, &trans, &diag, &n, (const BLASTYPE*)lhs_, &lda, (BLASTYPE*)x, &incx); \
232 \
233 /* Add op(a_tr)rhs into res*/ \
234 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, axpy), BLASPOSTFIX) \
235 (&n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)x, &incx, (BLASTYPE*)res_, &incy); \
236 /* Non-square case - doesn't fit to BLAS ?TRMV. Fall to default triangular product*/ \
237 if (size < (std::max)(rows, cols)) { \
238 EIGEN_IF_CONSTEXPR (ConjRhs) \
239 x_tmp = rhs.conjugate(); \
240 else \
241 x_tmp = rhs; \
242 x = x_tmp.data(); \
243 if (size < rows) { \
244 y = res_ + size * resIncr; \
245 a = lhs_ + size * lda; \
246 m = convert_index<BlasIndex>(rows - size); \
247 n = convert_index<BlasIndex>(size); \
248 } else { \
249 x += size; \
250 y = res_; \
251 a = lhs_ + size; \
252 m = convert_index<BlasIndex>(size); \
253 n = convert_index<BlasIndex>(cols - size); \
254 } \
255 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, gemv), BLASPOSTFIX) \
256 (&trans, &n, &m, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)a, &lda, (const BLASTYPE*)x, \
257 &incx, (const BLASTYPE*)&numext::real_ref(beta), (BLASTYPE*)y, &incy); \
258 } \
259 } \
260 };
261
262#ifdef EIGEN_USE_MKL
263EIGEN_BLAS_TRMV_RM(double, double, d, d, )
264EIGEN_BLAS_TRMV_RM(dcomplex, MKL_Complex16, cd, z, )
265EIGEN_BLAS_TRMV_RM(float, float, f, s, )
266EIGEN_BLAS_TRMV_RM(scomplex, MKL_Complex8, cf, c, )
267#else
268EIGEN_BLAS_TRMV_RM(double, double, d, d, EIGEN_BLAS_POSTFIX)
269EIGEN_BLAS_TRMV_RM(dcomplex, double, cd, z, EIGEN_BLAS_POSTFIX)
270EIGEN_BLAS_TRMV_RM(float, float, f, s, EIGEN_BLAS_POSTFIX)
271EIGEN_BLAS_TRMV_RM(scomplex, float, cf, c, EIGEN_BLAS_POSTFIX)
272#endif
273
274#undef EIGEN_BLAS_TRMV_RM
275#undef EIGEN_BLAS_TRMV_SPECIALIZE
276#undef EIGEN_BLAS_TRMV_CM
277} // namespace internal
278
279} // end namespace Eigen
280
281#endif // EIGEN_TRIANGULAR_MATRIX_VECTOR_BLAS_H