LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_classifier_for_segmentation.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_SRC_TRAIN_BIRNN_CLASSIFIER_FOR_SEGMENTATION_H
7#define DEEPLIMA_SRC_TRAIN_BIRNN_CLASSIFIER_FOR_SEGMENTATION_H
8
11
12namespace deeplima
13{
14namespace segmentation
15{
16namespace train
17{
18
20{
21public:
22 typedef torch::Tensor tensor_t;
24
29
31 const std::vector<impl::ngram_descr_t>& ngram_descr,
32 const std::vector<nets::embd_descr_t>& embd_descr,
33 const std::vector<nets::rnn_descr_t>& rnn_descr,
34 const std::string& output_name,
35 uint32_t num_classes,
36 float input_dropout_prob)
37 : BiRnnClassifierImpl(std::move(dicts),
38 embd_descr,
39 rnn_descr,
40 { output_name },
41 { num_classes },
42 input_dropout_prob),
43 m_ngram_descr(ngram_descr),
44 m_workers(0)
45 {
46 }
47
48 virtual void load(torch::serialize::InputArchive& archive);
49 virtual void save(torch::serialize::OutputArchive& archive) const;
50
51 void load(const std::string& fn)
52 {
53 torch::load(*this, fn);
54 }
55
56 size_t init_new_worker(size_t /*input_len*/)
57 {
58 return m_workers++;
59 }
60
61 const std::vector<impl::ngram_descr_t>& get_ngram_descr() const
62 {
63 return m_ngram_descr;
64 }
65
66protected:
67 std::vector<impl::ngram_descr_t> m_ngram_descr;
68 size_t m_workers;
69};
70
71inline torch::serialize::OutputArchive& operator<<(
72 torch::serialize::OutputArchive& archive,
74{
75 module.save(archive);
76 return archive;
77}
78
79inline torch::serialize::InputArchive& operator>>(
80 torch::serialize::InputArchive& archive,
82{
83 module.load(archive);
84 return archive;
85}
86
87TORCH_MODULE(BiRnnClassifierForSegmentation);
88
89} // namespace train
90} // namespace segmentation
91} // namespace deeplima
92
93#endif
94
BiRnnClassifierForSegmentationImpl(DictsHolder &&dicts, const std::vector< impl::ngram_descr_t > &ngram_descr, const std::vector< nets::embd_descr_t > &embd_descr, const std::vector< nets::rnn_descr_t > &rnn_descr, const std::string &output_name, uint32_t num_classes, float input_dropout_prob)
virtual void load(torch::serialize::InputArchive &archive)
virtual void save(torch::serialize::OutputArchive &archive) const
torch::serialize::OutputArchive & operator<<(torch::serialize::OutputArchive &archive, const BiRnnClassifierForSegmentationImpl &module)
TORCH_MODULE(BiRnnClassifierForSegmentation)
torch::serialize::InputArchive & operator>>(torch::serialize::InputArchive &archive, BiRnnClassifierForSegmentationImpl &module)
STL namespace.