LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
conllu_file_iterator.cpp
Go to the documentation of this file.
1// Copyright 2022 CEA LIST
2// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
8
9#include <iostream>
10#include <random>
11
12using namespace std;
13using namespace torch;
14using torch::indexing::Slice;
15
16namespace deeplima
17{
18namespace graph_dp
19{
20namespace train
21{
22
24 size_t batch_size,
26 const DictsHolder& morph_tag_dh,
27 shared_ptr<FeatureVectorizerBase<>> p_embd,
28 bool add_root,
29 shared_ptr<const std::map<std::string, int64_t>> deprel2id)
30 : m_add_root(add_root),
31 m_batch_size(batch_size),
32 m_morph_tag_dh(morph_tag_dh),
33 m_deprel2id(deprel2id),
34 m_annot(annot),
35 m_feat_vectorizers({ p_embd }),
36 m_feat_extractor(feat_extractor)
37{
38 for (auto p_embd : m_feat_vectorizers)
39 {
40 m_feat_descr.push_back({ CoNLLUToTorchMatrix::str_feature, "form", p_embd });
41 }
42
43 for (size_t i = 0; i < m_morph_tag_dh.size(); ++i)
44 {
45 if (morph_tag_dh[i]->size() > 0)
46 {
47 m_embd_feat_descr.push_back({
49 m_feat_extractor.feats()[i],
50 8,
51 morph_tag_dh[i]
52 });
53 }
54 }
55}
56
58{
59 //load_embeddings();
60 vectorize();
61}
62
63std::vector<nets::embd_descr_t> CoNLLUDataSet::get_embd_descr()
64{
65 assert(m_feat_descr.size() > 0 || m_embd_feat_descr.size() > 0);
66 CoNLLUToTorchMatrix vectorizer(m_feat_descr, m_embd_feat_descr, m_feat_extractor);
67 return vectorizer.get_embd_descr();
68}
69
71{
72 DictsHolder dicts;
73 for (const auto& fd : m_embd_feat_descr)
74 {
75 dicts.push_back(fd.m_dict);
76 }
77 return dicts;
78}
79
80void CoNLLUDataSet::vectorize()
81{
82 assert(m_feat_descr.size() > 0 || m_embd_feat_descr.size() > 0);
83 CoNLLUToTorchMatrix vectorizer(m_feat_descr, m_embd_feat_descr, m_feat_extractor);
84
85 // sort sentences by length (bucketing)
86 vector<vector<size_t>> sent_by_len;
87 sent_by_len.resize(256);
88
89 for (size_t i = 0; i < m_annot.get_num_sentences(); i++)
90 {
91 const CoNLLU::Sentence& sent = m_annot.get_sentence(i);
92 size_t len = sent.calc_num_of_words(m_annot);
93 assert(len > 0);
94 if (len < 4)
95 {
96 // nothing to train :)
97 continue;
98 }
99 if (len > sent_by_len.size() - 1)
100 {
101 sent_by_len.resize(len + 1);
102 }
103 sent_by_len[len].push_back(i);
104 }
105
106 m_bucket_keys.reserve(sent_by_len.size());
107 for (size_t i = 0; i < sent_by_len.size(); i++)
108 {
109 const vector<size_t>& sents = sent_by_len[i];
110 if (sents.empty())
111 {
112 continue;
113 }
114
115 size_t k = vectorize_bucket(i, sents, vectorizer);
116 m_bucket_keys.push_back(k);
117 }
118}
119
120size_t CoNLLUDataSet::vectorize_bucket(size_t len, const vector<size_t>& sents, const CoNLLUToTorchMatrix& vectorizer)
121{
122 uint64_t timepoint = 0;
123 uint64_t timepoints_per_sentence = m_add_root ? len + 1 : len;
124 uint64_t total_timepoints = timepoints_per_sentence * sents.size();
125 m_input_buckets[timepoints_per_sentence] = vectorizer.init_dst(total_timepoints);
126 // Two gold columns: column 0 = head (arc target), column 1 = deprel (rel target).
127 m_gold_buckets[timepoints_per_sentence] = std::make_shared<TorchMatrix<int64_t>>(total_timepoints, 2);
128
129 CoNLLU::Annotation m_root_annot;
130 stringstream root_line;
131 root_line << "1\t<ROOT>\t<ROOT>\tROOT\t_\t_\t0\troot\t0:root\t_\n";
132 m_root_annot.load(root_line);
133 CoNLLU::BoundedWordLevelAdapter root_generator(&m_root_annot, 0, 1);
134
135 for (size_t i = 0; i < sents.size(); i++)
136 {
137 const CoNLLU::Sentence& sent = m_annot.get_sentence(sents[i]);
138
139 CoNLLU::BoundedWordLevelAdapter adapter(&m_annot, sent.get_first_word_idx(), len);
140
141 if (m_add_root)
142 {
143 vectorizer.process(root_generator, m_input_buckets[timepoints_per_sentence], timepoint);
144 }
145
146 vectorizer.process(adapter, m_input_buckets[timepoints_per_sentence], m_add_root ? timepoint+1 : timepoint);
147 vectorize_bucket_gold(adapter, *(m_gold_buckets[timepoints_per_sentence].get()), timepoint);
148
149 timepoint += timepoints_per_sentence;
150 }
151
152 return timepoints_per_sentence;
153}
154
155void CoNLLUDataSet::vectorize_bucket_gold(const CoNLLU::BoundedWordLevelAdapter& src,
156 TorchMatrix<int64_t>& dst,
157 uint64_t timepoint)
158{
159 typename CoNLLU::WordLevelAdapter::const_iterator it = src.begin();
160 uint64_t current_timepoint = timepoint;
161 while (src.end() != it)
162 {
163 while(!(*it).is_word() && src.end() != it)
164 {
165 it++;
166 }
167
168 if (src.end() == it)
169 {
170 break;
171 }
172
173 const CoNLLU::idx_t& head = (*it).head();
174 const std::string& gold_rel = (*it).deprel();
175 assert(head.is_real_word());
176
177 const CoNLLU::idx_t idx = (*it).idx();
178 if (m_add_root && 1 == idx._first)
179 {
180 // The synthetic <ROOT> token is not a dependent: ignore it in both the
181 // arc and rel losses.
182 dst.set(current_timepoint, 0, IGNORE_INDEX);
183 dst.set(current_timepoint, 1, IGNORE_INDEX);
184 current_timepoint++;
185 }
186
187 // Column 0: head position. With <ROOT> prepended at position 0, the CoNLL-U
188 // HEAD value (0 = root, k = word k) maps directly to the timepoint position.
189 dst.set(current_timepoint, 0, head._first);
190
191 // Column 1: deprel class id (IGNORE_INDEX if no dict or label is unknown).
192 int64_t rel_id = IGNORE_INDEX;
193 if (m_deprel2id)
194 {
195 auto rel_it = m_deprel2id->find(gold_rel);
196 rel_id = (rel_it != m_deprel2id->end()) ? rel_it->second : IGNORE_INDEX;
197 }
198 dst.set(current_timepoint, 1, rel_id);
199
200 current_timepoint++;
201 if (current_timepoint == std::numeric_limits<uint64_t>::max())
202 {
203 throw std::overflow_error("Too much words in the dataset.");
204 }
205
206 it++;
207 }
208}
209
210std::shared_ptr<BatchIterator> CoNLLUDataSet::get_iterator() const
211{
212 return std::make_shared<Iterator>(*this);
213}
214
216{
217 m_batch_size = batch_size;
218}
219
221{
222 if (0 == m_batch_size)
223 {
224 throw std::runtime_error("CoNLLUDataSet::Iterator::start_epoch batch size cannot be 0.");
225 }
226 m_current_bucket = 0;
227 m_iter_counter = 0;
228}
229
231{
232 return m_current_bucket >= m_dataset.m_bucket_keys.size();
233}
234
236{
237 size_t k = m_dataset.m_bucket_keys[m_current_bucket];
238
239 auto it_input = m_dataset.m_input_buckets.find(k);
240 assert(m_dataset.m_input_buckets.end() != it_input);
241 const CoNLLUToTorchMatrix::vectorization_t& input = it_input->second;
242
243 auto it_gold = m_dataset.m_gold_buckets.find(k);
244 assert(m_dataset.m_gold_buckets.end() != it_gold);
245 const TorchMatrix<int64_t>& gold_bucket = *(it_gold->second);
246 //cerr << "gold_bucket.sizes() == " << gold_bucket.get_tensor().sizes() << endl;
247 //cerr << "gold_bucket.get_tensor().size(0) == " << gold_bucket.get_tensor().size(0) << endl;
248
249 int64_t seq_len = it_input->first;
250
251 //cerr << "input.first->get_tensor().sizes() == " << input.first->get_tensor().sizes() << endl;
252 //cerr << "input.second->get_tensor().sizes() == " << input.second->get_tensor().sizes() << endl;
253
254 int64_t batch_size = m_batch_size;
255 if (-1 == batch_size)
256 {
257 assert(0 == input.first->get_tensor().size(0) % seq_len);
258 batch_size = input.first->get_tensor().size(0) / seq_len; // batch == full bucket
259 }
260
261 //std::cerr << gold_bucket.get_tensor().sizes() << std::endl;
262
263 std::random_device r;
264 std::default_random_engine e1(r());
265 int64_t max_start_offset = ( gold_bucket.get_tensor().size(0) / seq_len ) - batch_size;
266 if (max_start_offset < 0)
267 {
268 max_start_offset = 0;
269 }
270 std::uniform_int_distribution<int> uniform_dist(0, max_start_offset);
271 int64_t batch_start_offset = uniform_dist(e1);
272
273 if (0 == batch_start_offset &&
274 seq_len * (batch_start_offset + batch_size) > gold_bucket.get_tensor().size(0))
275 {
276 m_current_bucket++;
278 }
279
280 /*std::cerr << "seq_len = " << seq_len
281 << " iter " << m_iter_counter
282 << " from " << batch_start_offset << " until "
283 << gold_bucket.get_tensor().size(0) / seq_len << std::endl;*/
284 const torch::Tensor trainable
285 = input.first->get_tensor().index({ Slice(seq_len * batch_start_offset,
286 seq_len * (batch_start_offset + batch_size)),
287 Slice() }).reshape({ batch_size, seq_len, -1 }).transpose(0, 1);
288 const torch::Tensor frozen
289 = input.second->get_tensor().index({ Slice(seq_len * batch_start_offset,
290 seq_len * (batch_start_offset + batch_size)),
291 Slice() }).reshape({ batch_size, seq_len, -1 }).transpose(0, 1);
292
293 //cerr << "trainable.sizes() == " << trainable.sizes() << endl;
294 //cerr << "frozen.sizes() == " << frozen.sizes() << endl;
295
296 const torch::Tensor gold
297 = gold_bucket.get_tensor().index({ Slice(seq_len * batch_start_offset,
298 seq_len * (batch_start_offset + batch_size)),
299 Slice() }).reshape({ batch_size, seq_len, -1 });//.transpose(0, 1);
300 const CoNLLUDataSet::Iterator::Batch batch(trainable, frozen, gold, 1);
301
302 // advance pointers
303 if (m_batch_size >= 0)
304 {
305 m_iter_counter++;
306 //std::cerr << gold_bucket.get_tensor().size(0) << std::endl;
307 if (m_iter_counter * m_batch_size > gold_bucket.get_tensor().size(0) / seq_len)
308 {
309 m_current_bucket++;
310 m_iter_counter = 0;
311 }
312 }
313 else
314 {
315 m_current_bucket++;
316 m_iter_counter = 0;
317 }
318
319 return batch;
320}
321
322} // train
323} // graph_dp
324} // deeplima
size_t get_num_sentences() const
Definition treebank.h:142
const Sentence & get_sentence(size_t idx) const
Definition treebank.cpp:380
size_t calc_num_of_words(const Annotation &annot) const
Definition treebank.cpp:25
std::vector< std::string > feats() const
const torch::Tensor & get_tensor() const
std::vector< nets::embd_descr_t > get_embd_descr()
virtual std::shared_ptr< BatchIterator > get_iterator() const
CoNLLUDataSet(const CoNLLU::Annotation &annot, size_t batch_size, const ConlluFeatExtractor< CoNLLU::WordLevelAdapter::token_t > &feat_extractor, const DictsHolder &morph_tag_dh, std::shared_ptr< FeatureVectorizerBase<> > p_embd, bool add_root=false, std::shared_ptr< const std::map< std::string, int64_t > > deprel2id=nullptr)
std::pair< std::shared_ptr< MatrixInt >, std::shared_ptr< MatrixFloat > > vectorization_t
const std::vector< deeplima::nets::embd_descr_t > get_embd_descr() const
STL namespace.