LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
train_tag.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
6#include <string>
7#include <algorithm>
8#include <cctype>
9
10#include "conllu/treebank.h"
13
17
18#include "../model/birnn_classifier_for_tag.h"
21
22#include "train_tag.h"
23
24using namespace std;
25using namespace deeplima::nets;
26
27namespace deeplima
28{
29namespace tagging
30{
31namespace train
32{
33
34template <typename Token>
36{
37public:
38 UPosFeatExtractor(const std::string&)
39 {
40 }
41
42 inline static bool needs_preprocessing()
43 {
44 return false;
45 }
46
47 inline void preprocess(const Token&)
48 {
49 }
50
51 inline static std::string feat_value(const Token& token, size_t feat_no)
52 {
53 switch (feat_no)
54 {
55 case 0:
56 return token.upos();
57 default:
58 throw runtime_error("Unknown feature id.");
59 }
60 }
61
62 inline static size_t size()
63 {
64 return 1;
65 }
66};
67
68template <class M, class FeatExtractor>
69std::shared_ptr<M> vectorize_gold(const CoNLLU::Annotation& annot, const FeatExtractor& fe, const DictsHolder& tag_dh)
70{
71 CoNLLU::WordLevelAdapter src(&annot);
72 int64_t len = 0;
74 while (src.end() != i)
75 {
76 if ((*i).is_word())
77 {
78 len++;
79 }
80 i++;
81 }
82
83 auto out = std::make_shared<M>(len, tag_dh.size());
84
86 uint64_t current_timepoint = 0;
87 while (src.end() != it)
88 {
89 while(!(*it).is_word() && src.end() != it)
90 {
91 it++;
92 }
93 if (src.end() == it)
94 {
95 break;
96 }
97
98 for (size_t feat_idx = 0; feat_idx < tag_dh.size(); ++feat_idx)
99 {
100 const std::string& feat_val = fe.feat_value(*it, feat_idx);
101 StringDict::key_t idx = dynamic_pointer_cast<StringDict>(tag_dh[feat_idx])->get_idx(feat_val);
102 out->set(current_timepoint, feat_idx, idx);
103 }
104
105 current_timepoint++;
106 if (current_timepoint == std::numeric_limits<uint64_t>::max())
107 {
108 throw std::overflow_error("Too much words in the dataset.");
109 }
110
111 it++;
112 }
113
114 return out;
115}
116
126
128
130{
131 // Load data sets
132 CoNLLU::Annotation train_data, dev_data;
133 train_data.load(params.m_train_set_fn);
134 dev_data.load(params.m_dev_set_fn);
135
136 // Find target classes
137 TagDictBuilderFromCoNLLU tag_dict_builder;
139 = tag_dict_builder.preprocess(CoNLLU::WordLevelAdapter(&train_data),
140 params.m_tasks_string);
141 DictsHolder tag_dh = tag_dict_builder.process(CoNLLU::WordLevelAdapter(&train_data),
142 feat_extractor, 0, "");
143
144 DictsHolder dh;
145 BiRnnClassifierForNer model(nullptr);
146
147 if (params.m_input_model_name.size() > 0)
148 {
149 model = BiRnnClassifierForNer();
150 model->load(params.m_input_model_name);
151 }
152 else
153 {
154 if (params.m_trainable_embeddings_dim > 0)
155 {
156 WordDictBuilderFromCoNLLU dict_builder;
157
158 dh = dict_builder.process(CoNLLU::WordLevelAdapter(&train_data),
160 dh.erase(dh.begin()); // cased forms aren't needed
161 }
162 }
163
164 // Input features
165 vector<CoNLLUToTorchMatrix::feature_descr_t> feat_descr;
166
167 std::shared_ptr<FastTextVectorizerToTorchMatrix> p_embd;
168 if (params.m_embeddings_fn.size() > 0)
169 {
170 try
171 {
172 p_embd = std::make_shared<FastTextVectorizerToTorchMatrix>(params.m_embeddings_fn);
173 assert(nullptr != p_embd.get());
174 feat_descr.push_back({ CoNLLUToTorchMatrix::str_feature, "form", p_embd });
175 }
176 catch (const exception& e)
177 {
178 cerr << e.what() << endl;
179 return -1;
180 }
181 catch (...)
182 {
183 cerr << "Something wrong happened while loading \""
184 << params.m_embeddings_fn << "\"" << endl;
185 return -1;
186 }
187 }
188
189 shared_ptr<DirectDict<TorchMatrix<float>>> p_eos;
190 if (params.m_use_eos)
191 {
192 if (string::npos != params.m_tasks_string.find("eos"))
193 {
194 throw std::invalid_argument("Can't use EOS as both input and output");
195 }
196
197 p_eos = std::make_shared<DirectDict<TorchMatrix<float>>>(2);
198 assert(nullptr != p_embd);
199 feat_descr.push_back({ CoNLLUToTorchMatrix::int_feature, "eos", p_eos });
200 }
201
202 vector<CoNLLUToTorchMatrix::embeddable_feature_descr_t> embd_feat_descr;
203 if (params.m_input_model_name.size() > 0)
204 {
205 embd_feat_descr.push_back({
206 CoNLLUToTorchMatrix::str_feature,
207 "lc(form)",
208 int(model->get_embd_descr()[0].m_dim),
209 model->get_dicts()[0]
210 });
211 }
212 else if (params.m_trainable_embeddings_dim > 0)
213 {
214 assert(dh[0]->size() > 1);
215 embd_feat_descr.push_back({
216 CoNLLUToTorchMatrix::str_feature,
217 "lc(form)",
218 int(params.m_trainable_embeddings_dim),
219 dh[0]
220 });
221 }
222
223 assert(feat_descr.size() > 0 || embd_feat_descr.size() > 0);
224
225 CoNLLUToTorchMatrix vectorizer(feat_descr, embd_feat_descr);
226
227 CoNLLUToTorchMatrix::vectorization_t train_input
228 = vectorizer.process(CoNLLU::WordLevelAdapter(&train_data));
229 CoNLLUToTorchMatrix::vectorization_t dev_input
230 = vectorizer.process(CoNLLU::WordLevelAdapter(&dev_data));
231
232 shared_ptr<TorchMatrix<int64_t>> train_gold
233 = vectorize_gold<TorchMatrix<int64_t>, ConlluFeatExtractor<CoNLLU::WordLevelAdapter::token_t>>(train_data,
234 feat_extractor,
235 tag_dh);
236
237 shared_ptr<TorchMatrix<int64_t>> dev_gold
238 = vectorize_gold<TorchMatrix<int64_t>, ConlluFeatExtractor<CoNLLU::WordLevelAdapter::token_t>>(dev_data,
239 feat_extractor,
240 tag_dh);
241
242 if (train_input.first && train_input.second && train_input.first->size() != train_input.second->size())
243 {
244 std::cerr << "ERROR: train set length missmatch: "
245 << train_input.first->size()
246 << " " << train_input.second->size() << std::endl;
247 throw runtime_error("Train set length missmatch");
248 }
249
250 if ((train_input.first && train_input.first->size() != train_gold->size())
251 || (train_input.second && train_input.second->size() != train_gold->size()))
252 {
253 std::cerr << "ERROR: train set length (input != gold): "
254 << (train_input.first ? train_input.first->size() : train_input.second->size())
255 << " " << train_gold->size() << std::endl;
256 throw runtime_error("Train set length missmatch (input != gold)");
257 }
258
259 if (params.m_input_model_name.size() == 0)
260 {
261 vector<embd_descr_t> embd_descr = vectorizer.get_embd_descr();
262 embd_descr.emplace_back("raw", train_input.second->get_tensor().size(1), 0);
263 vector<rnn_descr_t> rnn_descr = { rnn_descr_t(params.m_rnn_hidden_dim) /*, rnn_descr_t(32) */ };
264 model = BiRnnClassifierForNer(std::move(dh),
265 embd_descr,
266 rnn_descr,
267 feat_extractor.feats(),
268 std::move(tag_dh),
269 boost::filesystem::path(params.m_embeddings_fn).stem().string(),
270 params);
271 }
272
273 // cerr << model->get_script() << endl;
274
275 torch::Device device(params.m_device_string);
276
277 train_input.first->to(device);
278 train_input.second->to(device);
279 train_gold->to(device);
280
281 dev_input.first->to(device);
282 dev_input.second->to(device);
283 dev_gold->to(device);
284
285 model->to(device);
286
287 double min_perf = 0;
288
289 for (const string& opt_name : utils::split(params.m_optimizers, ','))
290 {
291 shared_ptr<torch::optim::Optimizer> optimizer;
292
293 if (opt_name == "adam")
294 {
295 optimizer = make_shared<torch::optim::Adam>(model->parameters(),
296 torch::optim::AdamOptions(params.m_learning_rate)
297 .weight_decay(params.m_weight_decay)
298 .betas({params.m_beta_one, params.m_beta_two}));
299 }
300 else if (opt_name == "sgd")
301 {
302 optimizer = make_shared<torch::optim::SGD>(model->parameters(),
303 torch::optim::SGDOptions(params.m_learning_rate * 1000)
304 .weight_decay(params.m_weight_decay));
305 }
306 else
307 {
308 throw runtime_error("Unknown optimizer: " + opt_name);
309 }
310
311 model->train(params,
312 feat_extractor.feats(),
313 *(train_input.first.get()), *(train_input.second.get()), *(train_gold.get()),
314 *(dev_input.first.get()), *(dev_input.second.get()), *(dev_gold.get()),
315 *optimizer, min_perf, device);
316
317 std::cerr << "train_tag: Optimizer " << opt_name << " stopped at " << min_perf << std::endl;
318 }
319
320 return 0;
321}
322
323} // namespace train
324} // namespace tagging
325} // namespace deeplima
326
void load(const std::string &fn)
Definition treebank.cpp:91
virtual const_iterator begin() const
Definition treebank.h:403
virtual const_iterator end() const
Definition treebank.h:409
void preprocess(const Token &token)
std::vector< std::string > feats() const
static std::string feat_value(const Token &token, size_t feat_no)
Definition train_tag.cpp:51
WordDictBuilderImpl< CoNLLU::WordLevelAdapter, deeplima::TokenStrFeatExtractor< CoNLLU::WordLevelAdapter::token_t > > WordDictBuilderFromCoNLLU
int train_entity_tagger(const train_params_tagging_t &params)
WordDictBuilderImpl< CoNLLU::WordLevelAdapter, ConlluFeatExtractor< CoNLLU::WordLevelAdapter::token_t > > TagDictBuilderFromCoNLLU
std::shared_ptr< M > vectorize_gold(const CoNLLU::Annotation &annot, const FeatExtractor &fe, const DictsHolder &tag_dh)
Definition train_tag.cpp:69
FastTextVectorizer< TorchMatrix< float > > FastTextVectorizerToTorchMatrix
WordSeqVectorizerImpl< CoNLLU::WordLevelAdapter, deeplima::TokenStrFeatExtractor< CoNLLU::WordLevelAdapter::token_t >, deeplima::TokenUIntFeatExtractor< CoNLLU::WordLevelAdapter::token_t >, TorchMatrix< int64_t >, TorchMatrix< float > > CoNLLUToTorchMatrix
std::vector< std::string > split(const std::string &str, char delim)
STL namespace.