LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
train_graph_dp.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 "../model/birnn_and_deep_biaffine_attention.h"
7
12
13#include "train_graph_dp.h"
15
16using namespace std;
17using namespace deeplima::nets;
18
19namespace deeplima
20{
21namespace graph_dp
22{
23namespace train
24{
25
28
30
32{
33 // Load data sets
34 CoNLLU::Annotation train_data, dev_data;
35 train_data.load(params.m_train_set_fn);
36 dev_data.load(params.m_dev_set_fn);
37
38 shared_ptr<FastTextVectorizerToTorchMatrix> p_embd;
39 if (params.m_embeddings_fn.size() > 0)
40 {
41 try
42 {
43 p_embd = std::make_shared<FastTextVectorizerToTorchMatrix>(params.m_embeddings_fn);
44 assert(nullptr != p_embd.get());
45 //feat_descr.push_back({ CoNLLUToTorchMatrix::str_feature, "form", p_embd.get() });
46 }
47 catch (const exception& e)
48 {
49 cerr << e.what() << endl;
50 return -1;
51 }
52 catch (...)
53 {
54 cerr << "Something wrong happened while loading \""
55 << params.m_embeddings_fn << "\"" << endl;
56 return -1;
57 }
58 }
59
60 TagDictBuilderFromCoNLLU tag_dict_builder;
61 // TODO retrieve tagger tasks list from model or config
62 // Currently the list below (form,upos…) is hard-coded while the tagging model can be trained with various tasks.
63 // This list should be saved somewhere and retrieved here later on
65 = tag_dict_builder.preprocess(CoNLLU::WordLevelAdapter(&train_data),
66 "*form,upos,feats,xpos,-Typo,-Foreign");
67 DictsHolder tag_dh = tag_dict_builder.process(CoNLLU::WordLevelAdapter(&train_data),
68 feat_extractor, 0, "");
69
70 DictsHolder dh;
71
72 // Build the deprel (syntactic relation label) dictionary from the training
73 // data. Shared with the dev set so class ids are consistent. Class id == the
74 // dict size at insertion time (stable, 0-based). This mapping will also be
75 // needed by the label decoder and saved with the model.
76 auto deprel2id = std::make_shared<std::map<std::string, int64_t>>();
77 {
78 CoNLLU::WordLevelAdapter adapter(&train_data);
80 while (adapter.end() != it)
81 {
82 if ((*it).is_word())
83 {
84 const std::string& rel = (*it).deprel();
85 if (deprel2id->find(rel) == deprel2id->end())
86 {
87 int64_t next_id = static_cast<int64_t>(deprel2id->size());
88 (*deprel2id)[rel] = next_id;
89 }
90 }
91 it++;
92 }
93 }
94 std::cerr << "Built deprel dictionary: " << deprel2id->size() << " labels" << std::endl;
95
96 // id -> deprel string, saved with the model so inference can label arcs.
97 std::vector<std::string> rel_class_names(deprel2id->size());
98 for (const auto& kv : *deprel2id)
99 {
100 rel_class_names[kv.second] = kv.first;
101 }
102
103 // The model's output tasks. Include "rel" when a deprel dict is available so the
104 // model is built (and saved) as a labeled parser; the inference output buffer is
105 // then sized to two columns (head, rel). Used for both construction and training.
106 std::vector<std::string> tasks = { "arc" };
107 if (!deprel2id->empty())
108 {
109 tasks.push_back("rel");
110 }
111
112 CoNLLUDataSet train_iterator(train_data,
113 params.m_batch_size,
114 feat_extractor,
115 tag_dh,
116 { p_embd },
118 deprel2id);
119 train_iterator.init();
120
121 CoNLLUDataSet dev_iterator(dev_data,
122 params.m_batch_size,
123 feat_extractor,
124 tag_dh,
125 { p_embd },
127 deprel2id);
128 dev_iterator.init();
129
130 BiRnnAndDeepBiaffineAttention model(nullptr);
131
132 if (params.m_input_model_name.size() == 0)
133 {
134 vector<embd_descr_t> embd_descr = train_iterator.get_embd_descr();
135 embd_descr.emplace(embd_descr.begin(), "raw", p_embd->dim(), 0);
136 vector<rnn_descr_t> rnn_descr;
137 rnn_descr.reserve(params.m_rnn_hidden_dims.size());
138 for (size_t d : params.m_rnn_hidden_dims)
139 {
140 rnn_descr.push_back(rnn_descr_t(d));
141 }
142
143 vector<deep_biaffine_attention_descr_t> decoder_descr = { deep_biaffine_attention_descr_t(128) };
144
145 // 1st arg = the dict holder used to size the input embeddings. The generated
146 // script numbers each embedding by its position in embd_descr, so this holder
147 // must be PARALLEL to embd_descr: a placeholder for the "raw" feature at index
148 // 0 (it has no Embedding) followed by the morph-feature dicts in the same
149 // order get_embd_descr() returns them. Passing the full tag dict holder here
150 // misaligned the indices whenever a feature was filtered out for having an
151 // empty dict, reading the wrong (often empty) dict and crashing.
152 // (2nd/6th arg = the output classes; a separate object, so no double-move.)
153 DictsHolder feat_dicts = train_iterator.get_embd_feature_dicts();
154 DictsHolder input_dicts;
155 if (!feat_dicts.empty())
156 {
157 input_dicts.push_back(feat_dicts.front()); // placeholder for "raw" (slot 0, unused)
158 for (const auto& d : feat_dicts)
159 {
160 input_dicts.push_back(d);
161 }
162 }
163 model = BiRnnAndDeepBiaffineAttention(std::move(input_dicts),
164 embd_descr,
165 rnn_descr,
166 decoder_descr,
167 tasks,
168 std::move(tag_dh),
169 boost::filesystem::path(params.m_embeddings_fn).stem().string(),
171 static_cast<int64_t>(deprel2id->size()),
172 rel_class_names);
173 }
174 else
175 {
176 model = BiRnnAndDeepBiaffineAttention();
177 model->load(params.m_input_model_name);
178 }
179
180 // cerr << model->get_script() << endl;
181
182 torch::optim::Adam optimizer(model->parameters(),
183 torch::optim::AdamOptions(params.m_learning_rate)
184 .weight_decay(params.m_weight_decay));
185
186 torch::Device device(params.m_device_string);
187
188 model->to(device);
189
190 double min_perf = 0;
191
192 for (const string& opt_name : utils::split(params.m_optimizers, ','))
193 {
194 shared_ptr<torch::optim::Optimizer> optimizer;
195
196 if (opt_name == "adam")
197 {
198 optimizer = make_shared<torch::optim::Adam>(model->parameters(),
199 torch::optim::AdamOptions(params.m_learning_rate)
200 .weight_decay(params.m_weight_decay));
201 }
202 else if (opt_name == "sgd")
203 {
204 optimizer = make_shared<torch::optim::SGD>(model->parameters(),
205 torch::optim::SGDOptions(params.m_learning_rate * 1000)
206 .weight_decay(params.m_weight_decay));
207 }
208 else
209 {
210 throw runtime_error("Unknown optimizer: " + opt_name);
211 }
212
213 model->train(params, tasks,
214 train_iterator, dev_iterator,
215 *optimizer, min_perf, device);
216
217 std::cerr << "train_graph_dp: Optimizer " << opt_name << " stopped at " << min_perf << std::endl;
218 }
219
220 return 0;
221}
222
223} // train
224} // graph_dp
225} // deeplima
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< nets::embd_descr_t > get_embd_descr()
WordDictBuilderImpl< CoNLLU::WordLevelAdapter, ConlluFeatExtractor< CoNLLU::WordLevelAdapter::token_t > > TagDictBuilderFromCoNLLU
FastTextVectorizer< TorchMatrix< float > > FastTextVectorizerToTorchMatrix
int train_graph_dp(const train_params_graph_dp_t &params)
std::vector< std::string > split(const std::string &str, char delim)
STL namespace.