LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_classifier_script_generator.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
6#include <sstream>
7
9
10using namespace std;
11
12namespace deeplima
13{
14namespace nets
15{
16
17string BiRnnClassifierImpl::generate_script(const std::vector<embd_descr_t>& embd_descr,
18 const std::vector<rnn_descr_t>& rnn_descr,
19 const std::vector<std::string>& output_names,
20 const std::vector<uint32_t>& classes,
21 float input_dropout_prob)
22{
23 stringstream ss;
24
25 ss << "input_dropout = def Dropout prob=" << input_dropout_prob << std::endl;
26
27 size_t rnn_input_size = 0;
28 for (size_t i = 0; i < embd_descr.size(); i++)
29 {
30 if (1 == embd_descr[i].m_type)
31 {
32 ss << "embd_" << embd_descr[i].m_name << " = def Embedding dict=" << i
33 << " dim=" << embd_descr[i].m_dim << std::endl;
34 }
35 rnn_input_size += embd_descr[i].m_dim;
36 }
37
38 ss << std::endl;
39
40 size_t input_size = rnn_input_size;
41 for (size_t i = 0; i < rnn_descr.size(); i++)
42 {
43 ss << "lstm_" << i << " = def LSTM input_size=" << input_size
44 << " hidden_size=" << rnn_descr[i].m_dim << " num_layers=1 bidirectional=true" << std::endl;
45 input_size = rnn_descr[i].m_dim * 2;
46 }
47
48 ss << std::endl;
49
50 for (size_t i = 0; i < output_names.size(); ++i)
51 {
52 ss << "fc_" << output_names[i] << " = def Linear input_size=" << input_size
53 << " output_size=" << classes[i] << std::endl;
54
55 ss << std::endl;
56 }
57
58 for (size_t i = 0; i < embd_descr.size(); i++)
59 {
60 ss << embd_descr[i].m_name << " = def Arg" << std::endl;
61 }
62
63 ss << std::endl;
64
65 for (size_t i = 0; i < embd_descr.size(); i++)
66 {
67 if (1 == embd_descr[i].m_type)
68 {
69 ss << "input_" << embd_descr[i].m_name << " = forward module=embd_"
70 << embd_descr[i].m_name
71 << " input=" << embd_descr[i].m_name << std::endl;
72 }
73 }
74
75 ss << std::endl;
76
77 ss << "input_brut = cat input=";
78 for (size_t i = 0; i < embd_descr.size(); i++)
79 {
80 if (0 == embd_descr[i].m_type)
81 {
82 ss << embd_descr[i].m_name;
83 }
84 else if (1 == embd_descr[i].m_type)
85 {
86 ss << "input_" << embd_descr[i].m_name;
87 }
88 if (i < embd_descr.size() - 1)
89 {
90 ss << ",";
91 }
92 }
93 ss << " dim=2" << std::endl;
94
95 ss << std::endl;
96
97 ss << "input = forward module=input_dropout input=input_brut" << std::endl;
98
99 ss << std::endl;
100
101 std::string last_output_name = "input";
102 for (size_t i = 0; i < rnn_descr.size(); i++)
103 {
104 ss << "rnn_out_" << i << " = forward module=lstm_" << i
105 << " input=" << last_output_name << std::endl;
106 std::stringstream t;
107 t << "rnn_out_" << i;
108 last_output_name = t.str();
109 }
110
111 ss << std::endl;
112
113 for (size_t i = 0; i < output_names.size(); ++i)
114 {
115 ss << output_names[i] << "_raw = forward module=fc_" << output_names[i]
116 << " input=" << last_output_name << std::endl;
117
118 ss << output_names[i] << " = log_softmax input=" << output_names[i] << "_raw" << std::endl;
119 ss << endl;
120 }
121
122 return ss.str();
123}
124
125} // namespace nets
126} // namespace deeplima
127
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)
STL namespace.