LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_seq2seq.h
Go to the documentation of this file.
1// Copyright 2002-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_SEQ2SEQ_H
7#define DEEPLIMA_LIBS_NN_BIRNN_SEQ2SEQ_H
8
10
11namespace deeplima
12{
13namespace nets
14{
15
17{
18public:
19 typedef torch::Tensor tensor_t;
20
22
24 const std::vector<embd_descr_t>& encoder_embd_descr,
25 const std::vector<rnn_descr_t>& encoder_rnn_descr,
26 const std::vector<embd_descr_t>& decoder_embd_descr,
27 const std::vector<rnn_descr_t>& decoder_rnn_descr,
28 const std::vector<embd_descr_t>& cat_embd_descr)
29 : BiRnnClassifierImpl(std::move(dicts), concatenate(encoder_embd_descr, decoder_embd_descr, cat_embd_descr),
30 generate_script(encoder_embd_descr, encoder_rnn_descr,
31 decoder_embd_descr, decoder_rnn_descr,
32 cat_embd_descr, dicts[1]->size())),
33 m_cat_embd_descr(cat_embd_descr)
34 {
35 }
36
37 virtual void load(torch::serialize::InputArchive& archive);
38 virtual void save(torch::serialize::OutputArchive& archive) const;
39
40 void evaluate(const std::vector<std::string>& output_names,
41 const TorchMatrix<int64_t>& input,
42 const TorchMatrix<int64_t>& gold,
43 epoch_stat_t& stat,
44 const torch::Device& device = torch::Device(torch::kCPU));
45
46protected:
47 std::string generate_script(const std::vector<embd_descr_t>& encoder_embd_descr,
48 const std::vector<rnn_descr_t>& encoder_rnn_descr,
49 const std::vector<embd_descr_t>& decoder_embd_descr,
50 const std::vector<rnn_descr_t>& decoder_rnn_descr,
51 const std::vector<embd_descr_t>& cat_embd_descr,
52 size_t n_output_classes);
53
54 void train_batch(const std::vector<std::string>& output_names,
55 const torch::Tensor& input,
56 const std::vector<torch::Tensor>& input_cat,
57 const torch::Tensor& target,
58 torch::optim::Optimizer& opt,
59 epoch_stat_t& stat,
60 const torch::Device& device);
61
62 template <class T>
63 static std::vector<T> concatenate(const std::vector<T>& a,
64 const std::vector<T>& b,
65 const std::vector<T>& c = {})
66 {
67 std::vector<T> out = a;
68 out.insert(out.end(), b.begin(), b.end());
69 out.insert(out.end(), c.begin(), c.end());
70 return out;
71 }
72
73 std::vector<embd_descr_t> m_cat_embd_descr;
74
75};
76
77inline torch::serialize::OutputArchive& operator<<(
78 torch::serialize::OutputArchive& archive,
79 const BiRnnSeq2SeqImpl& module)
80{
81 module.save(archive);
82 return archive;
83}
84
85inline torch::serialize::InputArchive& operator>>(
86 torch::serialize::InputArchive& archive,
87 BiRnnSeq2SeqImpl& module)
88{
89 module.load(archive);
90 return archive;
91}
92
93TORCH_MODULE(BiRnnSeq2Seq);
94
95} // namespace nets
96} // namespace deeplima
97
98#endif // DEEPLIMA_LIBS_NN_BIRNN_SEQ2SEQ_H
std::string generate_script(const std::vector< embd_descr_t > &encoder_embd_descr, const std::vector< rnn_descr_t > &encoder_rnn_descr, const std::vector< embd_descr_t > &decoder_embd_descr, const std::vector< rnn_descr_t > &decoder_rnn_descr, const std::vector< embd_descr_t > &cat_embd_descr, size_t n_output_classes)
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)
static std::vector< T > concatenate(const std::vector< T > &a, const std::vector< T > &b, const std::vector< T > &c={})
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))
BiRnnSeq2SeqImpl(DictsHolder &&dicts, const std::vector< embd_descr_t > &encoder_embd_descr, const std::vector< rnn_descr_t > &encoder_rnn_descr, const std::vector< embd_descr_t > &decoder_embd_descr, const std::vector< rnn_descr_t > &decoder_rnn_descr, const std::vector< embd_descr_t > &cat_embd_descr)
torch::serialize::OutputArchive & operator<<(torch::serialize::OutputArchive &archive, const BiRnnSeq2SeqImpl &module)
std::map< std::string, task_stat_t > epoch_stat_t
torch::serialize::InputArchive & operator>>(torch::serialize::InputArchive &archive, BiRnnSeq2SeqImpl &module)
TORCH_MODULE(BiRnnSeq2Seq)
STL namespace.