LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_classifier_for_segmentation.cpp
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
7
8using namespace torch;
9using namespace std;
10
11namespace deeplima
12{
13namespace segmentation
14{
15namespace train
16{
17
18void BiRnnClassifierForSegmentationImpl::load(serialize::InputArchive& archive)
19{
20 BiRnnClassifierImpl::load(archive);
21
22 assert(m_ngram_descr.size() == 0);
23
24 c10::IValue v;
25 if (archive.try_read("ngram_descr", v))
26 {
27 if (!v.isList())
28 {
29 throw std::runtime_error("ngram_descr must be a list.");
30 }
31
32 const c10::List<c10::IValue>& l = v.toList();
33 m_ngram_descr.reserve(l.size());
34 for (size_t i = 0; i < l.size(); i++)
35 {
36 if (!l.get(i).isString())
37 {
38 throw std::runtime_error("ngram_descr must be a list of strings.");
39 }
40 const std::string& str = l.get(i).toStringRef();
41 m_ngram_descr.emplace_back(impl::ngram_descr_t(str));
42 }
43 }
44 else
45 {
46 throw std::runtime_error("Can't load ngram_descr.");
47 }
48}
49
50void BiRnnClassifierForSegmentationImpl::save(serialize::OutputArchive& archive) const
51{
52 BiRnnClassifierImpl::save(archive);
53
54 // Save ngram descriptions
55 c10::List<std::string> ngram_descr_list;
56 for (size_t i = 0; i < m_ngram_descr.size(); i++)
57 {
58 std::string ngram_str = m_ngram_descr[i].to_string();
59 assert(impl::ngram_descr_t(ngram_str) == m_ngram_descr[i]);
60 ngram_descr_list.push_back(ngram_str);
61 }
62
63 archive.write("ngram_descr", ngram_descr_list);
64}
65
66} // namespace train
67} // namespace segmentation
68} // namespace deeplima
69
virtual void load(torch::serialize::InputArchive &archive)
virtual void save(torch::serialize::OutputArchive &archive) const
STL namespace.