LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
linear.h
Go to the documentation of this file.
1// Copyright 2021 CEA LIST
2// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
6#ifndef DEEPLIMA_SRC_INFERENCE_EIGEN_LINEAR_H
7#define DEEPLIMA_SRC_INFERENCE_EIGEN_LINEAR_H
8
9#include <eigen3/Eigen/Dense>
10#include "op_base.h"
11
12namespace deeplima
13{
14namespace eigen_impl
15{
16
17template<class M=Eigen::MatrixXf, class V=Eigen::VectorXf>
19{
22};
23
24template<class M, class V, class T>
25class Op_Linear : public Op_Base
26{
27protected:
29 {
31 {
32 }
33 virtual ~workbench_t() {}
34 };
35
36public:
37 typedef V Vector;
39
40 virtual std::shared_ptr<Op_Base::workbench_t> create_workbench([[maybe_unused]] uint32_t input_size,
41 [[maybe_unused]] const std::shared_ptr<param_base_t> params,
42 bool /*precomputed_input=false*/) const override
43 {
44 assert(input_size > 0);
45 assert(nullptr != params);
46 // TODO should it be used?
47 // const params_linear_t<M, V>& layer = *static_cast<const params_t*>(params);
48
49 return std::make_shared<workbench_t>();
50 }
51
52 virtual size_t execute([[maybe_unused]] std::shared_ptr<Op_Base::workbench_t> pwb,
53 const V& input,
54 const std::shared_ptr<param_base_t> params,
55 Vector& output)
56 {
57 assert(nullptr != pwb);
58 assert(nullptr != params);
59 const auto& layer = std::dynamic_pointer_cast<const params_t>(params);
60
61 // TODO should it be used?
62 // auto wb = std::dynamic_pointer_cast<workbench_t>(pwb);
63
64 output = (layer->weight * input).colwise() + layer->bias;
65
66 return 0;
67 }
68};
69
70} // namespace eigen_impl
71} // namespace deeplima
72
73#endif
params_linear_t< M, V > params_t
Definition linear.h:38
virtual size_t execute(std::shared_ptr< Op_Base::workbench_t > pwb, const V &input, const std::shared_ptr< param_base_t > params, Vector &output)
Definition linear.h:52
virtual std::shared_ptr< Op_Base::workbench_t > create_workbench(uint32_t input_size, const std::shared_ptr< param_base_t > params, bool) const override
Definition linear.h:40