LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_seq2seq.cpp
Go to the documentation of this file.
1// Copyright 2002-2022 CEA LIST
2// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
6#include <string>
7
8#include "birnn_seq2seq.h"
9
10using namespace std;
11using namespace torch;
12using torch::indexing::Slice;
13
14namespace deeplima
15{
16namespace nets
17{
18
19void BiRnnSeq2SeqImpl::train_batch(const std::vector<std::string>& output_names,
20 const torch::Tensor& input,
21 const vector<torch::Tensor>& input_cat,
22 const torch::Tensor& target,
23 torch::optim::Optimizer& opt,
24 epoch_stat_t& stat,
25 const torch::Device& device)
26{
27 opt.zero_grad();
28 int64_t l = target.sizes()[0];
29 torch::Tensor gold_input = torch::empty({1, target.sizes()[1]},
30 torch::TensorOptions().dtype(torch::kInt64)).fill_(2).to(device);
31 gold_input = torch::cat({gold_input, target.index({Slice(0, l-1), Slice()})}, 0);
32 //cerr << input.sizes() << " " << gold_input.sizes() << endl;
33 //cerr << gold_input << endl;
34 //cerr << target << endl;
35 std::map<string, torch::Tensor> input_map
36 = { { "enc_chars", input }, { "dec_chars", gold_input } };
37
38 assert(input_cat.size() == m_cat_embd_descr.size());
39 for (size_t feat_idx = 0; feat_idx < m_cat_embd_descr.size(); ++feat_idx)
40 {
41 //std::cerr << input_cat[feat_idx].sizes() << std::endl;
42 input_map[m_cat_embd_descr[feat_idx].m_name] = torch::squeeze(input_cat[feat_idx], 0).to(device);
43 }
44 auto output_map = forward(input_map, output_names.begin(), output_names.end());
45
46 for (size_t i = 0; i < output_names.size(); ++i)
47 {
48 const string& task_name = output_names[i];
49 task_stat_t& task_stat = stat[task_name];
50
51 auto output = output_map[task_name].to(device);
52
53 //cerr << "output.sizes() == " << output.sizes() << std::endl;
56 auto o = output.reshape({ -1, output.size(2) }).to(device);
57 //cerr << target.sizes() << endl;
58 auto this_task_target = target.reshape({ -1 }).to(device); //.index({ Slice(), Slice(i, i+1) });
59 //cerr << o.sizes() << endl;
60 //cerr << this_task_target.sizes() << endl;
61 torch::Tensor loss_tensor = torch::nn::functional::nll_loss(o, this_task_target);
62 loss_tensor.to(device);
63 //cerr << loss_tensor.sizes() << endl;
64 double loss_value = loss_tensor.mean().item<double>();
65 task_stat.m_loss += loss_value;
66 //std::cerr << "o.sizes() == " << o.sizes() << std::endl;
67 auto prediction = o.argmax(1);
68 int64_t correct_predictions = prediction.eq(this_task_target).sum().item<int64_t>();
69 task_stat.m_correct += correct_predictions;
70 task_stat.m_items += this_task_target.size(0);
71
72 loss_tensor.backward({}, true);
73 }
74 opt.step();
75}
76
77void BiRnnSeq2SeqImpl::evaluate(const vector<string>& ,
81 const torch::Device& )
82{
83 eval();
84}
85
86void BiRnnSeq2SeqImpl::load(torch::serialize::InputArchive& archive)
87{
89}
90
91void BiRnnSeq2SeqImpl::save(torch::serialize::OutputArchive& archive) const
92{
94}
95
96} // namespace lemmatization
97} // namespace deeplima
98
virtual void load(torch::serialize::InputArchive &archive)
virtual void save(torch::serialize::OutputArchive &archive) const
void train_batch(const std::vector< std::string > &output_names, const torch::Tensor &input, const std::vector< torch::Tensor > &input_cat, const torch::Tensor &target, torch::optim::Optimizer &opt, epoch_stat_t &stat, const torch::Device &device)
std::vector< embd_descr_t > m_cat_embd_descr
virtual void save(torch::serialize::OutputArchive &archive) const
virtual void load(torch::serialize::InputArchive &archive)
void evaluate(const std::vector< std::string > &output_names, const TorchMatrix< int64_t > &input, const TorchMatrix< int64_t > &gold, epoch_stat_t &stat, const torch::Device &device=torch::Device(torch::kCPU))
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.