LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
train_lemmatization.cpp
Go to the documentation of this file.
1// Copyright 2002-2022 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 <unordered_map>
8#include <algorithm>
9
10#include "conllu/treebank.h"
11
12#include "train_lemmatization.h"
13
15
17#include "static_graph/dict.h"
20
21#include <boost/functional/hash.hpp>
24
25using namespace std;
26using namespace deeplima::morph_model;
27using namespace deeplima::nets;
28
29namespace deeplima
30{
31namespace lemmatization
32{
33namespace train
34{
35
48
49inline bool operator==(const form_morph_t& a, const form_morph_t& b)
50{
51 return (a.m_str_id == b.m_str_id) && (a.m_feats == b.m_feats);
52}
53
55{
56 std::size_t operator()(form_morph_t const& a) const noexcept
57 {
58 size_t h = std::hash<StringIndex::idx_t>{}(a.m_str_id);
59 boost::hash_combine(h, a.m_feats.toBaseType());
60 return h;
61 }
62};
63
64typedef unordered_map<form_morph_t, unordered_map<StringIndex::idx_t, size_t>, form_morph_hash> form2lemma_t;
65
66const char32_t START = 0x10FFFE;
67const char32_t EOS = 0x10FFFF; // End of Supplementary Private Use Area-B
68
70{
72
73 unordered_map<char32_t, uint64_t> temp_encoder_dict;
74 unordered_map<char32_t, uint64_t> temp_decoder_dict;
75
76 for ( const auto& kv : form2lemma )
77 {
78 const form_morph_t& enc_form = kv.first;
79 const unordered_map<StringIndex::idx_t, size_t>& dec_forms = kv.second;
80
81 const u32string& form = str_idx.get_ustr(enc_form.m_str_id);
82 for (const char32_t ch : form)
83 {
84 temp_encoder_dict[ch] += 1;
85 }
86
87 for ( const auto& dec_kv : dec_forms )
88 {
89 const u32string& lemma = str_idx.get_ustr(dec_kv.first);
90 for (const char32_t ch : lemma)
91 {
92 temp_decoder_dict[ch] += 1;
93 }
94 }
95 }
96
97 d.resize(2);
98 d[0] = std::make_shared<Char32Dict>(0, EOS,
99 temp_encoder_dict.begin(), temp_encoder_dict.end(),
100 [](uint64_t c) {
101 return c > 1;
102 });
103
104 d[1] = std::make_shared<Char32Dict>(0, EOS, START,
105 temp_decoder_dict.begin(), temp_decoder_dict.end(),
106 [](uint64_t c) {
107 return c > 1;
108 });
109
110 return d;
111}
112
113// TODO add use shallow fixed classes
114set<morph_model::feat_base_t> find_fixed(const morph_model::morph_model_t& lang_morph_model,
115 const form2lemma_t& form2lemma,
116 /*const*/ StringIndex& str_idx // TODO: str_idx must be const here
117 )
118{
119 set<morph_model::feat_base_t> fixed_upos, non_fixed_upos;
120 for ( const auto& kv : form2lemma )
121 {
122 const form_morph_t& form = kv.first;
123 const unordered_map<StringIndex::idx_t, size_t>& dec_forms = kv.second;
124
125 assert(!dec_forms.empty());
126 if (dec_forms.size() > 1 || kv.first.m_str_id != dec_forms.begin()->first)
127 {
128 morph_model::feat_base_t upos = lang_morph_model.decode_upos(form.m_feats);
129 non_fixed_upos.insert(upos);
130 if (lang_morph_model.decode_upos_to_str(morph_model::morph_feats_t(upos)) == "PUNCT"
131 || lang_morph_model.decode_upos_to_str(morph_model::morph_feats_t(upos)) == "X"
132 || lang_morph_model.decode_upos_to_str(morph_model::morph_feats_t(upos)) == "SYM")
133 {
134 cout << str_idx.get_str(form.m_str_id)
135 << " -> " << str_idx.get_str(dec_forms.begin()->first) << endl;
136 }
137 }
138 }
139
140 for ( const auto& kv : form2lemma )
141 {
142 const form_morph_t& form = kv.first;
143 const unordered_map<StringIndex::idx_t, size_t>& dec_forms = kv.second;
144
145 if (1 == dec_forms.size() && kv.first.m_str_id == dec_forms.begin()->first)
146 {
147 morph_model::feat_base_t upos = lang_morph_model.decode_upos(form.m_feats);
148 if (non_fixed_upos.end() == non_fixed_upos.find(upos))
149 {
150 fixed_upos.insert(upos);
151 }
152 }
153 }
154
155 for ( const auto& k : fixed_upos )
156 {
157 std::cout << "Fixed UPOS: " << lang_morph_model.decode_upos_to_str(morph_model::morph_feats_t(k))
158 << std::endl;
159 }
160 std::cout << std::endl;
161
162 return fixed_upos;
163}
164
166 const form2lemma_t& form2lemma,
167 /*const*/ StringIndex& str_idx, // TODO: str_idx must be const here
168 const DictsHolder dh,
169 uint32_t max_len,
170 vector<TorchMatrix<int64_t>>& v_seq_input,
171 vector<vector<TorchMatrix<int64_t>>>& v_cat_input,
172 vector<TorchMatrix<int64_t>>& v_gold)
173{
174 map<u32string::size_type, uint32_t> n_samples; // for each form len
175 map<u32string::size_type, u32string::size_type> flen2llen; // form len -> max lemma len;
176 for ( const auto& kv : form2lemma )
177 {
178 const u32string& form = str_idx.get_ustr(kv.first.m_str_id);
179 const unordered_map<StringIndex::idx_t, size_t>& dec_forms = kv.second;
180 const u32string& lemma = str_idx.get_ustr(dec_forms.begin()->first);
181
182 if (1 == dec_forms.size() && form.size() < max_len && lemma.size() < max_len)
183 {
184 n_samples[form.size()]++;
185 flen2llen[form.size()] = max(flen2llen[form.size()], lemma.size());
186 }
187 }
188
189 v_seq_input.reserve(flen2llen.size());
190 v_cat_input.reserve(flen2llen.size());
191 v_gold.reserve(flen2llen.size());
192
193 for ( const auto& fl : flen2llen )
194 {
195 uint32_t form_len = fl.first, max_lemma_len = fl.second;
196 TorchMatrix<int64_t> seq_input, gold;
197 vector<TorchMatrix<int64_t>> cat_input(lang_morph_model.get_feats_count());
198
199 for (size_t feat_idx = 0; feat_idx < lang_morph_model.get_feats_count(); ++feat_idx)
200 {
201 // 1 x n_samples
202 cat_input[feat_idx].init(1, n_samples[form_len]);
203 }
204
205 // max_len x n_samples (x 1 (one feature))
206 seq_input.init(form_len /*+ 1*/, n_samples[form_len]); // max_len
207 gold.init(max_lemma_len + 1, n_samples[form_len]); // max_len + STOP
208
209 Char32Dict* p_enc_dict = dynamic_cast<Char32Dict*>(dh[0].get());
210 Char32Dict* p_dec_dict = dynamic_cast<Char32Dict*>(dh[1].get());
211
212 int64_t sample_no = 0;
213 for ( const auto& kv : form2lemma )
214 {
215 const form_morph_t& enc_form = kv.first;
216 const u32string& form = str_idx.get_ustr(enc_form.m_str_id);
217
218 const unordered_map<StringIndex::idx_t, size_t>& dec_forms = kv.second;
219 const u32string& lemma = str_idx.get_ustr(dec_forms.begin()->first);
220
221 if (1 == kv.second.size() && form_len == form.size() && lemma.size() <= max_lemma_len)
222 {
223 // input sequence
224 for (size_t char_no = 0; char_no < form.size(); char_no++)
225 {
226 char32_t ch = form[char_no];
227 seq_input.set(char_no, sample_no, p_enc_dict->get_idx(ch));
228 }
229 //seq_input.set(form.size(), sample_no, p_enc_dict->get_idx(EOS));
230
231 // gold output sequence
232 size_t char_no = 0;
233 for (; char_no < lemma.size(); char_no++)
234 {
235 char32_t ch = lemma[char_no];
236 gold.set(char_no, sample_no, p_dec_dict->get_idx(ch));
237 }
238 for (; char_no < max_lemma_len + 1; char_no++)
239 {
240 gold.set(char_no, sample_no, p_dec_dict->get_idx(EOS));
241 }
242
243 // input categories
244 /*std::cerr << str_idx.get_str(enc_form.m_str_id) << " "
245 << lang_morph_model.to_string(enc_form.m_feats) << std::endl;*/
246 for (size_t feat_idx = 0; feat_idx < lang_morph_model.get_feats_count(); ++feat_idx)
247 {
248 const auto& feat_value = lang_morph_model.decode_feat(enc_form.m_feats, feat_idx);
249 //std::cerr << " " << feat_value;
250 cat_input[feat_idx].set(0, sample_no, feat_value);
251 }
252 //std::cerr << std::endl;
253
254 sample_no++;
255 }
256 }
257
258 v_seq_input.emplace_back(seq_input);
259 v_cat_input.emplace_back(cat_input);
260 v_gold.emplace_back(gold);
261 }
262}
263
265{
266 // Load data sets
267 CoNLLU::Annotation train_data, dev_data;
268 train_data.load(params.m_train_set_fn);
269 dev_data.load(params.m_dev_set_fn);
270
271 Seq2SeqLemmatizer model(nullptr);
272 morph_model::morph_model_t lang_morph_model;
273 DictsHolder dh;
274
275 if (params.m_input_model_name.size() > 0)
276 {
277 model = Seq2SeqLemmatizer();
278 model->load(params.m_input_model_name);
279 lang_morph_model = model->get_morph_model();
280 }
281 else
282 {
283 lang_morph_model = morph_model::morph_model_builder::build(train_data, dev_data);
284 }
285
286 {
287 // Serialization test
288 string t1 = lang_morph_model.to_string();
289 // cerr << t1 << endl;
291 string t2 = m2.to_string();
292 if (t1 != t2)
293 {
294 // cerr << t2 << endl;
295 throw std::runtime_error("train_lemmatization: "+t2);
296 }
297 // End of serialization test
298 }
299
300 StringIndex str_idx;
301
302 form2lemma_t form2lemma, dev_form2lemma;
303 for (const auto& word: train_data.words())
304 {
305 const CoNLLU::CoNLLULine& line = train_data.get_line(word.m_line_idx);
306 if (line.is_foreign() || line.is_typo())
307 {
308 continue;
309 }
310 const string& form = line.form();
311 const string& lemma = line.lemma();
312 StringIndex::idx_t form_id = str_idx.get_idx(form);
313 StringIndex::idx_t lemma_id = str_idx.get_idx(lemma);
314
315 morph_model::morph_feats_t feats = lang_morph_model.convert(line.upos(), line.feats());
316 form2lemma[form_morph_t(form_id, feats)][lemma_id] += 1;
317 }
318
319 if (params.m_input_model_name.size() > 0)
320 {
321 dh = model->get_dicts();
322 }
323 else
324 {
325 dh = build_char_dicts(form2lemma, str_idx);
326 }
327
328 for (const auto& word: dev_data.words())
329 {
330 const CoNLLU::CoNLLULine& line = dev_data.get_line(word.m_line_idx);
331 if (line.is_foreign() || line.is_typo())
332 {
333 continue;
334 }
335 const string& form = line.form();
336 const string& lemma = line.lemma();
337 StringIndex::idx_t form_id = str_idx.get_idx(form);
338 StringIndex::idx_t lemma_id = str_idx.get_idx(lemma);
339 // store form in the strings index
340 str_idx.get_ustr(form_id);
341 // store lemma in the strings index
342 str_idx.get_ustr(lemma_id);
343
344 auto feats = lang_morph_model.convert(line.upos(), line.feats());
345 form_morph_t k(form_id, feats);
346 if (form2lemma.end() != form2lemma.find(k))
347 {
348 continue;
349 }
350 dev_form2lemma[k][lemma_id] += 1;
351 }
352
353 auto fixed_upos = find_fixed(lang_morph_model, form2lemma, str_idx);
354
355 // Remove fixed UPOS
356 auto is_fixed_pos = [&fixed_upos = static_cast<const set<morph_model::feat_base_t>&>(fixed_upos),
357 &lang_morph_model = static_cast<morph_model::morph_model_t&>(lang_morph_model)](form2lemma_t::const_reference v){
358 const form_morph_t& form = v.first;
359 feat_base_t upos = lang_morph_model.decode_upos(form.m_feats);
360 return fixed_upos.end() != fixed_upos.find(upos);
361 };
362
363 cerr << "Fixed tokens removed from train set: " << utils::backport::erase_if(form2lemma, is_fixed_pos) << endl;
364 cerr << "Fixed tokens removed from dev set: " << utils::backport::erase_if(dev_form2lemma, is_fixed_pos) << endl;
365
366 torch::Device device(params.m_device_string);
367
368 vector<TorchMatrix<int64_t>> train_input_seq, train_gold;
369 vector<TorchMatrix<int64_t>> dev_input_seq, dev_gold;
370 vector<vector<TorchMatrix<int64_t>>> train_input_cat, dev_input_cat;
371 vectorize_dataset(lang_morph_model, form2lemma, str_idx, dh, 31,
372 train_input_seq, train_input_cat, train_gold);
373 vectorize_dataset(lang_morph_model, dev_form2lemma, str_idx, dh, 31,
374 dev_input_seq, dev_input_cat, dev_gold);
375
376 for (auto& t : train_input_seq)
377 t.to(device);
378 for (auto& t : dev_input_seq)
379 t.to(device);
380
381 for (auto& t : train_gold)
382 t.to(device);
383 for (auto& t : dev_gold)
384 t.to(device);
385
386 for (auto& v : train_input_cat)
387 for (auto& t : v)
388 t.to(device);
389 for (auto& v : dev_input_cat)
390 for (auto& t : v)
391 t.to(device);
392
393 if (model.is_empty())
394 {
395 vector<embd_descr_t> encoder_embd_descr = { embd_descr_t("enc_chars", params.m_encoder_embd_dim) };
396 vector<embd_descr_t> decoder_embd_descr = { embd_descr_t("dec_chars", params.m_decoder_embd_dim) };
397 vector<rnn_descr_t> encoder_rnn_descr = { rnn_descr_t(params.m_encoder_rnn_hidden_dim) };
398 vector<rnn_descr_t> decoder_rnn_descr = { rnn_descr_t(params.m_decoder_rnn_hidden_dim) };
399 vector<embd_descr_t> cat_embd_descr;
400 for (size_t feat_idx = 0; feat_idx < lang_morph_model.get_feats_count(); ++feat_idx)
401 {
402 cat_embd_descr.emplace_back(embd_descr_t("cat_" + lang_morph_model.get_feat_name(feat_idx), 8));
403 const auto& values = lang_morph_model.get_feat_vec_ref(feat_idx);
404 dh.emplace_back(std::make_shared<StringDict>(values));
405 }
406 string str_fixed_upos = utils::join(fixed_upos.begin(),
407 fixed_upos.end(),
408 [&lang_morph_model = static_cast<morph_model::morph_model_t&>(lang_morph_model)]
409 (set<morph_model::feat_base_t>::const_reference v){
410 return lang_morph_model.decode_upos_to_str(v);
411 });
412
413 model = Seq2SeqLemmatizer(std::move(dh), lang_morph_model,
414 encoder_embd_descr, encoder_rnn_descr,
415 decoder_embd_descr, decoder_rnn_descr,
416 cat_embd_descr, str_fixed_upos);
417 }
418
419 // std::cerr << model->get_script() << std::endl;
420
421 torch::optim::Adam optimizer(model->parameters(),
422 torch::optim::AdamOptions(params.m_learning_rate)
423 .weight_decay(params.m_weight_decay));
424
425 model->to(device);
426 for(auto& matrix: train_input_seq) matrix.to(device);
427 for(auto& matrix: train_gold) matrix.to(device);
428 for(auto& vec: train_input_cat) for(auto& matrix:vec) matrix.to(device);
429 model->train(params, train_input_seq, train_input_cat, train_gold,
430 dev_input_seq, dev_input_cat, dev_gold, optimizer, device);
431
432 return 0;
433}
434
435} // namespace train
436} // namespace lemmatization
437} // namespace deeplima
438
const std::vector< word_t > & words() const
Definition treebank.h:168
void load(const std::string &fn)
Definition treebank.cpp:91
const CoNLLULine & get_line(size_t line_idx) const
Definition treebank.h:184
const std::map< std::string, std::set< std::string > > & feats() const
Definition line.h:178
const std::string & form() const
Definition line.h:153
const std::string & lemma() const
Definition line.h:158
bool is_foreign() const
Definition line.h:207
const std::string & upos() const
Definition line.h:163
const S & get_str(const idx_t idx) const
Definition str_index.h:59
idx_t get_idx(const char *p, size_t len)
Definition str_index.h:24
const U & get_ustr(const idx_t idx)
accessor with side effect.
Definition str_index.h:76
void init(int64_t max_time, int64_t max_feat)
void set(int64_t time, uint64_t feat, T value)
Encoding on one 64 bits integer of the set of morphological features for one token.
Definition morph_model.h:36
static morph_model_t build(const CoNLLU::Annotation &annotation1, const CoNLLU::Annotation &annotation2)
Helper class for morphology data (upos, features) binarization.
Definition morph_model.h:93
const std::string & decode_upos_to_str(const morph_feats_t &feats) const
morph_feats_t convert(const std::string &upos, const std::map< std::string, std::set< std::string > > &feats) const
const std::vector< std::string > & get_feat_vec_ref(size_t feat_id) const
const std::string & get_feat_name(size_t feat_id) const
feat_base_t decode_feat(const morph_feats_t &feats, size_t feat_id) const
feat_base_t decode_upos(const morph_feats_t &feats) const
DictsHolder build_char_dicts(const form2lemma_t &form2lemma, StringIndex &str_idx)
int train_lemmatization(const train_params_lemmatization_t &params)
bool operator==(const form_morph_t &a, const form_morph_t &b)
set< morph_model::feat_base_t > find_fixed(const morph_model::morph_model_t &lang_morph_model, const form2lemma_t &form2lemma, StringIndex &str_idx)
void vectorize_dataset(const morph_model::morph_model_t &lang_morph_model, const form2lemma_t &form2lemma, StringIndex &str_idx, const DictsHolder dh, uint32_t max_len, vector< TorchMatrix< int64_t > > &v_seq_input, vector< vector< TorchMatrix< int64_t > > > &v_cat_input, vector< TorchMatrix< int64_t > > &v_gold)
unordered_map< form_morph_t, unordered_map< StringIndex::idx_t, size_t >, form_morph_hash > form2lemma_t
std::unordered_map< K, V, H, A >::size_type erase_if(std::unordered_map< K, V, H, A > &c, P pred)
Definition backport.h:20
std::string join(InputIt begin, InputIt end, Pred f)
Dict< char32_t > Char32Dict
Definition dict.h:276
STL namespace.
std::size_t operator()(form_morph_t const &a) const noexcept
form_morph_t(StringIndex::idx_t str_id, morph_feats_t feats)