LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
convert_from_torch.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#include <string>
12#include <regex>
13#include "torch/script.h"
14
15// Dummy function below allows to avoid a crash when compiling with ASAN, as described here:
16// https://github.com/pytorch/pytorch/issues/49460
17// it's unnecessary to invoke this function, just enforce library compiled
18void dummy() {
19 std::regex regstr("Why");
20 std::string s = "Why crashed";
21 std::regex_search(s, regstr);
22}
23
24
25using namespace std;
26using namespace torch;
27using namespace deeplima::convert_from_torch;
28
29namespace deeplima
30{
31namespace eigen_impl
32{
33
35{
36 // dicts and embeddings
37 const vector<nets::embd_descr_t>& embd_descr = src.get_embd_descr();
38 size_t count_embd_uint = 0, count_embd_str = 0;
39 for (size_t i = 0; i < embd_descr.size(); i++)
40 {
41 if (/*0 == embd_descr[i].m_type*/ embd_descr[i].m_name == "raw" )
42 {
43 // There is no dictionary for "raw" embeddings
44 continue;
45 }
46 std::shared_ptr<UInt64Dict> sp_uint64_dict
47 = std::dynamic_pointer_cast<UInt64Dict, DictBase>(src.get_dicts()[i]);
48 std::shared_ptr<StringDict> sp_str_dict
49 = std::dynamic_pointer_cast<StringDict, DictBase>(src.get_dicts()[i]);
50 std::shared_ptr<Char32Dict> sp_char32_dict
51 = std::dynamic_pointer_cast<Char32Dict, DictBase>(src.get_dicts()[i]);
52
53 if (!sp_uint64_dict && !sp_char32_dict && sp_str_dict)
54 {
55 count_embd_str++;
56 }
57 else if ((sp_uint64_dict || sp_char32_dict) && !sp_str_dict)
58 {
59 count_embd_uint++;
60 }
61 else
62 {
63 throw runtime_error("Something wrong with dicts.");
64 }
65 }
66
67 if (count_embd_uint > 0)
68 {
69 m_input_uint_dicts.resize(count_embd_uint);
70 }
71
72 if (count_embd_str > 0)
73 {
74 m_input_str_dicts.resize(count_embd_str);
75 }
76
77 size_t uint_dict_idx = 0, str_dict_idx = 0;
78 for (size_t i = 0; i < embd_descr.size(); i++)
79 {
80 if (/*0 == embd_descr[i].m_type*/ embd_descr[i].m_name == "raw")
81 {
82 // There is no dictionary for "raw" embeddings
83 continue;
84 }
85
86 const string module_name = "embd_" + embd_descr[i].m_name;
87 nn::Embedding m = src.get_module_by_name(module_name);
88
89 std::shared_ptr<UInt64Dict> sp_uint64_dict = std::dynamic_pointer_cast<UInt64Dict, DictBase>(src.get_dicts()[i]);
90 std::shared_ptr<StringDict> sp_str_dict = std::dynamic_pointer_cast<StringDict, DictBase>(src.get_dicts()[i]);
91 std::shared_ptr<Char32Dict> sp_char32_dict = std::dynamic_pointer_cast<Char32Dict, DictBase>(src.get_dicts()[i]);
92
93 assert(sp_uint64_dict || sp_char32_dict || sp_str_dict);
94 if (sp_uint64_dict)
95 {
96 assert(uint_dict_idx < m_input_uint_dicts.size());
97 convert_module_from_torch(m, src.get_dicts()[i], m_input_uint_dicts[uint_dict_idx]);
98 uint_dict_idx++;
99 }
100 else if (sp_char32_dict)
101 {
102 assert(uint_dict_idx < m_input_uint_dicts.size());
103 convert_module_from_torch(m, src.get_dicts()[i], m_input_uint_dicts[uint_dict_idx]);
104 uint_dict_idx++;
105 }
106 else if (sp_str_dict)
107 {
108 assert(str_dict_idx < m_input_str_dicts.size());
109 convert_module_from_torch(m, src.get_dicts()[i], m_input_str_dicts[str_dict_idx]);
110 str_dict_idx++;
111 }
112 }
113}
114
115} // namespace eigen_impl
116
117using namespace eigen_impl;
118
119namespace segmentation
120{
121namespace eigen_impl
122{
123
125{
127 try
128 {
129 torch::load(src, fn, torch::Device(torch::kCPU));
130 }
131 catch (const c10::Error& e)
132 {
133 std::cerr << "Exception while trying to load Torch model file " << fn << std::endl
134 << e.what_without_backtrace();
135 throw std::runtime_error(e.what_without_backtrace());
136 }
137 // ngram_descr
139
140 // dicts and embeddings
142
143 // torch modules
144 Parent::m_lstm.reserve(src.get_layers_lstm().size());
145 for (size_t i = 0; i < src.get_layers_lstm().size(); i++)
146 {
147 const std::string name = src.get_module_name(i, "lstm");
148 Parent::m_lstm_idx[name] = i;
149
150 const nn::LSTM& m = src.get_layers_lstm()[i];
153
155 }
156
157 Parent::m_linear.reserve(src.get_layers_linear().size());
158 for (size_t i = 0; i < src.get_layers_linear().size(); i++)
159 {
160 const std::string name = src.get_module_name(i, "linear");
161 Parent::m_linear_idx[name] = i;
162
163 const nn::Linear& m = src.get_layers_linear()[i];
166
168 }
169
170 // temp: create exec plan
173 auto p = std::dynamic_pointer_cast<typename Op_BiLSTM_Dense_ArgMax<Eigen::MatrixXf, Eigen::VectorXf, float>::params_t>(Parent::m_params.back());
174 p->bilstm = Parent::m_lstm[0];
175 p->linear.push_back(Parent::m_linear[0]);
176 Parent::m_wb.resize(1);
177 p->precompute();
178
179 for (size_t i = 0; i < Parent::m_linear.size(); i++)
180 {
182 // Segmentation has no class-name strings, but record the number of output
183 // classes (the linear layer's row count) so downstream code can tell an
184 // MWT-aware model (> max_segm_tag classes) from a plain tokenizer.
185 Parent::m_output_str_dicts.emplace_back(
186 std::vector<std::string>(size_t(Parent::m_linear[i].weight.rows())));
187 }
188}
189
190} // namespace eigen_impl
191} // namespace segmentation
192} // namespace deeplima
193
std::vector< std::vector< std::shared_ptr< Op_Base::workbench_t > > > m_wb
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)
std::map< std::string, size_t > m_linear_idx
std::map< std::string, size_t > m_lstm_idx
std::vector< params_linear_t< Eigen::MatrixXf, Eigen::VectorXf > > m_linear
std::vector< std::vector< std::string > > m_output_str_dicts
std::vector< params_bilstm_spec_t > m_lstm
std::vector< std::shared_ptr< param_base_t > > m_params
Fusion of several torch modules implemented for inference in Eigen.
const std::vector< embd_descr_t > & get_embd_descr() const
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
torch::nn::Embedding get_module_by_name(const std::string &name) const
const DictsHolder & get_dicts() const
void dummy()
void convert_module_from_torch(const torch::nn::LSTM &src, eigen_impl::params_bilstm_t< M, V > &dst)
STL namespace.