LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
train_segmentation.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
9
12
14
15#include "char_dict_builder.h"
16#include "char_seq_vectorizer.h"
17#include "train_segmentation.h"
18
19#include <c10/util/Exception.h>
20
21#include <iostream>
22#include <memory>
23#include <vector>
24
25using namespace std;
26
27using namespace deeplima::segmentation;
28using namespace deeplima::segmentation::train;
29using namespace deeplima::segmentation::impl;
30using namespace deeplima::nets;
31using namespace deeplima;
32
33template <class M>
34std::shared_ptr<M> vectorize_gold(const CoNLLU::Annotation& annot, int64_t len, bool eos, bool mwt)
35{
36 auto out = std::make_shared<M>(len, 1);
37
38 const auto& tokens = annot.get_tokens();
39 size_t p = 0;
40 for (size_t i = 0; i < tokens.size(); i++)
41 {
42 const auto& t = tokens[i];
43 while (p < t.m_pos)
44 {
45 out->set(p, 0, segm_tag_t::X);
46 p++;
47 }
48
49 if (t.m_len == 0)
50 {
51 throw std::runtime_error("vectorize_gold length should be 0");
52 }
53
54 // A surface token that is a multiword token (e.g. "du") gets the MWT
55 // variant of its end tag so the model learns to predict, in context, which
56 // tokens should be expanded into sub-words.
57 const bool is_mwt = mwt && t.is_multiword();
58
59 if (t.m_len == 1)
60 {
61 if (eos && t.eos())
62 {
63 out->set(p, 0, is_mwt ? segm_tag_t::S_EOS_MWT : segm_tag_t::S_EOS);
64 }
65 else
66 {
67 out->set(p, 0, is_mwt ? segm_tag_t::S_MWT : segm_tag_t::S);
68 }
69 p++;
70 continue;
71 }
72
73 out->set(p, 0, segm_tag_t::B);
74 p++;
75 while (p < t.m_pos + t.m_len - 1)
76 {
77 out->set(p, 0, segm_tag_t::I);
78 p++;
79 }
80 if (eos && t.eos())
81 {
82 out->set(p, 0, is_mwt ? segm_tag_t::E_EOS_MWT : segm_tag_t::E_EOS);
83 }
84 else
85 {
86 out->set(p, 0, is_mwt ? segm_tag_t::E_MWT : segm_tag_t::E);
87 }
88
89 p++;
90 }
91
92 return out;
93}
94
98
101 int gpuid)
102{
103 std::vector<ngram_descr_t> ngram_descr = {
106 { -1, 3, ngram_descr_t::char_ngram },
109 { -1, 3, ngram_descr_t::type_ngram },
110 { 0, 1, ngram_descr_t::script_ngram } };
111
112 const CoNLLU::Document& train_doc = tb.get_doc("train");
113 uint64_t train_char_counter = 0;
114
115 DictionaryBuilder dict_builder(ngram_descr);
116 auto dicts = dict_builder.process(train_doc.get_original_text(), 100, train_char_counter);
117
118 if (train_char_counter >= std::numeric_limits<int64_t>::max())
119 {
120 throw std::overflow_error("Too much characters in training set.");
121 }
122
123 Utf8CharSeqToTorchMatrix vectorizer(ngram_descr);
124 vectorizer.set_dicts(dicts);
125 auto train_input = vectorizer.process(train_doc.get_original_text(),
126 (int64_t)train_char_counter);
127
128 const auto& dev_doc = tb.get_doc("dev");
129 auto dev_input = vectorizer.process(dev_doc.get_original_text(), dev_doc.get_text().size() + 1);
130
131 auto train_gold = vectorize_gold<TorchMatrix<int64_t>>(tb.get_annot("train"), (int64_t)train_char_counter,
132 params.train_ss, params.train_mwt);
133
134 auto dev_gold = vectorize_gold<TorchMatrix<int64_t>>(tb.get_annot("dev"), dev_doc.get_text().size() + 1,
135 params.train_ss, params.train_mwt);
136
137 std::vector<embd_descr_t> embd_descr = { { "char1gram", 2 }, { "char2gram", 3 }, { "char3gram", 4 },
138 { "class1gram", 2 }, { "class2gram", 2 }, { "class3gram", 2 },
139 { "scriptchange", 1 } };
140
141 std::vector<rnn_descr_t> rnn_descr = { rnn_descr_t( params.m_rnn_hidden_dim ) };
142
143
144 // Class count: 5 (tokenization only), 7 (+ sentence segmentation), or 11
145 // (+ the four multiword-token end-tag variants). MWT implies sentence seg.
146 const uint32_t num_classes = params.train_mwt ? 11 : (params.train_ss ? 7 : 5);
147
148 BiRnnClassifierForSegmentation model(std::move(dicts),
149 ngram_descr,
150 embd_descr,
151 rnn_descr,
152 "tokens",
153 num_classes,
154 params.m_input_dropout_prob);
155
156 torch::optim::Adam optimizer(model->parameters(),
157 torch::optim::AdamOptions(params.m_learning_rate)
158 .weight_decay(params.m_weight_decay)
159 .betas({params.m_beta_one, params.m_beta_two}));
160
161 std::string dev = "cpu";
162 if (gpuid >= 0)
163 {
164 std::ostringstream oss;
165 oss << "cuda:" << gpuid;
166 dev = oss.str();
167 }
168 torch::Device device(dev);
169
170 train_input->to(device);
171 train_gold->to(device);
172
173 dev_input->to(device);
174 dev_gold->to(device);
175
176 model->to(device);
177
178 try
179 {
180 model->train(params.m_max_epochs, params.m_batch_size, params.m_sequence_length, { "tokens" },
181 *(train_input.get()), *(train_gold.get()),
182 *(dev_input.get()), *(dev_gold.get()),
183 optimizer, params.m_output_model_name, device);
184 }
185 catch (const c10::Error& e)
186 {
187 std::cerr << "Exception in model training: " << e.what() << std::endl;
188 return 1;
189 }
190
191 return 0;
192}
193
const std::vector< token_t > & get_tokens() const
Definition treebank.h:147
const std::string & get_original_text() const
Definition treebank.h:72
const Document & get_doc(const std::string &name) const
Definition treebank.h:479
const Annotation & get_annot(const std::string &name) const
Definition treebank.h:489
STL namespace.
DictionaryBuilderImpl< CharNgramEncoder< Utf8Reader<> > > DictionaryBuilder
int train_segmentation_model(const CoNLLU::Treebank &tb, deeplima::segmentation::train::train_params_segmentation_t &params, int gpuid)
std::shared_ptr< M > vectorize_gold(const CoNLLU::Annotation &annot, int64_t len, bool eos, bool mwt)
CharSeqVectorizerImpl< CharNgramEncoder< Utf8Reader<> >, TorchMatrix< int64_t >, DictHolderAdapter< UInt64Dict, TorchMatrix< int64_t > > > Utf8CharSeqToTorchMatrix