73 unordered_map<char32_t, uint64_t> temp_encoder_dict;
74 unordered_map<char32_t, uint64_t> temp_decoder_dict;
76 for (
const auto& kv : form2lemma )
79 const unordered_map<StringIndex::idx_t, size_t>& dec_forms = kv.second;
82 for (
const char32_t ch : form)
84 temp_encoder_dict[ch] += 1;
87 for (
const auto& dec_kv : dec_forms )
89 const u32string& lemma = str_idx.
get_ustr(dec_kv.first);
90 for (
const char32_t ch : lemma)
92 temp_decoder_dict[ch] += 1;
98 d[0] = std::make_shared<Char32Dict>(0,
EOS,
99 temp_encoder_dict.begin(), temp_encoder_dict.end(),
104 d[1] = std::make_shared<Char32Dict>(0,
EOS,
START,
105 temp_decoder_dict.begin(), temp_decoder_dict.end(),
119 set<morph_model::feat_base_t> fixed_upos, non_fixed_upos;
120 for (
const auto& kv : form2lemma )
123 const unordered_map<StringIndex::idx_t, size_t>& dec_forms = kv.second;
125 assert(!dec_forms.empty());
126 if (dec_forms.size() > 1 || kv.first.m_str_id != dec_forms.begin()->first)
129 non_fixed_upos.insert(upos);
135 <<
" -> " << str_idx.
get_str(dec_forms.begin()->first) << endl;
140 for (
const auto& kv : form2lemma )
143 const unordered_map<StringIndex::idx_t, size_t>& dec_forms = kv.second;
145 if (1 == dec_forms.size() && kv.first.m_str_id == dec_forms.begin()->first)
148 if (non_fixed_upos.end() == non_fixed_upos.find(upos))
150 fixed_upos.insert(upos);
155 for (
const auto& k : fixed_upos )
160 std::cout << std::endl;
174 map<u32string::size_type, uint32_t> n_samples;
175 map<u32string::size_type, u32string::size_type> flen2llen;
176 for (
const auto& kv : form2lemma )
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);
182 if (1 == dec_forms.size() && form.size() < max_len && lemma.size() < max_len)
184 n_samples[form.size()]++;
185 flen2llen[form.size()] = max(flen2llen[form.size()], lemma.size());
189 v_seq_input.reserve(flen2llen.size());
190 v_cat_input.reserve(flen2llen.size());
191 v_gold.reserve(flen2llen.size());
193 for (
const auto& fl : flen2llen )
195 uint32_t form_len = fl.first, max_lemma_len = fl.second;
197 vector<TorchMatrix<int64_t>> cat_input(lang_morph_model.
get_feats_count());
199 for (
size_t feat_idx = 0; feat_idx < lang_morph_model.
get_feats_count(); ++feat_idx)
202 cat_input[feat_idx].init(1, n_samples[form_len]);
206 seq_input.
init(form_len , n_samples[form_len]);
207 gold.
init(max_lemma_len + 1, n_samples[form_len]);
212 int64_t sample_no = 0;
213 for (
const auto& kv : form2lemma )
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);
221 if (1 == kv.second.size() && form_len == form.size() && lemma.size() <= max_lemma_len)
224 for (
size_t char_no = 0; char_no < form.size(); char_no++)
226 char32_t ch = form[char_no];
227 seq_input.
set(char_no, sample_no, p_enc_dict->get_idx(ch));
233 for (; char_no < lemma.size(); char_no++)
235 char32_t ch = lemma[char_no];
236 gold.
set(char_no, sample_no, p_dec_dict->get_idx(ch));
238 for (; char_no < max_lemma_len + 1; char_no++)
240 gold.
set(char_no, sample_no, p_dec_dict->get_idx(
EOS));
246 for (
size_t feat_idx = 0; feat_idx < lang_morph_model.
get_feats_count(); ++feat_idx)
250 cat_input[feat_idx].set(0, sample_no, feat_value);
258 v_seq_input.emplace_back(seq_input);
259 v_cat_input.emplace_back(cat_input);
260 v_gold.emplace_back(gold);
271 Seq2SeqLemmatizer model(
nullptr);
277 model = Seq2SeqLemmatizer();
279 lang_morph_model = model->get_morph_model();
295 throw std::runtime_error(
"train_lemmatization: "+
t2);
303 for (
const auto& word: train_data.
words())
310 const string& form = line.
form();
311 const string& lemma = line.
lemma();
316 form2lemma[
form_morph_t(form_id, feats)][lemma_id] += 1;
321 dh = model->get_dicts();
328 for (
const auto& word: dev_data.
words())
335 const string& form = line.
form();
336 const string& lemma = line.
lemma();
346 if (form2lemma.end() != form2lemma.find(k))
350 dev_form2lemma[k][lemma_id] += 1;
353 auto fixed_upos =
find_fixed(lang_morph_model, form2lemma, str_idx);
356 auto is_fixed_pos = [&fixed_upos =
static_cast<const set<morph_model::feat_base_t>&
>(fixed_upos),
360 return fixed_upos.end() != fixed_upos.find(upos);
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;
372 train_input_seq, train_input_cat, train_gold);
374 dev_input_seq, dev_input_cat, dev_gold);
376 for (
auto& t : train_input_seq)
378 for (
auto& t : dev_input_seq)
381 for (
auto& t : train_gold)
383 for (
auto& t : dev_gold)
386 for (
auto& v : train_input_cat)
389 for (
auto& v : dev_input_cat)
393 if (model.is_empty())
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)
404 dh.emplace_back(std::make_shared<StringDict>(values));
406 string str_fixed_upos =
utils::join(fixed_upos.begin(),
409 (set<morph_model::feat_base_t>::const_reference v){
410 return lang_morph_model.decode_upos_to_str(v);
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);
421 torch::optim::Adam optimizer(model->parameters(),
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);
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)