LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_and_deep_biaffine_attention.h
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
6#ifndef DEEPLIMA_LIBS_TASKS_GRAPH_DP_BIRNN_AND_DEEP_BIAFFINE_ATTENTION_H
7#define DEEPLIMA_LIBS_TASKS_GRAPH_DP_BIRNN_AND_DEEP_BIAFFINE_ATTENTION_H
8
12#include "iterable_dataset.h"
13
14namespace deeplima
15{
16namespace graph_dp
17{
18namespace train
19{
20
22{
23public:
24 typedef torch::Tensor tensor_t;
26
31
33 const std::vector<nets::embd_descr_t>& embd_descr,
34 const std::vector<nets::rnn_descr_t>& rnn_descr,
35 const std::vector<nets::deep_biaffine_attention_descr_t>& decoder_descr,
36 const std::vector<std::string>& output_names,
37 DictsHolder&& classes,
38 const std::string& embd_fn,
39 bool input_includes_root,
40 int64_t num_labels = 0,
41 const std::vector<std::string>& rel_class_names = {})
42 : BiRnnClassifierImpl(std::move(dicts),
43 embd_descr,
44 generate_script(embd_descr, rnn_descr, decoder_descr, output_names, input_includes_root, num_labels)
45 /*rnn_descr, output_names, classes.get_counters()*/),
46 m_workers(0),
47 m_num_labels(num_labels),
48 m_rel_class_names(rel_class_names),
49 m_output_class_names(output_names),
50 m_embd_fn(embd_fn)
51 {
52 m_input_classes = classes;
53 m_input_class_names.reserve(embd_descr.size());
54 for ( const auto& d : embd_descr )
55 {
56 m_input_class_names.emplace_back(d.m_name);
57 }
58 }
59
60 virtual void load(torch::serialize::InputArchive& archive);
61 virtual void save(torch::serialize::OutputArchive& archive) const;
62
63 void load(const std::string& fn)
64 {
65 torch::load(*this, fn, torch::Device(torch::kCPU));
66 }
67
68 size_t init_new_worker(size_t )
69 {
70 return m_workers++;
71 }
72
73 void train(const train_params_graph_dp_t& params,
74 const std::vector<std::string>& output_names,
75 const IterableDataSet& train_batches,
76 const IterableDataSet& eval_batches,
77 torch::optim::Optimizer& opt,
78 double& best_eval_accuracy,
79 const torch::Device& device = torch::Device(torch::kCPU));
80
81 void evaluate(const std::vector<std::string>& output_names,
82 std::shared_ptr<BatchIterator> dataset_iterator,
84 const torch::Device& device = torch::Device(torch::kCPU));
85
86 void predict(size_t worker_id,
87 const torch::Tensor& inputs,
88 int64_t input_begin,
89 int64_t input_end,
90 int64_t output_begin,
91 int64_t output_end,
92 std::shared_ptr< StdMatrix<uint8_t> >& output,
93 const std::vector<std::string>& outputs_names,
94 const torch::Device& device = torch::Device(torch::kCPU));
95
96 const DictsHolder& get_classes() const
97 {
98 return m_output_classes;
99 }
100
101 const std::vector<std::string>& get_output_class_names() const
102 {
104 }
105
106 const std::vector<std::string>& get_input_class_names() const
107 {
108 return m_input_class_names;
109 }
110
111 // deprel id -> string mapping (empty if the model has no label decoder)
112 const std::vector<std::string>& get_rel_class_names() const
113 {
114 return m_rel_class_names;
115 }
116
117 const std::string& get_embd_fn([[maybe_unused]] size_t idx) const
118 {
119 assert(0 == idx);
120 return m_embd_fn;
121 }
122
123protected:
124 void train_epoch(size_t batch_size,
125 size_t seq_len,
126 const std::vector<std::string>& output_names,
127 std::shared_ptr<BatchIterator> train_iterator,
128 torch::optim::Optimizer& opt,
129 nets::epoch_stat_t& stat,
130 const torch::Device& device);
131
132 void train_batch(size_t batch_size,
133 size_t seq_len,
134 const std::vector<std::string>& output_names,
135 const torch::Tensor& trainable_input,
136 const torch::Tensor& nontrainable_input,
137 const torch::Tensor& gold,
138 torch::optim::Optimizer& opt,
139 nets::epoch_stat_t& stat,
140 const torch::Device& device);
141
142 void evaluate(const std::vector<std::string>& output_names,
143 const torch::Tensor& trainable_input,
144 const torch::Tensor& nontrainable_input,
145 const torch::Tensor& gold,
146 nets::epoch_stat_t& stat,
147 const torch::Device& device);
148
149 static std::string generate_script(const std::vector<nets::embd_descr_t>& embd_descr,
150 const std::vector<nets::rnn_descr_t>& rnn_descr,
151 const std::vector<nets::deep_biaffine_attention_descr_t>& decoder_descr,
152 const std::vector<std::string>& output_names,
153 bool input_includes_root=false,
154 int64_t num_labels=0/*,
155 const std::vector<uint32_t>& classes*/);
156
157 size_t m_workers;
158 int64_t m_num_labels = 0; // number of deprel classes; 0 = label decoder disabled
159 std::vector<std::string> m_rel_class_names; // deprel id -> string; saved with the model
160 std::vector<std::string> m_input_class_names;
162 std::vector<std::string> m_output_class_names;
164 std::string m_embd_fn;
165};
166
167inline torch::serialize::OutputArchive& operator<<(
168 torch::serialize::OutputArchive& archive,
170{
171 module.save(archive);
172 return archive;
173}
174
175inline torch::serialize::InputArchive& operator>>(
176 torch::serialize::InputArchive& archive,
178{
179 module.load(archive);
180 return archive;
181}
182
183TORCH_MODULE(BiRnnAndDeepBiaffineAttention);
184
185} // train
186} // graph_dp
187} // deeplima
188
189#endif // DEEPLIMA_LIBS_TASKS_GRAPH_DP_BIRNN_AND_DEEP_BIAFFINE_ATTENTION_H
BiRnnAndDeepBiaffineAttentionImpl(DictsHolder &&dicts, const std::vector< nets::embd_descr_t > &embd_descr, const std::vector< nets::rnn_descr_t > &rnn_descr, const std::vector< nets::deep_biaffine_attention_descr_t > &decoder_descr, const std::vector< std::string > &output_names, DictsHolder &&classes, const std::string &embd_fn, bool input_includes_root, int64_t num_labels=0, const std::vector< std::string > &rel_class_names={})
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))
static std::string generate_script(const std::vector< nets::embd_descr_t > &embd_descr, const std::vector< nets::rnn_descr_t > &rnn_descr, const std::vector< nets::deep_biaffine_attention_descr_t > &decoder_descr, const std::vector< std::string > &output_names, bool input_includes_root=false, int64_t num_labels=0)
torch::serialize::OutputArchive & operator<<(torch::serialize::OutputArchive &archive, const BiRnnAndDeepBiaffineAttentionImpl &module)
TORCH_MODULE(BiRnnAndDeepBiaffineAttention)
torch::serialize::InputArchive & operator>>(torch::serialize::InputArchive &archive, BiRnnAndDeepBiaffineAttentionImpl &module)
std::map< std::string, task_stat_t > epoch_stat_t
STL namespace.