LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_inference_base.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_EIGEN_WRP_BIRNN_INFERENCE_BASE_H
7#define DEEPLIMA_EIGEN_WRP_BIRNN_INFERENCE_BASE_H
8
9#include <eigen3/Eigen/Dense>
10
11#include "bilstm.h"
12#include "bilstm_and_dense.h"
13
14#include "embd_dict.h"
15
17
18namespace deeplima
19{
20namespace eigen_impl
21{
22
24{
25public:
26 typedef Eigen::MatrixXf Matrix;
27 typedef Eigen::VectorXf Vector;
28 typedef float Scalar;
29 typedef Eigen::MatrixXf tensor_t;
32
33 virtual ~BiRnnInferenceBase() = default;
34
35 virtual void load(const std::string& fn) = 0;
36
38 {
39 return m_input_uint_dicts;
40 }
41
42 inline const std::vector<std::string>& get_input_uint_dicts_names() const
43 {
45 }
46
48 {
49 return m_input_str_dicts;
50 }
51
52 inline const std::vector<std::string>& get_input_str_dicts_names() const
53 {
55 }
56
57 inline const std::vector<std::vector<std::string>>& get_output_str_dicts() const
58 {
59 return m_output_str_dicts;
60 }
61
62 inline const std::vector<std::string>& get_output_str_dicts_names() const
63 {
65 }
66
67 virtual size_t init_new_worker(size_t input_len, bool precomputed_input=false)
68 {
69 assert(m_wb.size() > 0);
70 assert(m_ops.size() > 0);
71 assert(m_ops.size() == m_params.size());
72
73 size_t new_worker_idx = m_wb[0].size();
74 for (size_t i = 0; i < m_ops.size(); i++)
75 {
76 assert(m_wb[i].size() == new_worker_idx);
77 assert(nullptr != m_params[i]);
78 m_wb[i].push_back(m_ops[i]->create_workbench(input_len, m_params[i], precomputed_input));
79 }
80
81 return new_worker_idx;
82 }
83
84 virtual void precompute_inputs(
85 const Eigen::MatrixXf& inputs,
86 Eigen::MatrixXf& outputs,
87 int64_t input_size
88 ) = 0;
89
90 virtual void predict(
91 size_t worker_id,
92 const Eigen::MatrixXf& inputs,
93 int64_t input_begin,
94 int64_t input_end,
95 int64_t output_begin,
96 int64_t output_end,
97 std::shared_ptr< StdMatrix<uint8_t> >& output,
98 const std::vector<std::string>& outputs_names
99 ) = 0;
100
101protected:
102 std::vector<std::shared_ptr<Op_Base>> m_ops;
103 std::vector<std::shared_ptr<param_base_t>> m_params; // TODO: replace this
104
105 std::vector<std::vector<std::shared_ptr<Op_Base::workbench_t>>> m_wb; // outer - calculation step, inner - worker id
106
108 std::vector<std::string> m_input_uint_dicts_names;
110 std::vector<std::string> m_input_str_dicts_names;
111
112 std::vector<std::vector<std::string>> m_output_str_dicts;
113 std::vector<std::string> m_output_str_dicts_names;
114
116 std::vector<params_bilstm_spec_t> m_lstm;
117 std::map<std::string, size_t> m_lstm_idx;
118
120 std::vector<std::shared_ptr<params_multilayer_bilstm_spec_t>> m_multi_bilstm;
121 std::map<std::string, size_t> m_multi_bilstm_idx;
122
123 std::vector<params_linear_t<Eigen::MatrixXf, Eigen::VectorXf>> m_linear;
124 std::map<std::string, size_t> m_linear_idx;
125
126 virtual void convert_from_torch(const std::string& fn) = 0;
128};
129
130} // namespace eigen_impl
131} // namespace deeplima
132
133#endif
std::map< std::string, size_t > m_multi_bilstm_idx
std::vector< std::vector< std::shared_ptr< Op_Base::workbench_t > > > m_wb
virtual void convert_from_torch(const std::string &fn)=0
std::vector< std::shared_ptr< Op_Base > > m_ops
std::vector< std::string > m_output_str_dicts_names
virtual void convert_dicts_and_embeddings(const nets::BiRnnClassifierImpl &src)
virtual void predict(size_t worker_id, const Eigen::MatrixXf &inputs, int64_t input_begin, int64_t input_end, int64_t output_begin, int64_t output_end, std::shared_ptr< StdMatrix< uint8_t > > &output, const std::vector< std::string > &outputs_names)=0
params_bilstm_t< Eigen::MatrixXf, Eigen::VectorXf > params_bilstm_spec_t
const std::vector< std::vector< std::string > > & get_output_str_dicts() const
const uint_dicts_holder_t & get_input_uint_dicts() const
std::map< std::string, size_t > m_linear_idx
virtual size_t init_new_worker(size_t input_len, bool precomputed_input=false)
std::vector< std::string > m_input_uint_dicts_names
const std::vector< std::string > & get_output_str_dicts_names() const
const std::vector< std::string > & get_input_uint_dicts_names() const
std::vector< std::string > m_input_str_dicts_names
std::vector< std::shared_ptr< params_multilayer_bilstm_spec_t > > m_multi_bilstm
const str_dicts_holder_t & get_input_str_dicts() const
std::map< std::string, size_t > m_lstm_idx
virtual void load(const std::string &fn)=0
std::vector< params_linear_t< Eigen::MatrixXf, Eigen::VectorXf > > m_linear
std::vector< std::vector< std::string > > m_output_str_dicts
const std::vector< std::string > & get_input_str_dicts_names() const
std::vector< params_bilstm_spec_t > m_lstm
std::vector< std::shared_ptr< param_base_t > > m_params
virtual void precompute_inputs(const Eigen::MatrixXf &inputs, Eigen::MatrixXf &outputs, int64_t input_size)=0
params_multilayer_bilstm_t< Eigen::MatrixXf, Eigen::VectorXf > params_multilayer_bilstm_spec_t