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"
36 casadi_int m, casadi_int n, casadi_int k,
38 const double* A, casadi_int lda,
39 const double* B, casadi_int ldb,
41 double* C, casadi_int ldc) {
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.");
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);
54 dgemm_(&ta, &tb, &m_, &n_, &k_, &alpha,
55 A, &lda_, B, &ldb_, &beta,
C, &ldc_ CASADI_BLAS_CLASSIC_CHARLEN_ARGS);
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);
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);
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);
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);
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);
97 if (x) std::memcpy(y, x, n *
sizeof(
double));
98 else std::memset(y, 0, n *
sizeof(
double));
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"
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"
112 "#define CASADI_BLAS_CLASSIC_CHARLEN_DECL\n"
113 "#define CASADI_BLAS_CLASSIC_CHARLEN_ARGS\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);";
124 "#ifndef CASADI_BLAS_DAXPY\n"
125 "#define CASADI_BLAS_DAXPY daxpy_\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);";
132 "#ifndef CASADI_BLAS_DDOT\n"
133 "#define CASADI_BLAS_DDOT ddot_\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);";
140 "#ifndef CASADI_BLAS_DSCAL\n"
141 "#define CASADI_BLAS_DSCAL dscal_\n"
143 "extern void CASADI_BLAS_DSCAL(const int* n, const double* alpha,\n"
144 " double* x, const int* incx);";
147 "#ifndef CASADI_BLAS_DNRM2\n"
148 "#define CASADI_BLAS_DNRM2 dnrm2_\n"
150 "extern double CASADI_BLAS_DNRM2(const int* n, const double* x, const int* incx);";
153 "#ifndef CASADI_BLAS_DASUM\n"
154 "#define CASADI_BLAS_DASUM dasum_\n"
156 "extern double CASADI_BLAS_DASUM(const int* n, const double* x, const int* incx);";
159 const std::string& A, casadi_int m, casadi_int k,
160 const std::string& B, casadi_int n,
const std::string& C) {
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";
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"
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"
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"
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"
228 const std::vector<std::string>& inst) {
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"
241 extern "C" int CASADI_BLAS_CLASSIC_EXPORT
243 plugin->name =
"classic";
245 plugin->version = CASADI_VERSION;
259 plugin->options =
nullptr;
260 plugin->deserialize =
nullptr;
261 plugin->creator =
nullptr;
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() {
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.
static const char * CLASSIC_DAXPY_DECL
static const char * CLASSIC_DSCAL_DECL
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