25 #include "blas_blasfeo.hpp"
26 #include "casadi/core/code_generator.hpp"
27 #include "casadi/core/exception.hpp"
34 casadi_int m, casadi_int n, casadi_int k,
36 const double* A, casadi_int lda,
37 const double* B, casadi_int ldb,
39 double* C, casadi_int ldc) {
40 casadi_assert(m <= INT_MAX && n <= INT_MAX && k <= INT_MAX
41 && lda <= INT_MAX && ldb <= INT_MAX && ldc <= INT_MAX,
42 "BLAS 'blasfeo' plugin: matrix dimension exceeds 32-bit BLAS ABI.");
46 int m_ =
static_cast<int>(m), n_ =
static_cast<int>(n), k_ =
static_cast<int>(k);
47 int lda_ =
static_cast<int>(lda), ldb_ =
static_cast<int>(ldb), ldc_ =
static_cast<int>(ldc);
51 blasfeo_blas_dgemm(&ta, &tb, &m_, &n_, &k_, &alpha,
52 const_cast<double*
>(A), &lda_,
53 const_cast<double*
>(B), &ldb_,
60 const double* x,
double* y) {
61 casadi_assert(n <= INT_MAX,
62 "BLAS 'blasfeo' plugin: vector length exceeds 32-bit BLAS ABI.");
63 int n_ =
static_cast<int>(n), inc = 1;
64 blasfeo_blas_daxpy(&n_, &alpha,
65 const_cast<double*
>(x), &inc, y, &inc);
69 const double* x,
const double* y) {
70 casadi_assert(n <= INT_MAX,
71 "BLAS 'blasfeo' plugin: vector length exceeds 32-bit BLAS ABI.");
72 int n_ =
static_cast<int>(n), inc = 1;
73 return blasfeo_blas_ddot(&n_,
74 const_cast<double*
>(x), &inc,
75 const_cast<double*
>(y), &inc);
79 "/* BLAS \"blasfeo\" plugin: namespaced Fortran ABI, link with -lblasfeo */\n"
80 "extern void blasfeo_blas_dgemm(char* transa, char* transb,\n"
81 " int* m, int* n, int* k,\n"
83 " double* A, int* lda,\n"
84 " double* B, int* ldb,\n"
86 " double* C, int* ldc);";
89 "extern void blasfeo_blas_daxpy(int* n, double* alpha,\n"
90 " double* x, int* incx,\n"
91 " double* y, int* incy);";
94 "extern double blasfeo_blas_ddot(int* n,\n"
95 " double* x, int* incx,\n"
96 " double* y, int* incy);";
99 const std::string& A, casadi_int m, casadi_int k,
100 const std::string& B, casadi_int n,
const std::string& C) {
104 g.
local(
"blas_m",
"int");
105 g.
local(
"blas_n",
"int");
106 g.
local(
"blas_k",
"int");
107 g <<
"blas_m = " << m <<
"; blas_n = " << n <<
"; blas_k = " << k <<
";\n";
108 g <<
"blasfeo_blas_dgemm(&blas_tn, &blas_tn, &blas_m, &blas_n, &blas_k, "
109 "&blas_one, (double*)" << A <<
", &blas_m, (double*)" << B
110 <<
", &blas_k, &blas_one, " <<
C <<
", &blas_m);\n";
118 const std::vector<std::string>& inst) {
119 static const char* SRC =
120 "// SYMBOL \"axpy\"\n"
121 "void casadi_axpy(casadi_int n, casadi_real alpha,\n"
122 " const casadi_real* x, casadi_real* y) {\n"
123 " int n_ = (int)n, inc = 1;\n"
124 " blasfeo_blas_daxpy(&n_, &alpha, (double*)x, &inc, y, &inc);\n"
135 const std::vector<std::string>& inst) {
136 static const char* SRC =
137 "// SYMBOL \"dot\"\n"
138 "casadi_real casadi_dot(casadi_int n,\n"
139 " const casadi_real* x, const casadi_real* y) {\n"
140 " int n_ = (int)n, inc = 1;\n"
141 " return blasfeo_blas_ddot(&n_, (double*)x, &inc, (double*)y, &inc);\n"
147 extern "C" int CASADI_BLAS_BLASFEO_EXPORT
149 plugin->name =
"blasfeo";
151 plugin->version = CASADI_VERSION;
156 plugin->exposed.dscal =
nullptr;
157 plugin->exposed.dnrm2 =
nullptr;
158 plugin->exposed.dasum =
nullptr;
159 plugin->exposed.dcopy =
nullptr;
162 plugin->exposed.codegen_scal_aux =
nullptr;
163 plugin->exposed.codegen_nrm2_aux =
nullptr;
164 plugin->exposed.codegen_asum_aux =
nullptr;
165 plugin->options =
nullptr;
166 plugin->deserialize =
nullptr;
167 plugin->creator =
nullptr;
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)
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.
int CASADI_BLAS_BLASFEO_EXPORT casadi_register_blas_blasfeo(Blas::Plugin *plugin)
static const char * BLASFEO_DDOT_DECL
static void blasfeo_daxpy(casadi_int n, double alpha, const double *x, double *y)
static void blasfeo_codegen_dot_aux(CodeGenerator &g, const std::vector< std::string > &inst)
void CASADI_BLAS_BLASFEO_EXPORT casadi_load_blas_blasfeo()
static void blasfeo_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 blasfeo_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 void blasfeo_codegen_axpy_aux(CodeGenerator &g, const std::vector< std::string > &inst)
static const char * BLASFEO_DECL
static double blasfeo_ddot(casadi_int n, const double *x, const double *y)
static const char * BLASFEO_DAXPY_DECL