LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
entity_tagging_model.cpp
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
8
10
11
12using namespace std;
13using namespace torch;
14using namespace deeplima::convert_from_torch;
15using namespace deeplima::eigen_impl;
16
17namespace deeplima
18{
19namespace tagging
20{
21namespace eigen_impl
22{
23
24void convert_classes(const DictsHolder& src, vector<vector<string>>& classes);
25
26template class BiRnnEigenInferenceForTagging<float>;
27template class BiRnnEigenInferenceForTagging<int16_t>;
28
29template <typename AuxScalar>
31{
33 torch::load(src, fn, torch::Device(torch::kCPU));
34
35 // dicts and embeddings
36 Parent::convert_dicts_and_embeddings(src);
37 m_embd_fn.push_back(src.get_embd_fn(0));
38 assert(m_embd_fn[0].size() > 0);
39
40 // classes
41 convert_classes(src.get_classes(), Parent::m_output_str_dicts);
42 Parent::m_output_str_dicts_names = src.get_class_names();
43
44 // torch modules
45 Parent::m_lstm.reserve(src.get_layers_lstm().size());
46 for (size_t i = 0; i < src.get_layers_lstm().size(); i++)
47 {
48 const std::string name = src.get_module_name(i, "lstm");
49 Parent::m_lstm_idx[name] = i;
50
51 const nn::LSTM& m = src.get_layers_lstm()[i];
52 Parent::m_lstm.emplace_back(typename Parent::params_bilstm_spec_t());
53 typename Parent::params_bilstm_spec_t& layer = Parent::m_lstm.back();
54
56 }
57
58 Parent::m_linear.reserve(src.get_layers_linear().size());
59 for (size_t i = 0; i < src.get_layers_linear().size(); i++)
60 {
61 const std::string name = src.get_module_name(i, "linear");
62 Parent::m_linear_idx[name] = i;
63
64 const nn::Linear& m = src.get_layers_linear()[i];
65 Parent::m_linear.emplace_back(params_linear_t<Eigen::MatrixXf, Eigen::VectorXf>());
66 params_linear_t<Eigen::MatrixXf, Eigen::VectorXf>& layer = Parent::m_linear.back();
67
69 }
70
71 // temp: create exec plan
72 Parent::m_ops.push_back(std::make_shared<op_bilstm_dense_argmax_t>());
73 Parent::m_params.push_back(std::make_shared<typename op_bilstm_dense_argmax_t::params_t>());
74 auto p = std::dynamic_pointer_cast<typename op_bilstm_dense_argmax_t::params_t>(Parent::m_params.back());
75 p->bilstm = Parent::m_lstm[0];
76 for (size_t i = 0; i < Parent::m_linear.size(); ++i)
77 {
78 p->linear.push_back(Parent::m_linear[i]);
79 }
80 p->precompute();
81 Parent::m_wb.resize(1);
82
83 // tags
84 /*cerr << "TAGS:" << endl;
85 for ( const auto& it : src.get_tags() )
86 {
87 cerr << "\t" << it.first << " = " << it.second << endl;
88 }
89 cerr << endl;*/
90}
91
92template <typename AuxScalar>
94 const std::string& fn,
95 std::vector<std::string>& class_names,
96 std::vector<std::vector<std::string>>& classes) {
98 torch::load(src, fn, torch::Device(torch::kCPU));
99
100 // dicts and embeddings
101 Parent::convert_dicts_and_embeddings(src);
102 // classes
103 convert_classes(src.get_classes(), classes);
104 class_names = src.get_class_names();
105}
106
107void convert_classes(const DictsHolder& src, vector<vector<string>>& classes)
108{
109 classes.resize(src.size());
110 for (size_t i = 0; i < classes.size(); ++i)
111 {
112 shared_ptr<StringDict> d = dynamic_pointer_cast<StringDict, DictBase>(src[i]);
113 classes[i].reserve(d->size());
114 for (size_t j = 0; j < d->size(); ++j)
115 {
116 classes[i].push_back(d->get_value(j));
117 }
118 }
119}
120
121} // namespace eigen_impl
122} // namespace tagging
123} // namespace deeplima
124
const std::vector< torch::nn::Linear > & get_layers_linear() const
const std::string get_module_name(size_t idx, const std::string &type) const
const std::vector< torch::nn::LSTM > & get_layers_lstm() const
virtual void convert_classes_from_fn(const std::string &fn, std::vector< std::string > &classes_names, std::vector< std::vector< std::string > > &classes)
const std::vector< std::string > & get_class_names() const
void convert_module_from_torch(const torch::nn::LSTM &src, eigen_impl::params_bilstm_t< M, V > &dst)
void convert_classes(const DictsHolder &src, vector< vector< string > > &classes)
STL namespace.