LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_and_deep_biaffine_attention.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 <algorithm>
7#include <chrono>
8#include <iostream>
9#include <limits>
10
12#include "static_graph/dict.h"
13
14using namespace std;
15using namespace torch;
16using torch::indexing::Slice;
17
18using namespace deeplima::nets;
19
20namespace deeplima
21{
22namespace graph_dp
23{
24namespace train
25{
26
27#define SERIALIZATION_KEY_TASKS "tasks"
28#define SERIALIZATION_KEY_INPUT_FEATURES "input_features"
29#define SERIALIZATION_KEY_INPUT_FEATURES_NAMES "input_features_names"
30#define SERIALIZATION_KEY_EMBD_FN "embd_fn"
31#define SERIALIZATION_KEY_REL_CLASSES "rel_classes"
32
33void BiRnnAndDeepBiaffineAttentionImpl::load(serialize::InputArchive& archive)
34{
36
37 //assert(m_classes.size() == 0);
38
39 c10::IValue v;
40 if (archive.try_read(SERIALIZATION_KEY_TASKS, v))
41 {
42 if (!v.isList())
43 {
44 throw std::runtime_error("List of tasks must be a list.");
45 }
46
47 const c10::List<c10::IValue>& l = v.toList();
48 m_output_class_names.reserve(l.size());
49
50 for (size_t i = 0; i < l.size(); i++)
51 {
52 if (!l.get(i).isString())
53 {
54 throw std::runtime_error("List of tasks must be a list of strings.");
55 }
56
57 m_output_class_names.push_back(l.get(i).toStringRef());
58 }
59 }
60 else
61 {
62 throw std::runtime_error("Can't load list of tasks.");
63 }
64
65 if (archive.try_read(SERIALIZATION_KEY_INPUT_FEATURES_NAMES, v))
66 {
67 if (!v.isList())
68 {
69 throw std::runtime_error("List of input feature names must be a list.");
70 }
71
72 const c10::List<c10::IValue>& l = v.toList();
73 m_input_class_names.reserve(l.size());
74
75 for (size_t i = 0; i < l.size(); i++)
76 {
77 if (!l.get(i).isString())
78 {
79 throw std::runtime_error("List of input features names must be a list of strings.");
80 }
81
82 m_input_class_names.push_back(l.get(i).toStringRef());
83 }
84 }
85 else
86 {
87 throw std::runtime_error("Can't load list of input feature names.");
88 }
89
90 if (archive.try_read(SERIALIZATION_KEY_INPUT_FEATURES, v))
91 {
92 if (!v.isList())
93 {
94 throw std::runtime_error("List of input features must be a list.");
95 }
96
97 const c10::List<c10::IValue>& l = v.toList();
98 m_input_classes.reserve(l.size());
99
100 for (size_t i = 0; i < l.size(); i++)
101 {
102 if (!l.get(i).isList())
103 {
104 throw std::runtime_error("List of input features must be a list of lists of strings.");
105 }
106
107 auto d = std::make_shared<StringDict>();
108 d->fromIValue(l.get(i));
109 m_input_classes.push_back(d);
110 }
111 }
112 else
113 {
114 throw std::runtime_error("Can't list of input features.");
115 }
116
117 if (archive.try_read(SERIALIZATION_KEY_EMBD_FN, v))
118 {
119 if (!v.isList())
120 {
121 throw std::runtime_error("embd_fn must be a list.");
122 }
123
124 const c10::List<c10::IValue>& l = v.toList();
125 const c10::IValue& s = l.get(0);
126
127 if (!s.isString())
128 {
129 throw std::runtime_error("embd_fm must be a list of strings.");
130 }
131
132 m_embd_fn = s.toStringRef();
133 }
134 else
135 {
136 throw std::runtime_error("Can't load embd_fn.");
137 }
138
139 // deprel (rel) class names: optional, absent in arc-only models.
140 if (archive.try_read(SERIALIZATION_KEY_REL_CLASSES, v))
141 {
142 if (!v.isList())
143 {
144 throw std::runtime_error("List of rel classes must be a list.");
145 }
146 const c10::List<c10::IValue>& l = v.toList();
147 m_rel_class_names.clear();
148 m_rel_class_names.reserve(l.size());
149 for (size_t i = 0; i < l.size(); i++)
150 {
151 if (!l.get(i).isString())
152 {
153 throw std::runtime_error("List of rel classes must be a list of strings.");
154 }
155 m_rel_class_names.push_back(l.get(i).toStringRef());
156 }
157 m_num_labels = (int64_t) m_rel_class_names.size();
158 }
159}
160
161void BiRnnAndDeepBiaffineAttentionImpl::save(serialize::OutputArchive& archive) const
162{
164
165 // Save output class names
166 {
167 c10::List<std::string> list_of_class_names;
168 list_of_class_names.reserve(m_output_class_names.size());
169 for (size_t i = 0; i < m_output_class_names.size(); i++)
170 {
171 list_of_class_names.push_back(m_output_class_names[i]);
172 }
173 archive.write(SERIALIZATION_KEY_TASKS, list_of_class_names);
174 }
175
176 // Save input class names
177 {
178 c10::List<std::string> list_of_class_names;
179 list_of_class_names.reserve(m_input_class_names.size());
180 for (size_t i = 0; i < m_input_class_names.size(); i++)
181 {
182 list_of_class_names.push_back(m_input_class_names[i]);
183 }
184 archive.write(SERIALIZATION_KEY_INPUT_FEATURES_NAMES, list_of_class_names);
185 }
186
187 // Save input classes
188 c10::List<c10::List<std::string>> list_of_classes;
189 for (size_t i = 0; i < m_input_classes.size(); i++)
190 {
191 c10::List<std::string> current_class;
192 shared_ptr<StringDict> d = dynamic_pointer_cast<StringDict, DictBase>(m_input_classes[i]);
193 for (size_t j = 0; j < d->size(); j++)
194 {
195 current_class.push_back(d->get_value(j));
196 }
197 list_of_classes.push_back(current_class);
198 }
199 archive.write(SERIALIZATION_KEY_INPUT_FEATURES, list_of_classes);
200
201 // Save embeddings file names
202 c10::List<std::string> list_of_embd_fn;
203 list_of_embd_fn.reserve(1);
204 list_of_embd_fn.push_back(m_embd_fn);
205 archive.write(SERIALIZATION_KEY_EMBD_FN, list_of_embd_fn);
206
207 // Save deprel (rel) class names so inference can map predicted ids to labels.
208 c10::List<std::string> list_of_rel_classes;
209 list_of_rel_classes.reserve(m_rel_class_names.size());
210 for (size_t i = 0; i < m_rel_class_names.size(); i++)
211 {
212 list_of_rel_classes.push_back(m_rel_class_names[i]);
213 }
214 archive.write(SERIALIZATION_KEY_REL_CLASSES, list_of_rel_classes);
215}
216
218 const std::vector<std::string>& output_names,
219 const IterableDataSet& train_batches,
220 const IterableDataSet& eval_batches,
221 torch::optim::Optimizer& opt,
222 double& best_eval_accuracy,
223 const torch::Device& device)
224{
225 auto train_iterator = train_batches.get_iterator();
226 train_iterator->set_batch_size(params.m_batch_size);
227
228 double best_eval_loss = std::numeric_limits<double>::max();
229 size_t count_below_best = 0;
230 double lr_copy = 0;
231 unsigned int epoch = 0;
232 while (true)
233 {
234 for (auto &group : opt.param_groups())
235 {
236 if (group.has_options())
237 {
238 auto &options = static_cast<torch::optim::AdamOptions &>(group.options());
239 // cout << "LR == " << options.get_lr() << endl;
240 lr_copy = options.get_lr();
241 }
242 }
243
244 nets::epoch_stat_t train_stat, eval_stat;
245
246 chrono::steady_clock::time_point begin = chrono::steady_clock::now();
247
249 0,
250 output_names,
251 train_iterator,
252 opt,
253 train_stat,
254 device);
255
256 for (const std::string& tn : output_names)
257 {
258 if (train_stat[tn].m_items > 0)
259 {
260 train_stat[tn].m_accuracy = double(train_stat[tn].m_correct) / train_stat[tn].m_items;
261 train_stat[tn].m_loss /= train_stat[tn].m_items;
262 }
263 }
264
265 chrono::steady_clock::time_point train_end = chrono::steady_clock::now();
266
267 evaluate(output_names, eval_batches.get_iterator(), eval_stat, device);
268
269 for (const std::string& tn : output_names)
270 {
271 if (eval_stat[tn].m_items > 0)
272 {
273 eval_stat[tn].m_accuracy = double(eval_stat[tn].m_correct) / eval_stat[tn].m_items;
274 eval_stat[tn].m_loss /= eval_stat[tn].m_items;
275 }
276 }
277
278 chrono::steady_clock::time_point eval_end = chrono::steady_clock::now();
279
280 auto train_duration = std::chrono::duration_cast<std::chrono::milliseconds>(train_end - begin).count();
281 auto eval_duration = std::chrono::duration_cast<std::chrono::milliseconds>(eval_end - train_end).count();
282
283 for (const string& task_name : output_names)
284 {
285 const auto& train = train_stat[task_name];
286 const auto& eval = eval_stat[task_name];
287 std::cout << "EPOCH " << epoch << " | " << task_name << " | LR=" << lr_copy
288 << " | TRAIN LOSS=" << train.m_loss << " ACC=" << train.m_accuracy
289 << " | EVAL LOSS=" << eval.m_loss << " ACC=" << eval.m_accuracy;
290 if (2 == eval.m_num_classes)
291 {
292 std::cout << " P=" << eval.m_precision << " R=" << eval.m_recall<< " F1=" << eval.m_f1;
293 }
294 else
295 {
296 std::cout << " CORRECT=" << eval.m_correct;
297 }
298 std::cout << std::endl << std::flush;
299 }
300 cout << "TIME: train=" << train_duration << "[ms] eval=" << eval_duration << "[ms]" << endl;
301
302 task_stat_t& main_task_eval = eval_stat[output_names[0]];
303
304 best_eval_loss = min(best_eval_loss, main_task_eval.m_loss);
305 if (main_task_eval.m_accuracy > best_eval_accuracy)
306 {
307 best_eval_accuracy = main_task_eval.m_accuracy;
308 if (params.m_output_model_name.size() > 0)
309 {
310 torch::save(*this, params.m_output_model_name + ".pt");
311 }
312 count_below_best = 0;
313 }
314 else if (main_task_eval.m_accuracy < best_eval_accuracy)
315 {
316 for (auto &group : opt.param_groups())
317 {
318 if (group.has_options())
319 {
320 auto &options = static_cast<torch::optim::AdamOptions &>(group.options());
321 options.lr(options.lr() * (0.9));
322 lr_copy = options.lr();
323 }
324 }
325 if (lr_copy < 0.0000001)
326 {
327 return;
328 }
329 count_below_best++;
330 if (main_task_eval.m_loss > best_eval_loss && count_below_best > params.m_max_epochs_without_improvement)
331 {
332 return;
333 }
334 }
335 epoch++;
336 }
337}
338
340 size_t /*seq_len*/,
341 const vector<string>& output_names,
342 shared_ptr<BatchIterator> dataset_iterator,
343 torch::optim::Optimizer& opt,
344 epoch_stat_t& stat,
345 const torch::Device& device)
346{
347 Module::train(true);
348
349 dataset_iterator->start_epoch();
350
351 while (!dataset_iterator->end())
352 {
353 const BatchIterator::Batch batch = dataset_iterator->next_batch();
354 if (batch.empty())
355 {
356 continue;
357 }
359 0,
360 output_names,
361 batch.trainable_input(),
362 batch.frozen_input(),
363 batch.gold(),
364 opt,
365 stat,
366 device);
367 }
368}
369
371 size_t seq_len,
372 const vector<string>& output_names,
373 const torch::Tensor& trainable_input,
374 const torch::Tensor& nontrainable_input,
375 const torch::Tensor& gold,
376 torch::optim::Optimizer& opt,
377 epoch_stat_t& stat,
378 const torch::Device& device)
379{
380 using torch::indexing::Slice;
381 static constexpr int64_t kIgnore = -100;
382
383 map<string, torch::Tensor> current_batch_inputs;
384 split_input(trainable_input, current_batch_inputs, device);
385 current_batch_inputs["raw"] = nontrainable_input.to(device);
386
387 const int64_t B = gold.size(0);
388 const int64_t D = gold.size(1);
389 auto target = gold.reshape({-1, gold.size(2)}).to(device); // [B*D, ncols]
390
391 const bool train_rel = (m_num_labels > 0)
392 && (std::find(output_names.begin(), output_names.end(), std::string("rel")) != output_names.end());
393
394 std::vector<std::string> requested = { "arc" };
395 if (train_rel) { requested.push_back("rel_logits"); }
396
397 opt.zero_grad();
398 auto out = forward(current_batch_inputs, requested.begin(), requested.end());
399
400 // ---- arc (head prediction) ----
401 {
402 torch::Tensor o = out["arc"].reshape({ -1, out["arc"].size(2) }); // [N, heads]
403 torch::Tensor tgt = target.index({ Slice(), Slice(0, 1) }).reshape({ -1 });
404 torch::Tensor loss = torch::nn::functional::nll_loss(
405 o, tgt, torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnore));
406 torch::Tensor valid = tgt.ne(kIgnore);
407 torch::Tensor pred = o.argmax(1);
408 task_stat_t& s = stat["arc"];
409 s.m_loss += loss.sum().item<double>();
410 s.m_correct += pred.eq(tgt).logical_and(valid).sum().item<int64_t>();
411 s.m_items += valid.sum().item<int64_t>();
412 loss.backward({}, /*retain_graph=*/train_rel);
413 }
414
415 // ---- rel (deprel label), scored at the gold head ----
416 if (train_rel)
417 {
418 torch::Tensor rel_logits = out["rel_logits"]; // [B, dep, head, num_labels]
419 const int64_t L = rel_logits.size(3);
420 torch::Tensor gold_head = target.index({ Slice(), 0 }).reshape({ B, D });
421 torch::Tensor gh = gold_head.clamp_min(0); // ignored positions -> 0 (masked by loss)
422 torch::Tensor idx = gh.unsqueeze(-1).unsqueeze(-1).expand({ B, D, 1, L });
423 torch::Tensor gathered = rel_logits.gather(2, idx).squeeze(2); // [B, dep, num_labels]
424 torch::Tensor rel_log = torch::log_softmax(gathered, 2).reshape({ -1, L });
425 torch::Tensor tgt = target.index({ Slice(), Slice(1, 2) }).reshape({ -1 });
426 torch::Tensor loss = torch::nn::functional::nll_loss(
427 rel_log, tgt, torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnore));
428 torch::Tensor valid = tgt.ne(kIgnore);
429 torch::Tensor pred = rel_log.argmax(1);
430 task_stat_t& s = stat["rel"];
431 s.m_loss += loss.sum().item<double>();
432 s.m_correct += pred.eq(tgt).logical_and(valid).sum().item<int64_t>();
433 s.m_items += valid.sum().item<int64_t>();
434 loss.backward();
435 }
436
437 opt.step();
438}
439
440void BiRnnAndDeepBiaffineAttentionImpl::evaluate(const vector<string>& output_names,
441 shared_ptr<BatchIterator> dataset_iterator,
442 epoch_stat_t& stat,
443 const torch::Device& device)
444{
445 eval();
446 dataset_iterator->set_batch_size(-1);
447
448 dataset_iterator->start_epoch();
449
450 while (!dataset_iterator->end())
451 {
452 const BatchIterator::Batch batch = dataset_iterator->next_batch();
453 if (batch.empty())
454 {
455 continue;
456 }
457
458 epoch_stat_t t;
459
460 evaluate(output_names,
461 batch.trainable_input(),
462 batch.frozen_input(),
463 batch.gold(),
464 t,
465 device);
466 for (const std::string& task_name : output_names)
467 {
468 stat[task_name].m_correct += t[task_name].m_correct;
469 stat[task_name].m_items += t[task_name].m_items;
470 stat[task_name].m_loss += t[task_name].m_loss;
471 }
472 }
473}
474
475void BiRnnAndDeepBiaffineAttentionImpl::evaluate(const vector<string>& output_names,
476 const torch::Tensor& trainable_input,
477 const torch::Tensor& nontrainable_input,
478 const torch::Tensor& gold,
479 epoch_stat_t& stat,
480 const torch::Device& device)
481{
482 using torch::indexing::Slice;
483 static constexpr int64_t kIgnore = -100;
484
485 map<string, torch::Tensor> current_inputs;
486 split_input(trainable_input, current_inputs, device);
487 current_inputs["raw"] = nontrainable_input.to(device);
488
489 const int64_t B = gold.size(0);
490 const int64_t D = gold.size(1);
491 auto target = gold.reshape({-1, gold.size(-1)}).to(device);
492
493 const bool eval_rel = (m_num_labels > 0)
494 && (std::find(output_names.begin(), output_names.end(), std::string("rel")) != output_names.end());
495
496 std::vector<std::string> requested = { "arc" };
497 if (eval_rel) { requested.push_back("rel_logits"); }
498
499 torch::NoGradGuard no_grad;
500 auto out = forward(current_inputs, requested.begin(), requested.end());
501
502 // ---- arc (UAS) ----
503 {
504 torch::Tensor o = out["arc"].reshape({ -1, out["arc"].size(2) });
505 torch::Tensor tgt = target.index({ Slice(), Slice(0, 1) }).reshape({ -1 });
506 torch::Tensor loss = torch::nn::functional::nll_loss(
507 o, tgt, torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnore));
508 torch::Tensor valid = tgt.ne(kIgnore);
509 torch::Tensor pred = o.argmax(1);
510 task_stat_t& s = stat["arc"];
511 s.m_loss = loss.sum().item<double>();
512 s.m_correct = pred.eq(tgt).logical_and(valid).sum().item<int64_t>();
513 s.m_items = valid.sum().item<int64_t>();
514 }
515
516 // ---- rel (label accuracy at the gold head) ----
517 if (eval_rel)
518 {
519 torch::Tensor rel_logits = out["rel_logits"]; // [B, dep, head, num_labels]
520 const int64_t L = rel_logits.size(3);
521 torch::Tensor gold_head = target.index({ Slice(), 0 }).reshape({ B, D });
522 torch::Tensor gh = gold_head.clamp_min(0);
523 torch::Tensor idx = gh.unsqueeze(-1).unsqueeze(-1).expand({ B, D, 1, L });
524 torch::Tensor gathered = rel_logits.gather(2, idx).squeeze(2);
525 torch::Tensor rel_log = torch::log_softmax(gathered, 2).reshape({ -1, L });
526 torch::Tensor tgt = target.index({ Slice(), Slice(1, 2) }).reshape({ -1 });
527 torch::Tensor loss = torch::nn::functional::nll_loss(
528 rel_log, tgt, torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnore));
529 torch::Tensor valid = tgt.ne(kIgnore);
530 torch::Tensor pred = rel_log.argmax(1);
531 task_stat_t& s = stat["rel"];
532 s.m_loss = loss.sum().item<double>();
533 s.m_correct = pred.eq(tgt).logical_and(valid).sum().item<int64_t>();
534 s.m_items = valid.sum().item<int64_t>();
535 }
536}
537
539 const torch::Tensor& /*inputs*/,
540 int64_t /*input_begin*/,
541 int64_t /*input_end*/,
542 int64_t /*output_begin*/,
543 int64_t /*output_end*/,
544 std::shared_ptr< StdMatrix<uint8_t> >& /*output*/,
545 const std::vector<std::string>& /*outputs_names*/,
546 const torch::Device& /*device*/)
547{
548// TODO should it be implemented?
549}
550
551} // namespace train
552} // namespace graph_dp
553} // namespace deeplima
#define SERIALIZATION_KEY_TASKS
#define SERIALIZATION_KEY_EMBD_FN
#define SERIALIZATION_KEY_INPUT_FEATURES
#define SERIALIZATION_KEY_INPUT_FEATURES_NAMES
#define SERIALIZATION_KEY_REL_CLASSES
void train(const train_params_graph_dp_t &params, const std::vector< std::string > &output_names, const IterableDataSet &train_batches, const IterableDataSet &eval_batches, torch::optim::Optimizer &opt, double &best_eval_accuracy, const torch::Device &device=torch::Device(torch::kCPU))
void train_batch(size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const torch::Tensor &trainable_input, const torch::Tensor &nontrainable_input, const torch::Tensor &gold, torch::optim::Optimizer &opt, nets::epoch_stat_t &stat, const torch::Device &device)
void train_epoch(size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, std::shared_ptr< BatchIterator > train_iterator, torch::optim::Optimizer &opt, nets::epoch_stat_t &stat, const torch::Device &device)
void predict(size_t worker_id, const torch::Tensor &inputs, int64_t input_begin, int64_t input_end, int64_t output_begin, int64_t output_end, std::shared_ptr< StdMatrix< uint8_t > > &output, const std::vector< std::string > &outputs_names, const torch::Device &device=torch::Device(torch::kCPU))
virtual void load(torch::serialize::InputArchive &archive)
virtual void save(torch::serialize::OutputArchive &archive) const
void evaluate(const std::vector< std::string > &output_names, std::shared_ptr< BatchIterator > dataset_iterator, nets::epoch_stat_t &stat, const torch::Device &device=torch::Device(torch::kCPU))
virtual std::shared_ptr< BatchIterator > get_iterator() const =0
void split_input(const torch::Tensor &src, std::map< std::string, torch::Tensor > &dst, const torch::Device &device)
virtual void load(torch::serialize::InputArchive &archive)
virtual void save(torch::serialize::OutputArchive &archive) const
virtual void to(torch::Device device, bool non_blocking=false) override
std::map< std::string, torch::Tensor > forward(const std::map< std::string, torch::Tensor > &inputs, const std::string &output_name)
std::map< std::string, task_stat_t > epoch_stat_t
STL namespace.