LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
seq2seq_for_lemmatization.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 <limits>
7#include <string>
8
10
12
13using namespace std;
14using namespace torch;
15using torch::indexing::Slice;
16using namespace deeplima::utils;
17
18namespace deeplima
19{
20namespace lemmatization
21{
22namespace train
23{
24
26 const vector<TorchMatrix<int64_t>>& train_input,
27 const vector<vector<TorchMatrix<int64_t>>>& train_input_cat,
28 const vector<TorchMatrix<int64_t>>& train_gold,
29 [[maybe_unused]] const vector<TorchMatrix<int64_t>>& eval_input,
30 const vector<vector<TorchMatrix<int64_t>>>& /*eval_input_cat*/,
31 [[maybe_unused]] const vector<TorchMatrix<int64_t>>& eval_gold,
32 torch::optim::Optimizer& opt,
33 const torch::Device& device)
34{
35 assert(train_input.size() == train_gold.size());
36 assert(eval_input.size() == eval_gold.size());
37
38 set_tags(params.m_tags);
39
40 // double best_eval_accuracy = 0;
41 // double best_eval_loss = std::numeric_limits<double>::max();
42 // size_t count_below_best = 0;
43 // double lr_copy = 0;
44 for (size_t e = 1; e < params.m_max_epochs; e++)
45 {
46 nets::epoch_stat_t train_stat, eval_stat;
47 Module::train(true);
48 for (size_t i = 0; i < train_input.size(); ++i)
49 {
50 const auto& input = train_input[i];
51 const auto& input_cat = train_input_cat[i];
52 const auto& gold = train_gold[i];
53
54 assert(input.get_max_feat() == gold.get_max_feat());
55 assert(input.get_max_feat() == input_cat[0].get_max_feat());
56 auto train_subset_stat = train_on_subset(params, input, input_cat, gold, opt, device);
57 train_stat["output"].m_loss = train_subset_stat["output"].m_loss;
58 train_stat["output"].m_correct += train_subset_stat["output"].m_correct;
59 train_stat["output"].m_items += train_subset_stat["output"].m_items;
60 }
61 if (train_stat["output"].m_items > 0)
62 {
63 train_stat["output"].m_accuracy = float(train_stat["output"].m_correct) / train_stat["output"].m_items;
64 }
65
66
67 // evaluate(eval_input, eval_gold, eval_stat, device);
68 std::cout << "EPOCH " << e << " | lemmatization"
69 << " | LR=" << params.m_learning_rate << " |"
70 // << " LEN=" << input_tensor.sizes()[0]
71 << " TRAIN LOSS=" << train_stat["output"].m_loss
72 << " ACC=" << train_stat["output"].m_accuracy
73 << " CORRECT=" << train_stat["output"].m_correct
74 << " TOTAL=" << train_stat["output"].m_items
75 << std::endl << std::flush;
76
77
78 if (!params.m_output_model_name.empty())
79 {
80 torch::save(*this, params.m_output_model_name + ".pt");
81 }
82 }
83}
84
86 const train_params_lemmatization_t& params,
87 const TorchMatrix<int64_t>& train_input,
88 const vector<TorchMatrix<int64_t>>& train_input_cat,
89 const TorchMatrix<int64_t>& train_gold,
90 torch::optim::Optimizer& opt,
91 const torch::Device& device)
92{
93 //cerr << train_input.get_tensor().sizes() << endl;
94 //cerr << train_gold.get_tensor().sizes() << endl;
95 //cerr << train_input.size() << endl;
96 const auto& input_tensor = train_input.get_tensor().to(device);
97 const auto& gold_tensor = train_gold.get_tensor().to(device);
98 int64_t n_samples = input_tensor.sizes()[1];
100 for (int64_t i = 0; i < n_samples; i += params.m_batch_size)
101 {
102 int64_t end_sample = cast_to_signed<decltype(n_samples)>(params.m_batch_size) > n_samples ? n_samples : i + params.m_batch_size;
103 assert(i < end_sample);
104 const auto batch_input = input_tensor.index({ Slice(), Slice(i, end_sample) }).to(device);
105 const auto batch_gold = gold_tensor.index({ Slice(), Slice(i, end_sample) }).to(device);
106
107 vector<TorchMatrix<int64_t>::tensor_t> batch_input_cat(train_input_cat.size());
108 for (size_t feat_idx = 0; feat_idx < train_input_cat.size(); ++feat_idx)
109 {
110 const auto& input_cat_tensor = train_input_cat[feat_idx].get_tensor().to(device);
111 batch_input_cat[feat_idx] = input_cat_tensor.index({ Slice(), Slice(i, end_sample) }).to(device);
112 }
113
114 train_batch({ "output" }, batch_input, batch_input_cat, batch_gold, opt, stat, device);
115 }
116
117 if (stat["output"].m_items > 0)
118 {
119 stat["output"].m_accuracy = float(stat["output"].m_correct) / stat["output"].m_items;
120 }
121 return stat;
122}
123
125 const vector<TorchMatrix<int64_t>>& gold,
126 nets::epoch_stat_t& stat,
127 const torch::Device& device)
128{
129 for (size_t i = 0; i < input.size(); ++i)
130 {
131 const TorchMatrix<int64_t>& input_bucket = input[i];
132 const TorchMatrix<int64_t>& gold_bucket = gold[i];
133 assert(input_bucket.get_max_feat() == gold_bucket.get_max_feat());
134
135 BiRnnSeq2SeqImpl::evaluate({ "output" }, input_bucket, gold_bucket, stat, device);
136 }
137}
138
139void Seq2SeqLemmatizerImpl::load(torch::serialize::InputArchive& archive)
140{
141 BiRnnSeq2SeqImpl::load(archive);
142 c10::IValue val;
143 archive.read("morph_model", val);
144 string serialized_morph_model = *(val.toString().get());
145 m_morph_model = morph_model::morph_model_t(serialized_morph_model);
146
147 archive.read("fixed_upos", val);
148 m_fixed_upos = *(val.toString().get());
149}
150
151void Seq2SeqLemmatizerImpl::save(torch::serialize::OutputArchive& archive) const
152{
153 BiRnnSeq2SeqImpl::save(archive);
154 string serialized_morph_model = m_morph_model.to_string();
155 archive.write("morph_model", serialized_morph_model);
156 archive.write("fixed_upos", m_fixed_upos);
157}
158
159} // namespace train
160} // namespace lemmatization
161} // namespace deeplima
162
uint64_t get_max_feat() const
const torch::Tensor & get_tensor() const
virtual void save(torch::serialize::OutputArchive &archive) const
void evaluate(const std::vector< TorchMatrix< int64_t > > &input, const std::vector< TorchMatrix< int64_t > > &gold, nets::epoch_stat_t &stat, const torch::Device &device=torch::Device(torch::kCPU))
void train(const train_params_lemmatization_t &params, const std::vector< TorchMatrix< int64_t > > &train_input, const std::vector< std::vector< TorchMatrix< int64_t > > > &train_input_cat, const std::vector< TorchMatrix< int64_t > > &train_gold, const std::vector< TorchMatrix< int64_t > > &eval_input, const std::vector< std::vector< TorchMatrix< int64_t > > > &eval_input_cat, const std::vector< TorchMatrix< int64_t > > &eval_gold, torch::optim::Optimizer &opt, const torch::Device &device=torch::Device(torch::kCPU))
virtual void load(torch::serialize::InputArchive &archive)
nets::epoch_stat_t train_on_subset(const train_params_lemmatization_t &params, const TorchMatrix< int64_t > &train_input, const std::vector< TorchMatrix< int64_t > > &train_input_cat, const TorchMatrix< int64_t > &train_gold, torch::optim::Optimizer &opt, const torch::Device &device=torch::Device(torch::kCPU))
Helper class for morphology data (upos, features) binarization.
Definition morph_model.h:93
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)
virtual void set_tags(const std::map< std::string, std::string > &tags)
virtual void to(torch::Device device, bool non_blocking=false) override
std::map< std::string, task_stat_t > epoch_stat_t
STL namespace.
std::map< std::string, std::string > m_tags