LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_seq_classifier.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 <iostream>
7#include <limits>
8#include <map>
9#include <string>
10
12
13using namespace std;
14using torch::indexing::Slice;
15
16namespace deeplima
17{
18namespace nets
19{
20
21torch::Tensor BiRnnClassifierImpl::predict(const string& output_name,
22 const torch::Tensor& input,
23 const torch::Device& device)
24{
25 vector<string> output_names = { output_name };
26 return predict(output_names, input, device);
27}
28
29torch::Tensor BiRnnClassifierImpl::predict(const vector<string>& output_names,
30 const torch::Tensor& input,
31 const torch::Device& device)
32{
33 map<string, torch::Tensor> current_inputs;
34 split_input(input, current_inputs, device);
35
36 auto output_map = forward(current_inputs, output_names.begin(), output_names.end());
37
38 torch::Tensor prediction = torch::zeros({ input.size(0), long(output_names.size()) },
39 torch::TensorOptions().dtype(torch::kInt64));
40 for (size_t i = 0; i < output_names.size(); ++i)
41 {
42 const string& task_name = output_names[i];
43 torch::Tensor& output = output_map[task_name];
44 torch::Tensor o = output.reshape({-1, output.size(2)});
45 prediction.index({ Slice(), Slice(i, i+1) }) = o.argmax(1);
46 }
47
48 return prediction;
49}
50
51void BiRnnClassifierImpl::predict(size_t /*worker_id*/,
52 const torch::Tensor& inputs,
53 int64_t input_begin,
54 int64_t input_end,
55 int64_t output_begin,
56 int64_t output_end,
57 std::shared_ptr< StdMatrix<uint8_t> >& output,
58 const std::vector<std::string>& outputs_names,
59 const torch::Device& device)
60{
61 assert(output->size() == outputs_names.size());
62
63 int64_t start_shift = output_begin - input_begin;
64 assert(start_shift >= 0);
65 int64_t end_shift = input_end - output_end;
66 assert(end_shift >= 0);
67
68 const torch::Tensor inputs_slice = inputs.index({ Slice(input_begin, input_end), Slice() });
69
70 map<string, torch::Tensor> current_inputs;
71 split_input(inputs_slice, current_inputs, device);
72
73 auto output_map = forward(current_inputs, outputs_names.begin(), outputs_names.end());
74 for (size_t i = 0; i < outputs_names.size(); i++)
75 {
76 torch::Tensor& one_task_output = output_map[outputs_names[i]];
77 torch::Tensor o = one_task_output.reshape({-1, one_task_output.size(2)});
78 const torch::Tensor output_tensor = o.argmax(1);
79 torch::TensorAccessor<int64_t, 1> accessor = output_tensor.accessor<int64_t, 1>();
80 for (int64_t p = start_shift; p < output_tensor.size(0) - end_shift; p++)
81 {
82 (*output)[i][input_begin + p] = accessor[p];
83 }
84 }
85}
86
87void BiRnnClassifierImpl::split_input(const torch::Tensor& src,
88 map<string, torch::Tensor>& dst,
89 const torch::Device& device)
90{
91 int64_t c = 0;
92 for (size_t i = 0; i < m_embd_descr.size(); i++)
93 {
94 if (dst.end() != dst.find(m_embd_descr[i].m_name))
95 {
96 throw std::runtime_error("Duplicates in embd list.");
97 }
98 if (src.sizes().size() == 2)
99 {
100 if (m_embd_descr[i].m_type == 1)
101 {
102 dst[m_embd_descr[i].m_name]
103 = src.index({Slice(), Slice(c, c+1) }).reshape({ (int64_t)src.size(0), 1 }).to(device);
104 c++;
105 }
106 else if (m_embd_descr[i].m_type == 0)
107 {
108 //dst[m_embd_descr[i].m_name]
109 // = src.index({Slice(), Slice(i, i+m_embd_descr[i].m_dim) }).reshape({ (int64_t)src.size(0), m_embd_descr[i].m_dim });
110 }
111 }
112 else if (src.sizes().size() == 3)
113 {
114 //std::cout << "src.sizes() == " << src.sizes() << std::endl;
115 if (m_embd_descr[i].m_type == 1)
116 {
117 dst[m_embd_descr[i].m_name]
118 = src.index({Slice(), Slice(), Slice(c, c+1) }).reshape({ (int64_t)src.size(0), -1 }).to(device);
119 c++;
120 }
121 else if (m_embd_descr[i].m_type == 0)
122 {
123 //dst[m_embd_descr[i].m_name]
124 // = src.index({Slice(), Slice(), Slice(i, i+m_embd_descr[i].m_dim) }).reshape({ (int64_t)src.size(0), m_embd_descr[i].m_dim });
125 }
126 }
127 //std::cout << "dst[" << m_embd_descr[i].m_name << "]=" << dst[m_embd_descr[i].m_name].sizes() << std::endl;
128 }
129}
130
131void BiRnnClassifierImpl::evaluate(const vector<string>& output_names,
132 const TorchMatrix<int64_t>& input,
133 const TorchMatrix<int64_t>& gold,
134 epoch_stat_t& stat,
135 const torch::Device& device)
136{
137 map<string, torch::Tensor> current_inputs;
138 split_input(input.get_tensor(), current_inputs, device);
139
140 auto target = gold.get_tensor().reshape({ -1, long(output_names.size()) }).to(device);
141
142 evaluate(output_names, current_inputs, target, stat, device);
143}
144
145void BiRnnClassifierImpl::evaluate(const vector<string>& output_names,
146 const map<string, torch::Tensor>& input,
147 const torch::Tensor& target,
148 epoch_stat_t& stat,
149 const torch::Device& /*device*/)
150{
151 eval();
152 auto output_map = forward(input, output_names.begin(), output_names.end());
153
154 for (size_t i = 0; i < output_names.size(); ++i)
155 {
156 const string& task_name = output_names[i];
157 task_stat_t& task_stat = stat[task_name];
158
159 torch::Tensor& output = output_map[task_name];
160 torch::Tensor o = output.reshape({-1, output.size(2)});
161 auto this_task_target = target.index({ Slice(), Slice(i, i+1) }).reshape({ -1 });
162 //cerr << o.sizes() << endl;
163 //cerr << this_task_target.sizes() << endl;
164
165 static constexpr int64_t kIgnoreIndex = -100;
166 torch::Tensor loss_tensor = torch::nn::functional::nll_loss(
167 o, this_task_target,
168 torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnoreIndex));
169 task_stat.m_loss = loss_tensor.sum().item<double>();
170 auto valid = this_task_target.ne(kIgnoreIndex);
171 auto prediction = o.argmax(1);
172 task_stat.m_correct = prediction.eq(this_task_target).logical_and(valid).sum().item<int64_t>();
173 task_stat.m_items = valid.sum().item<int64_t>();
174 task_stat.m_accuracy = double(task_stat.m_correct) / task_stat.m_items;
175
176 if (o.size(-1) == 2)
177 {
178 //cerr << this_task_target.sizes() << endl;
179 //cerr << prediction.sizes() << endl;
180 //cerr << this_task_target << endl;
181 //cerr << prediction << endl;
182 task_stat.m_num_classes = o.size(-1);
183 auto gold_positive = prediction.eq(1);
184 auto gold_negative = prediction.eq(0);
185 auto pred_positive = this_task_target.eq(1);
186 auto pred_negative = this_task_target.eq(0);
187 auto tp = torch::logical_and(gold_positive, pred_positive).count_nonzero().item<int64_t>();
188 // auto fp = torch::logical_and(gold_negative, pred_positive).count_nonzero().item<int64_t>();
189 // auto fn = torch::logical_and(gold_positive, pred_negative).count_nonzero().item<int64_t>();
190 task_stat.m_precision = tp / (prediction.eq(1).sum().item<int64_t>() + 0.00001);
191 task_stat.m_recall = tp / (this_task_target.eq(1).sum().item<int64_t>() + 0.00001);
192 task_stat.m_f1 = (2 * task_stat.m_precision * task_stat.m_recall) / (task_stat.m_precision + task_stat.m_recall + 0.00001);
193 }
194 else if (o.size(-1) > 2)
195 {
196 // Multi-class task (e.g. the 7-class segmentation "tokens" task): the
197 // binary positive-class P/R/F1 above doesn't apply, so report a
198 // macro-averaged P/R/F1 over all classes instead of leaving them at the
199 // default 0. ACC alone is a poor signal here because one class (e.g.
200 // "inside") dominates, pushing accuracy ~0.999 even for a mediocre
201 // boundary detector.
202 const int64_t num_classes = o.size(-1);
203 task_stat.m_num_classes = num_classes;
204 double sum_p = 0, sum_r = 0, sum_f1 = 0;
205 for (int64_t c = 0; c < num_classes; ++c)
206 {
207 // Ignored (-100) positions: excluded from gold_c (target never equals c)
208 // and masked out of pred_c so they can't inflate the predicted count.
209 auto pred_c = prediction.eq(c).logical_and(valid);
210 auto gold_c = this_task_target.eq(c);
211 auto tp_c = torch::logical_and(pred_c, gold_c).count_nonzero().item<int64_t>();
212 double prec = tp_c / (pred_c.sum().item<int64_t>() + 0.00001);
213 double rec = tp_c / (gold_c.sum().item<int64_t>() + 0.00001);
214 sum_p += prec;
215 sum_r += rec;
216 sum_f1 += (2 * prec * rec) / (prec + rec + 0.00001);
217 }
218 task_stat.m_precision = sum_p / num_classes;
219 task_stat.m_recall = sum_r / num_classes;
220 task_stat.m_f1 = sum_f1 / num_classes;
221 }
222 }
223}
224
225void BiRnnClassifierImpl::train(size_t epochs,
226 size_t batch_size,
227 size_t seq_len,
228 const vector<string>& output_names,
229 const TorchMatrix<int64_t>& train_input,
230 const TorchMatrix<int64_t>& train_gold,
231 const TorchMatrix<int64_t>& eval_input,
232 const TorchMatrix<int64_t>& eval_gold,
233 torch::optim::Optimizer& opt,
234 const std::string& model_name,
235 const torch::Device& device)
236{
237 int64_t num_batches = train_input.size() / seq_len;
238 int64_t num_features = (int64_t)train_input.get_max_feat();
239 int64_t seq_len_i64 = (int64_t)seq_len;
240
241 auto aligned_input = train_input.get_tensor().index({ Slice(0, num_batches * seq_len_i64), Slice() });
242 auto aligned_gold = train_gold.get_tensor().index({ Slice(0, num_batches * seq_len_i64), Slice() });
243 auto input_batches
244 = aligned_input.reshape({ num_batches, seq_len_i64, num_features }).transpose(0, 1);
245 auto gold_batches
246 = aligned_gold.reshape({ num_batches, seq_len_i64 }).transpose(0, 1);
247
248 double best_eval_accuracy = 0;
249 double best_eval_loss = std::numeric_limits<double>::max();
250 size_t count_below_best = 0;
251 // Seed lr_copy with the optimizer's actual learning rate so the EPOCH log
252 // prints the real LR from epoch 0. It was previously left at 0 and only set
253 // inside the accuracy-decay branch, so every epoch before the first decay
254 // (and every improving epoch) reported a bogus LR=0.
255 double lr_copy = 0;
256 for (auto& group : opt.param_groups())
257 {
258 if (group.has_options())
259 {
260 lr_copy = static_cast<torch::optim::AdamOptions&>(group.options()).lr();
261 }
262 }
263 for (size_t e = 0; e < epochs; e++)
264 {
265 epoch_stat_t train_stat, eval_stat;
266
267 Module::train(true);
268 train_epoch(batch_size, seq_len, output_names, input_batches, gold_batches, opt, train_stat, device);
269
270 evaluate(output_names, eval_input, eval_gold, eval_stat, device);
271
272 for (const string& task_name : output_names)
273 {
274 const task_stat_t& train = train_stat[task_name];
275 const task_stat_t& eval = eval_stat[task_name];
276 std::cout << "EPOCH " << e << " | " << task_name << " | LR=" << lr_copy;
277 std::cout << " | TRAIN LOSS=" << train.m_loss << " ACC=" << train.m_accuracy
278 << " | EVAL LOSS=" << eval.m_loss << " ACC=" << eval.m_accuracy
279 << " P=" << eval.m_precision << " R=" << eval.m_recall<< " F1=" << eval.m_f1
280 << std::endl << std::flush;
281 }
282
283 task_stat_t& main_task_eval = eval_stat[output_names[0]];
284
285 best_eval_loss = min(best_eval_loss, main_task_eval.m_loss);
286 if (main_task_eval.m_accuracy > best_eval_accuracy)
287 {
288 best_eval_accuracy = main_task_eval.m_accuracy;
289 if (model_name.size() > 0)
290 {
291 torch::save(*this, model_name + ".pt");
292 }
293 count_below_best = 0;
294
295 if (1 == main_task_eval.m_accuracy)
296 {
297 return;
298 }
299 }
300 else if (main_task_eval.m_accuracy < best_eval_accuracy)
301 {
302 for (auto &group : opt.param_groups())
303 {
304 if (group.has_options())
305 {
306 auto &options = static_cast<torch::optim::AdamOptions &>(group.options());
307 options.lr(options.lr() * (0.9));
308 lr_copy = options.lr();
309 }
310 }
311 if (lr_copy < 0.0000001)
312 {
313 return;
314 }
315 count_below_best++;
316 if (main_task_eval.m_loss > best_eval_loss && count_below_best > 10)
317 {
318 return;
319 }
320 if (count_below_best > 50)
321 {
322 return;
323 }
324 }
325 }
326}
327
328void BiRnnClassifierImpl::train_epoch(size_t batch_size,
329 size_t seq_len,
330 const vector<string>& output_names,
331 const torch::Tensor& input_batches,
332 const torch::Tensor& gold_batches,
333 torch::optim::Optimizer& opt,
334 epoch_stat_t& stat,
335 const torch::Device& device)
336{
337 for (int64_t b = 0; b < input_batches.size(1); b += batch_size)
338 {
339 auto current_batch_size
340 = ((b + (int64_t)batch_size) > input_batches.size(1)) ? input_batches.size(1) - b : batch_size;
341
342 train_batch(current_batch_size, seq_len, output_names,
343 input_batches.index({Slice(), Slice(b, b + current_batch_size), Slice()}),
344 gold_batches.index({Slice(), Slice(b, b + current_batch_size) }),
345 opt, stat, device);
346 }
347}
348
349void BiRnnClassifierImpl::train_batch(size_t batch_size,
350 size_t seq_len,
351 const vector<string>& output_names,
352 const torch::Tensor& input,
353 const torch::Tensor& gold,
354 torch::optim::Optimizer& opt,
355 epoch_stat_t& stat,
356 const torch::Device& device)
357{
358 map<string, torch::Tensor> current_batch_inputs;
359 split_input(input, current_batch_inputs, device);
360
361 auto target = gold.reshape({-1, 1}).to(device);
362
363 train_batch(batch_size, seq_len, output_names, current_batch_inputs, target, opt, stat, device);
364}
365
366void BiRnnClassifierImpl::train_batch(size_t /*batch_size*/,
367 size_t /*seq_len*/,
368 const vector<string>& output_names,
369 const map<string, torch::Tensor>& input,
370 const torch::Tensor& target,
371 torch::optim::Optimizer& opt,
372 epoch_stat_t& stat,
373 const torch::Device& /*device*/)
374{
375 opt.zero_grad();
376
377 auto output_map = forward(input, output_names.begin(), output_names.end());
378
379 //torch::Tensor loss_tensor = torch::zeros({ batch_size * seq_len },
380 // torch::TensorOptions().dtype(torch::kFloat64));
381 //cerr << loss_tensor.sizes() << endl;
382
383 for (size_t i = 0; i < output_names.size(); ++i)
384 {
385 const string& task_name = output_names[i];
386 task_stat_t& task_stat = stat[task_name];
387
388 auto output = output_map[task_name];
389
390 auto o = output.reshape({ -1, output.size(2) });
391 auto this_task_target = target.index({ Slice(), Slice(i, i+1) }).reshape({ -1 });
392 //cerr << o.sizes() << endl;
393 //cerr << this_task_target.sizes() << endl;
394 // Positions whose target == kIgnoreIndex (e.g. the synthetic <ROOT> token)
395 // are excluded from the loss and the accuracy. Harmless for tasks whose
396 // targets are always valid class ids (the mask is then all-true).
397 static constexpr int64_t kIgnoreIndex = -100;
398 torch::Tensor loss_tensor = torch::nn::functional::nll_loss(
399 o, this_task_target,
400 torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnoreIndex));
401 double loss_value = loss_tensor.sum().item<double>();
402 task_stat.m_loss += loss_value;
403 //std::cerr << "o.sizes() == " << o.sizes() << std::endl;
404 auto valid = this_task_target.ne(kIgnoreIndex);
405 auto prediction = o.argmax(1);
406 int64_t correct_predictions = prediction.eq(this_task_target).logical_and(valid).sum().item<int64_t>();
407 task_stat.m_correct += correct_predictions;
408 task_stat.m_items += valid.sum().item<int64_t>();
409
410 loss_tensor.backward({}, true);
411 }
412 opt.step();
413}
414
415void BiRnnClassifierImpl::load(torch::serialize::InputArchive& archive)
416{
417 StaticGraphImpl::load(archive);
418
419 assert(m_embd_descr.size() == 0);
420
421 c10::IValue v;
422 if (archive.try_read("embd_descr", v))
423 {
424 if (!v.isList())
425 {
426 throw std::runtime_error("embd_descr must be a list.");
427 }
428
429 const c10::List<c10::IValue>& l = v.toList();
430 m_embd_descr.reserve(l.size());
431 for (size_t i = 0; i < l.size(); i++)
432 {
433 if (!l.get(i).isString())
434 {
435 throw std::runtime_error("embd_descr must be a list of strings.");
436 }
437 const string& str = l.get(i).toStringRef();
438 m_embd_descr.emplace_back(embd_descr_t(str));
439 }
440 }
441 else
442 {
443 throw std::runtime_error("Can't load embd_descr.");
444 }
445}
446
447void BiRnnClassifierImpl::save(torch::serialize::OutputArchive& archive) const
448{
449 StaticGraphImpl::save(archive);
450
451 // Save embd descriptions
452 c10::List<std::string> embd_descr_list;
453 for (size_t i = 0; i < m_embd_descr.size(); i++)
454 {
455 string embd_str = m_embd_descr[i].to_string();
456 assert(embd_descr_t(embd_str) == m_embd_descr[i]);
457 embd_descr_list.push_back(embd_str);
458 }
459
460 archive.write("embd_descr", embd_descr_list);
461}
462
463} // namespace nets
464} // namespace deeplima
465
uint64_t get_max_feat() const
const torch::Tensor & get_tensor() const
uint64_t size() const
void split_input(const torch::Tensor &src, std::map< std::string, torch::Tensor > &dst, const torch::Device &device)
std::vector< embd_descr_t > m_embd_descr
void evaluate(const std::vector< std::string > &output_name, const TorchMatrix< int64_t > &input, const TorchMatrix< int64_t > &gold, epoch_stat_t &stat, const torch::Device &device=torch::Device(torch::kCPU))
torch::Tensor predict(const std::string &output_name, const torch::Tensor &input, const torch::Device &device=torch::Device(torch::kCPU))
virtual void load(torch::serialize::InputArchive &archive)
virtual void save(torch::serialize::OutputArchive &archive) const
void train_epoch(size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const torch::Tensor &input_batches, const torch::Tensor &gold_batches, torch::optim::Optimizer &opt, epoch_stat_t &stat, const torch::Device &device)
void train_batch(size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const torch::Tensor &input, const torch::Tensor &gold, torch::optim::Optimizer &opt, epoch_stat_t &stat, const torch::Device &device)
void train(size_t epochs, size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const TorchMatrix< int64_t > &train_input, const TorchMatrix< int64_t > &train_gold, const TorchMatrix< int64_t > &eval_input, const TorchMatrix< int64_t > &eval_gold, torch::optim::Optimizer &opt, const std::string &model_name="", const torch::Device &device=torch::Device(torch::kCPU))
virtual void to(torch::Device device, bool non_blocking=false) override
virtual void load(torch::serialize::InputArchive &archive) override
std::map< std::string, torch::Tensor > forward(const std::map< std::string, torch::Tensor > &inputs, const std::string &output_name)
virtual void save(torch::serialize::OutputArchive &archive) const override
std::map< std::string, task_stat_t > epoch_stat_t
STL namespace.