LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_seq_classifier.h
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#ifndef DEEPLIMA_LIBS_NN_BIRNN_SEQ_CLASSIFIER_H
7#define DEEPLIMA_LIBS_NN_BIRNN_SEQ_CLASSIFIER_H
8
9#include <tuple>
10#include <vector>
11
16
17namespace deeplima
18{
19namespace nets
20{
21
23{
24 uint64_t m_items;
25 uint64_t m_correct;
26 uint32_t m_num_classes;
27 double m_accuracy;
28 double m_loss;
30 double m_recall;
31 double m_f1;
32
35};
36
37typedef std::map<std::string, task_stat_t> epoch_stat_t;
38
40{
41public:
42
44
46 const std::vector<embd_descr_t>& embd_descr,
47 const std::vector<rnn_descr_t>& rnn_descr,
48 const std::vector<std::string>& output_names,
49 const std::vector<uint32_t>& classes,
50 float input_dropout_prob)
51 : StaticGraphImpl(dicts,
52 generate_script(embd_descr, rnn_descr, output_names, classes, input_dropout_prob)),
53 m_embd_descr(embd_descr)
54 {
55 }
56
58 const std::vector<embd_descr_t>& embd_descr,
59 const std::string& script)
60 : StaticGraphImpl(dicts, script),
61 m_embd_descr(embd_descr)
62 {
63 }
64
65 const std::vector<embd_descr_t>& get_embd_descr() const
66 {
67 return m_embd_descr;
68 }
69
70 virtual void load(torch::serialize::InputArchive& archive);
71 virtual void save(torch::serialize::OutputArchive& archive) const;
72
73 void train(size_t epochs,
74 size_t batch_size,
75 size_t seq_len,
76 const std::vector<std::string>& output_names,
77 const TorchMatrix<int64_t>& train_input,
78 const TorchMatrix<int64_t>& train_gold,
79 const TorchMatrix<int64_t>& eval_input,
80 const TorchMatrix<int64_t>& eval_gold,
81 torch::optim::Optimizer& opt,
82 const std::string& model_name = "",
83 const torch::Device& device = torch::Device(torch::kCPU));
84
85 void evaluate(const std::vector<std::string>& output_name,
86 const TorchMatrix<int64_t>& input,
87 const TorchMatrix<int64_t>& gold,
88 epoch_stat_t& stat,
89 const torch::Device& device = torch::Device(torch::kCPU));
90
91 torch::Tensor predict(const std::string& output_name,
92 const torch::Tensor& input,
93 const torch::Device& device = torch::Device(torch::kCPU));
94
95 torch::Tensor predict(const std::vector<std::string>& output_names,
96 const torch::Tensor& input,
97 const torch::Device& device = torch::Device(torch::kCPU));
98
99 void predict(size_t worker_id,
100 const torch::Tensor& inputs,
101 int64_t input_begin,
102 int64_t input_end,
103 int64_t output_begin,
104 int64_t output_end,
105 std::shared_ptr< StdMatrix<uint8_t> >& output,
106 const std::vector<std::string>& outputs_names,
107 const torch::Device& device = torch::Device(torch::kCPU));
108
109protected:
110
111 void split_input(const torch::Tensor& src,
112 std::map<std::string, torch::Tensor>& dst,
113 const torch::Device& device);
114
115 void evaluate(const std::vector<std::string>& output_names,
116 const std::map<std::string, torch::Tensor>& input,
117 const torch::Tensor& target,
118 epoch_stat_t& stat,
119 const torch::Device& device);
120
121 void train_epoch(size_t batch_size,
122 size_t seq_len,
123 const std::vector<std::string>& output_names,
124 const torch::Tensor& input_batches,
125 const torch::Tensor& gold_batches,
126 torch::optim::Optimizer& opt,
127 epoch_stat_t& stat,
128 const torch::Device& device);
129
130 void train_batch(size_t batch_size,
131 size_t seq_len,
132 const std::vector<std::string>& output_names,
133 const torch::Tensor& input,
134 const torch::Tensor& gold,
135 torch::optim::Optimizer& opt,
136 epoch_stat_t& stat,
137 const torch::Device& device);
138
139 void train_batch(size_t batch_size,
140 size_t seq_len,
141 const std::vector<std::string>& output_names,
142 const std::map<std::string, torch::Tensor>& input,
143 const torch::Tensor& target,
144 torch::optim::Optimizer& opt,
145 epoch_stat_t& stat,
146 const torch::Device& device);
147
148 static std::string generate_script(const std::vector<embd_descr_t>& embd_descr,
149 const std::vector<rnn_descr_t>& rnn_descr,
150 const std::vector<std::string>& output_names,
151 const std::vector<uint32_t>& classes,
152 float input_dropout_prob);
153
154 std::vector<embd_descr_t> m_embd_descr;
155};
156
157inline torch::serialize::OutputArchive& operator<<(
158 torch::serialize::OutputArchive& archive,
159 const BiRnnClassifierImpl& module)
160{
161 module.save(archive);
162 return archive;
163}
164
165TORCH_MODULE(BiRnnClassifier);
166
167} // namespace nets
168} // namespace deeplima
169
170#endif
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
const std::vector< embd_descr_t > & get_embd_descr() const
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))
BiRnnClassifierImpl(DictsHolder &&dicts, const std::vector< embd_descr_t > &embd_descr, const std::string &script)
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)
BiRnnClassifierImpl(DictsHolder &&dicts, const std::vector< embd_descr_t > &embd_descr, const std::vector< rnn_descr_t > &rnn_descr, const std::vector< std::string > &output_names, const std::vector< uint32_t > &classes, float input_dropout_prob)
static std::string generate_script(const std::vector< embd_descr_t > &embd_descr, const std::vector< rnn_descr_t > &rnn_descr, const std::vector< std::string > &output_names, const std::vector< uint32_t > &classes, float input_dropout_prob)
Implementation for Torch of the Tensorflow execution graph Used in training only.
torch::serialize::OutputArchive & operator<<(torch::serialize::OutputArchive &archive, const BiRnnSeq2SeqImpl &module)
std::map< std::string, task_stat_t > epoch_stat_t
TORCH_MODULE(BiRnnSeq2Seq)