blas_classic.cpp
1 /*
2  * This file is part of CasADi.
3  *
4  * CasADi -- A symbolic framework for dynamic optimization.
5  * Copyright (C) 2010-2023 Joel Andersson, Joris Gillis, Moritz Diehl,
6  * KU Leuven. All rights reserved.
7  *
8  * CasADi is free software; you can redistribute it and/or
9  * modify it under the terms of the GNU Lesser General Public
10  * License as published by the Free Software Foundation; either
11  * version 3 of the License, or (at your option) any later version.
12  *
13  * CasADi is distributed in the hope that it will be useful,
14  * but WITHOUT ANY WARRANTY; without even the implied warranty of
15  * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the GNU
16  * Lesser General Public License for more details.
17  *
18  * You should have received a copy of the GNU Lesser General Public
19  * License along with CasADi; if not, write to the Free Software
20  * Foundation, Inc., 51 Franklin Street, Fifth Floor, Boston, MA 02110-1301 USA
21  *
22  */
23 
24 
25 #include "blas_classic.hpp"
26 #include "casadi/core/code_generator.hpp"
27 #include "casadi/core/exception.hpp"
28 #include "casadi/core/runtime/casadi_runtime.hpp"
29 
30 #include <climits>
31 #include <cstring>
32 
33 namespace casadi {
34 
35  static void classic_dgemm(int transa, int transb,
36  casadi_int m, casadi_int n, casadi_int k,
37  double alpha,
38  const double* A, casadi_int lda,
39  const double* B, casadi_int ldb,
40  double beta,
41  double* C, casadi_int ldc) {
42  // Fortran BLAS-32 takes int by reference. Guard against ILP64 / huge
43  // problems before silently truncating.
44  casadi_assert(m <= INT_MAX && n <= INT_MAX && k <= INT_MAX
45  && lda <= INT_MAX && ldb <= INT_MAX && ldc <= INT_MAX,
46  "BLAS 'classic' plugin: matrix dimension exceeds 32-bit BLAS ABI. "
47  "Rebuild against an ILP64 BLAS or use the \"reference\" plugin.");
48 
49  const char ta = (transa == CASADI_BLAS_TRANS) ? 'T' : 'N';
50  const char tb = (transb == CASADI_BLAS_TRANS) ? 'T' : 'N';
51  int m_ = static_cast<int>(m), n_ = static_cast<int>(n), k_ = static_cast<int>(k);
52  int lda_ = static_cast<int>(lda), ldb_ = static_cast<int>(ldb), ldc_ = static_cast<int>(ldc);
53 
54  dgemm_(&ta, &tb, &m_, &n_, &k_, &alpha,
55  A, &lda_, B, &ldb_, &beta, C, &ldc_ CASADI_BLAS_CLASSIC_CHARLEN_ARGS);
56  }
57 
58  static void classic_daxpy(casadi_int n, double alpha,
59  const double* x, double* y) {
60  casadi_assert(n <= INT_MAX,
61  "BLAS 'classic' plugin: vector length exceeds 32-bit BLAS ABI.");
62  int n_ = static_cast<int>(n), inc = 1;
63  daxpy_(&n_, &alpha, x, &inc, y, &inc);
64  }
65 
66  static double classic_ddot(casadi_int n,
67  const double* x, const double* y) {
68  casadi_assert(n <= INT_MAX,
69  "BLAS 'classic' plugin: vector length exceeds 32-bit BLAS ABI.");
70  int n_ = static_cast<int>(n), inc = 1;
71  return ddot_(&n_, x, &inc, y, &inc);
72  }
73 
74  static void classic_dscal(casadi_int n, double alpha, double* x) {
75  casadi_assert(n <= INT_MAX,
76  "BLAS 'classic' plugin: vector length exceeds 32-bit BLAS ABI.");
77  int n_ = static_cast<int>(n), inc = 1;
78  dscal_(&n_, &alpha, x, &inc);
79  }
80 
81  static double classic_dnrm2(casadi_int n, const double* x) {
82  casadi_assert(n <= INT_MAX,
83  "BLAS 'classic' plugin: vector length exceeds 32-bit BLAS ABI.");
84  int n_ = static_cast<int>(n), inc = 1;
85  return dnrm2_(&n_, x, &inc);
86  }
87 
88  static double classic_dasum(casadi_int n, const double* x) {
89  casadi_assert(n <= INT_MAX,
90  "BLAS 'classic' plugin: vector length exceeds 32-bit BLAS ABI.");
91  int n_ = static_cast<int>(n), inc = 1;
92  return dasum_(&n_, x, &inc);
93  }
94 
95  static void classic_dcopy(const double* x, casadi_int n, double* y) {
96  if (!y) return;
97  if (x) std::memcpy(y, x, n * sizeof(double));
98  else std::memset(y, 0, n * sizeof(double));
99  }
100 
101  static const char* CLASSIC_DECL =
102  "/* BLAS \"classic\" plugin: Fortran ABI dgemm_, link with -lblas/-lopenblas/-lmkl. */\n"
103  "/* Override the symbol at compile time with -DCASADI_BLAS_DGEMM=my_dgemm. */\n"
104  "#ifndef CASADI_BLAS_DGEMM\n"
105  "#define CASADI_BLAS_DGEMM dgemm_\n"
106  "#endif\n"
107  "#ifdef __EMSCRIPTEN__\n"
108  "#include <stddef.h>\n"
109  "#define CASADI_BLAS_CLASSIC_CHARLEN_DECL , size_t, size_t\n"
110  "#define CASADI_BLAS_CLASSIC_CHARLEN_ARGS , 1, 1\n"
111  "#else\n"
112  "#define CASADI_BLAS_CLASSIC_CHARLEN_DECL\n"
113  "#define CASADI_BLAS_CLASSIC_CHARLEN_ARGS\n"
114  "#endif\n"
115  "extern void CASADI_BLAS_DGEMM(const char* transa, const char* transb,\n"
116  " const int* m, const int* n, const int* k,\n"
117  " const double* alpha,\n"
118  " const double* A, const int* lda,\n"
119  " const double* B, const int* ldb,\n"
120  " const double* beta,\n"
121  " double* C, const int* ldc CASADI_BLAS_CLASSIC_CHARLEN_DECL);";
122 
123  static const char* CLASSIC_DAXPY_DECL =
124  "#ifndef CASADI_BLAS_DAXPY\n"
125  "#define CASADI_BLAS_DAXPY daxpy_\n"
126  "#endif\n"
127  "extern void CASADI_BLAS_DAXPY(const int* n, const double* alpha,\n"
128  " const double* x, const int* incx,\n"
129  " double* y, const int* incy);";
130 
131  static const char* CLASSIC_DDOT_DECL =
132  "#ifndef CASADI_BLAS_DDOT\n"
133  "#define CASADI_BLAS_DDOT ddot_\n"
134  "#endif\n"
135  "extern double CASADI_BLAS_DDOT(const int* n,\n"
136  " const double* x, const int* incx,\n"
137  " const double* y, const int* incy);";
138 
139  static const char* CLASSIC_DSCAL_DECL =
140  "#ifndef CASADI_BLAS_DSCAL\n"
141  "#define CASADI_BLAS_DSCAL dscal_\n"
142  "#endif\n"
143  "extern void CASADI_BLAS_DSCAL(const int* n, const double* alpha,\n"
144  " double* x, const int* incx);";
145 
146  static const char* CLASSIC_DNRM2_DECL =
147  "#ifndef CASADI_BLAS_DNRM2\n"
148  "#define CASADI_BLAS_DNRM2 dnrm2_\n"
149  "#endif\n"
150  "extern double CASADI_BLAS_DNRM2(const int* n, const double* x, const int* incx);";
151 
152  static const char* CLASSIC_DASUM_DECL =
153  "#ifndef CASADI_BLAS_DASUM\n"
154  "#define CASADI_BLAS_DASUM dasum_\n"
155  "#endif\n"
156  "extern double CASADI_BLAS_DASUM(const int* n, const double* x, const int* incx);";
157 
159  const std::string& A, casadi_int m, casadi_int k,
160  const std::string& B, casadi_int n, const std::string& C) {
162  // Locals shared with other Fortran-ABI BLAS plugins. g.local is
163  // idempotent on (name,type) match, so multiple plugins coexisting in
164  // the same Function reuse the same locals safely.
165  g.local("blas_tn", "char"); g.init_local("blas_tn", "'N'");
166  g.local("blas_one", "double"); g.init_local("blas_one", "1.0");
167  g.local("blas_m", "int");
168  g.local("blas_n", "int");
169  g.local("blas_k", "int");
170  g << "blas_m = " << m << "; blas_n = " << n << "; blas_k = " << k << ";\n";
171  g << "CASADI_BLAS_DGEMM(&blas_tn, &blas_tn, &blas_m, &blas_n, &blas_k, "
172  "&blas_one, " << A << ", &blas_m, " << B << ", &blas_k, "
173  "&blas_one, " << C << ", &blas_m CASADI_BLAS_CLASSIC_CHARLEN_ARGS);\n";
174  }
175 
176 
178  const std::vector<std::string>& inst) {
179  static const char* SRC =
180  "// SYMBOL \"axpy\"\n"
181  "void casadi_axpy(casadi_int n, casadi_real alpha,\n"
182  " const casadi_real* x, casadi_real* y) {\n"
183  " int n_ = (int)n, inc = 1;\n"
184  " CASADI_BLAS_DAXPY(&n_, &alpha, x, &inc, y, &inc);\n"
185  "}\n";
186  g.auxiliaries << CLASSIC_DAXPY_DECL << "\n";
187  g.auxiliaries << g.sanitize_source(SRC, inst);
188  }
189 
191  const std::vector<std::string>& inst) {
192  static const char* SRC =
193  "// SYMBOL \"scal\"\n"
194  "void casadi_scal(casadi_int n, casadi_real alpha, casadi_real* x) {\n"
195  " int n_ = (int)n, inc = 1;\n"
196  " CASADI_BLAS_DSCAL(&n_, &alpha, x, &inc);\n"
197  "}\n";
198  g.auxiliaries << CLASSIC_DSCAL_DECL << "\n";
199  g.auxiliaries << g.sanitize_source(SRC, inst);
200  }
201 
203  const std::vector<std::string>& inst) {
204  static const char* SRC =
205  "// SYMBOL \"dot\"\n"
206  "casadi_real casadi_dot(casadi_int n,\n"
207  " const casadi_real* x, const casadi_real* y) {\n"
208  " int n_ = (int)n, inc = 1;\n"
209  " return CASADI_BLAS_DDOT(&n_, x, &inc, y, &inc);\n"
210  "}\n";
211  g.auxiliaries << CLASSIC_DDOT_DECL << "\n";
212  g.auxiliaries << g.sanitize_source(SRC, inst);
213  }
214 
216  const std::vector<std::string>& inst) {
217  static const char* SRC =
218  "// SYMBOL \"norm_2\"\n"
219  "casadi_real casadi_norm_2(casadi_int n, const casadi_real* x) {\n"
220  " int n_ = (int)n, inc = 1;\n"
221  " return CASADI_BLAS_DNRM2(&n_, x, &inc);\n"
222  "}\n";
223  g.auxiliaries << CLASSIC_DNRM2_DECL << "\n";
224  g.auxiliaries << g.sanitize_source(SRC, inst);
225  }
226 
228  const std::vector<std::string>& inst) {
229  // Match casadi_norm_1<T>'s null-pointer semantics: returns 0 on x==NULL.
230  static const char* SRC =
231  "// SYMBOL \"norm_1\"\n"
232  "casadi_real casadi_norm_1(casadi_int n, const casadi_real* x) {\n"
233  " int n_ = (int)n, inc = 1;\n"
234  " if (!x) return 0;\n"
235  " return CASADI_BLAS_DASUM(&n_, x, &inc);\n"
236  "}\n";
237  g.auxiliaries << CLASSIC_DASUM_DECL << "\n";
238  g.auxiliaries << g.sanitize_source(SRC, inst);
239  }
240 
241  extern "C" int CASADI_BLAS_CLASSIC_EXPORT
242  casadi_register_blas_classic(Blas::Plugin* plugin) {
243  plugin->name = "classic";
244  plugin->doc = BlasClassic::meta_doc.c_str();
245  plugin->version = CASADI_VERSION;
246  plugin->exposed.dgemm = &classic_dgemm;
247  plugin->exposed.codegen_mtimes = &classic_codegen_mtimes;
248  plugin->exposed.daxpy = &classic_daxpy;
249  plugin->exposed.ddot = &classic_ddot;
250  plugin->exposed.dscal = &classic_dscal;
251  plugin->exposed.dnrm2 = &classic_dnrm2;
252  plugin->exposed.dasum = &classic_dasum;
253  plugin->exposed.dcopy = &classic_dcopy;
254  plugin->exposed.codegen_axpy_aux = &classic_codegen_axpy_aux;
255  plugin->exposed.codegen_dot_aux = &classic_codegen_dot_aux;
256  plugin->exposed.codegen_scal_aux = &classic_codegen_scal_aux;
257  plugin->exposed.codegen_nrm2_aux = &classic_codegen_nrm2_aux;
258  plugin->exposed.codegen_asum_aux = &classic_codegen_asum_aux;
259  plugin->options = nullptr;
260  plugin->deserialize = nullptr;
261  plugin->creator = nullptr;
262  return 0;
263  }
264 
265  extern "C" void CASADI_BLAS_CLASSIC_EXPORT casadi_load_blas_classic() {
267  }
268 
269 #if defined(CASADI_CORE_BLAS_DEPENDENCY) && defined(CASADI_L1_BLAS)
270  extern "C" void CASADI_BLAS_CLASSIC_EXPORT casadi_blas_classic_set_l1_hooks() {
271  casadi_daxpy_hook = &classic_daxpy;
272  casadi_ddot_hook = &classic_ddot;
273  casadi_dscal_hook = &classic_dscal;
274  casadi_dnrm2_hook = &classic_dnrm2;
275  casadi_dasum_hook = &classic_dasum;
276  casadi_dcopy_hook = &classic_dcopy;
277  }
278 #endif
279 
280 } // namespace casadi
static const std::string meta_doc
A documentation string.
Helper class for C code generation.
void local(const std::string &name, const std::string &type, const std::string &ref="")
Declare a local variable.
void add_external(const std::string &new_external, const std::string &name="")
Add an external function declaration.
void init_local(const std::string &name, const std::string &def)
Specify the default value for a local variable.
std::string sanitize_source(const std::string &src, const std::vector< std::string > &inst, bool add_shorthand=true)
Sanitize source files for codegen.
std::stringstream auxiliaries
static void registerPlugin(const Plugin &plugin, bool needs_lock=true)
Register an integrator in the factory.
The casadi namespace.
Definition: archiver.cpp:28
static const char * CLASSIC_DAXPY_DECL
static const char * CLASSIC_DSCAL_DECL
@ CASADI_BLAS_TRANS
Definition: blas_impl.hpp:42
static void classic_codegen_nrm2_aux(CodeGenerator &g, const std::vector< std::string > &inst)
void CASADI_BLAS_CLASSIC_EXPORT casadi_load_blas_classic()
static void classic_codegen_asum_aux(CodeGenerator &g, const std::vector< std::string > &inst)
static const char * CLASSIC_DECL
static const char * CLASSIC_DNRM2_DECL
static void classic_codegen_dot_aux(CodeGenerator &g, const std::vector< std::string > &inst)
static double classic_dasum(casadi_int n, const double *x)
static const char * CLASSIC_DASUM_DECL
static void classic_codegen_mtimes(CodeGenerator &g, const std::string &A, casadi_int m, casadi_int k, const std::string &B, casadi_int n, const std::string &C)
static void classic_dcopy(const double *x, casadi_int n, double *y)
int CASADI_BLAS_CLASSIC_EXPORT casadi_register_blas_classic(Blas::Plugin *plugin)
static void classic_dgemm(int transa, int transb, casadi_int m, casadi_int n, casadi_int k, double alpha, const double *A, casadi_int lda, const double *B, casadi_int ldb, double beta, double *C, casadi_int ldc)
static double classic_ddot(casadi_int n, const double *x, const double *y)
static void classic_codegen_axpy_aux(CodeGenerator &g, const std::vector< std::string > &inst)
static void classic_daxpy(casadi_int n, double alpha, const double *x, double *y)
static void classic_codegen_scal_aux(CodeGenerator &g, const std::vector< std::string > &inst)
static void classic_dscal(casadi_int n, double alpha, double *x)
static double classic_dnrm2(casadi_int n, const double *x)
static const char * CLASSIC_DDOT_DECL