LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_seq2seq_script_generator.cpp
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#include <string>
7#include <vector>
8#include <sstream>
9
10#include "birnn_seq2seq.h"
11
12using namespace std;
13
14namespace deeplima
15{
16namespace nets
17{
18
19string BiRnnSeq2SeqImpl::generate_script(const vector<embd_descr_t>& encoder_embd_descr,
20 const vector<rnn_descr_t>& encoder_rnn_descr,
21 const vector<embd_descr_t>& decoder_embd_descr,
22 const vector<rnn_descr_t>& decoder_rnn_descr,
23 const vector<embd_descr_t>& cat_embd_descr,
24 size_t n_output_classes)
25{
26 stringstream ss;
27
28 // Encoder definitions
29 ss << "encoder_embd_dropout = def Dropout prob=0.7" << std::endl;
30
31 size_t encoder_rnn_input_size = 0;
32 for (size_t i = 0; i < encoder_embd_descr.size(); i++)
33 {
34 if (1 == encoder_embd_descr[i].m_type)
35 {
36 ss << "embd_" << encoder_embd_descr[i].m_name
37 << " = def Embedding dict=" << i
38 << " dim=" << encoder_embd_descr[i].m_dim << std::endl;
39 }
40 encoder_rnn_input_size += encoder_embd_descr[i].m_dim;
41 }
42
43 ss << std::endl;
44
45 size_t encoder_input_size = encoder_rnn_input_size;
46 for (size_t i = 0; i < encoder_rnn_descr.size(); i++)
47 {
48 ss << "encoder_lstm_" << i
49 << " = def LSTM input_size=" << encoder_input_size
50 << " hidden_size=" << encoder_rnn_descr[i].m_dim
51 << " num_layers=1 bidirectional=true" << std::endl;
52 encoder_input_size = encoder_rnn_descr[i].m_dim * 2;
53 }
54
55 ss << std::endl;
56
57 // Decoder definitions
58 ss << "decoder_embd_dropout = def Dropout prob=0.7" << std::endl;
59
60 size_t decoder_rnn_input_size = 0;
61 for (size_t i = 0; i < decoder_embd_descr.size(); i++)
62 {
63 if (1 == decoder_embd_descr[i].m_type)
64 {
65 ss << "embd_" << decoder_embd_descr[i].m_name
66 << " = def Embedding dict=" << encoder_embd_descr.size() + i
67 << " dim=" << decoder_embd_descr[i].m_dim << std::endl;
68 }
69 decoder_rnn_input_size += decoder_embd_descr[i].m_dim;
70 }
71
72 ss << std::endl;
73
74 size_t decoder_input_size = decoder_rnn_input_size;
75 for (size_t i = 0; i < decoder_rnn_descr.size(); i++)
76 {
77 ss << "decoder_lstm_" << i
78 << " = def LSTM input_size=" << decoder_input_size
79 << " hidden_size=" << decoder_rnn_descr[i].m_dim
80 << " num_layers=1 bidirectional=false" << std::endl;
81 decoder_input_size = decoder_rnn_descr[i].m_dim;
82 }
83
84 ss << std::endl;
85
86 vector<string> output_names = { "output" };
87
88 size_t input_size = decoder_input_size;
89 for (size_t i = 0; i < output_names.size(); ++i)
90 {
91 ss << "fc_" << output_names[i] << " = def Linear input_size=" << input_size
92 << " output_size=" << n_output_classes << std::endl;
93
94 ss << std::endl;
95 }
96
97 ss << std::endl;
98
99 size_t cat_embd_input_dim = 0;
100 for (size_t i = 0; i < cat_embd_descr.size(); i++)
101 {
102 if (1 == cat_embd_descr[i].m_type)
103 {
104 ss << "embd_" << cat_embd_descr[i].m_name
105 << " = def Embedding dict=" << encoder_embd_descr.size() + decoder_embd_descr.size() + i
106 << " dim=" << cat_embd_descr[i].m_dim << std::endl;
107 }
108 cat_embd_input_dim += cat_embd_descr[i].m_dim;
109 }
110
111 ss << std::endl;
112
113 ss << "interm_fc_h0 = def Linear input_size=" << encoder_input_size * 2 + cat_embd_input_dim
114 << " output_size=" << decoder_rnn_descr[0].m_dim << std::endl;
115 ss << "interm_fc_c0 = def Linear input_size=" << encoder_input_size * 2 + cat_embd_input_dim
116 << " output_size=" << decoder_rnn_descr[0].m_dim << std::endl;
117
118 ss << std::endl;
119
120 if (cat_embd_descr.size() > 0)
121 {
122 ss << "cat_embd_dropout = def Dropout prob=0.7" << std::endl << std::endl;
123
124 ss << "fc_cat2decoder = def Linear input_size=" << cat_embd_input_dim
125 << " output_size=" << cat_embd_input_dim << std::endl << std::endl;
126
127
128 // present categories as initial state of encoder
129 ss << "cat_embd_2_encoder_dropout = def Dropout prob=0.7" << std::endl << std::endl;
130
131 ss << "fc_cat2encoder = def Linear input_size=" << cat_embd_input_dim
132 << " output_size=" << encoder_input_size << std::endl << std::endl;
133 }
134
135 // Args definitions
136 for (size_t i = 0; i < encoder_embd_descr.size(); i++)
137 {
138 ss << encoder_embd_descr[i].m_name << " = def Arg" << std::endl;
139 }
140
141 for (size_t i = 0; i < decoder_embd_descr.size(); i++)
142 {
143 ss << decoder_embd_descr[i].m_name << " = def Arg" << std::endl;
144 }
145
146 for (size_t i = 0; i < cat_embd_descr.size(); i++)
147 {
148 ss << cat_embd_descr[i].m_name << " = def Arg" << std::endl;
149 }
150
151 ss << std::endl;
152
153 // Encoder forward pass
154 if (cat_embd_descr.size() > 0)
155 {
156 for (size_t i = 0; i < cat_embd_descr.size(); i++)
157 {
158 ss << "encoder_input_categories_" << cat_embd_descr[i].m_name
159 << " = forward module=embd_" << cat_embd_descr[i].m_name
160 << " input=" << cat_embd_descr[i].m_name << std::endl;
161 }
162
163 ss << "categories_embd_brut = cat input=";
164 for (size_t i = 0; i < cat_embd_descr.size(); i++)
165 {
166 ss << "encoder_input_categories_" << cat_embd_descr[i].m_name;
167 if (i < cat_embd_descr.size() - 1)
168 {
169 ss << ",";
170 }
171 }
172 ss << " dim=1" << std::endl;
173
174 // for decoder
175 ss << "categories_embd = forward module=cat_embd_dropout input=categories_embd_brut" << std::endl;
176 ss << std::endl;
177 ss << "categories_encoded = forward module=fc_cat2decoder input=categories_embd" << endl;
178
179 // for encoder
180 ss << "categories_enc_embd = forward module=cat_embd_2_encoder_dropout input=categories_embd_brut" << std::endl;
181 ss << std::endl;
182 ss << "encoder_init_state_ = forward module=fc_cat2encoder input=categories_enc_embd" << endl;
183 ss << "encoder_init_state = reshape input=encoder_init_state_ dims=2,-1,"
184 << encoder_input_size / 2 << endl;
185 }
186
187 ss << std::endl;
188
189 for (size_t i = 0; i < encoder_embd_descr.size(); i++)
190 {
191 if (1 == encoder_embd_descr[i].m_type)
192 {
193 ss << "encoder_input_" << encoder_embd_descr[i].m_name
194 << " = forward module=embd_"
195 << encoder_embd_descr[i].m_name
196 << " input=" << encoder_embd_descr[i].m_name << std::endl;
197 }
198 }
199
200 ss << std::endl;
201
202 ss << "encoder_input_brut = cat input=";
203 for (size_t i = 0; i < encoder_embd_descr.size(); i++)
204 {
205 if (0 == encoder_embd_descr[i].m_type)
206 {
207 ss << encoder_embd_descr[i].m_name;
208 }
209 else if (1 == encoder_embd_descr[i].m_type)
210 {
211 ss << "encoder_input_" << encoder_embd_descr[i].m_name;
212 }
213 if (i < encoder_embd_descr.size() - 1)
214 {
215 ss << ",";
216 }
217 }
218 ss << " dim=2" << std::endl;
219
220 ss << std::endl;
221
222 ss << "encoder_input = forward module=encoder_embd_dropout input=encoder_input_brut" << std::endl;
223
224 ss << std::endl;
225
226 std::string last_output_name = "encoder_input";
227 for (size_t i = 0; i < encoder_rnn_descr.size(); i++)
228 {
229 ss << "encoder_rnn_out_" << i
230 << ",encoder_rnn_h_" << i
231 << ",encoder_rnn_c_" << i
232 << " = forward module=encoder_lstm_" << i
233 << " input=" << last_output_name << ",encoder_init_state,encoder_init_state" << std::endl;
234 std::stringstream t;
235 t << "encoder_rnn_out_" << i;
236 last_output_name = t.str();
237 }
238
239 // encoder_rnn_h_X = [2*num_layers, batch_size, enc_hidden_size]
240 // encoder_rnn_c_X = [2*num_layers, batch_size, enc_hidden_size]
241
242 // decoder_rnn_h_X = [1*num_layers, batch_size, dec_hidden_size]
243 // decoder_rnn_c_X = [1*num_layers, batch_size, dec_hidden_size]
244
245 ss << std::endl;
246
247 ss << "encoder_final_concat = cat input=encoder_rnn_h_0,encoder_rnn_c_0 dim=2";
248 ss << std::endl;
249 // encoder_final_concat = [2*num_layers, batch_size, 2*enc_hidden_size]
250
251 for (size_t i = 0; i < 2 * encoder_rnn_descr.size(); i++)
252 {
253 ss << "encoder_final_l" << i;
254 if (i < 2 * encoder_rnn_descr.size() - 1)
255 {
256 ss << ",";
257 }
258 }
259 ss << " = unbind input=encoder_final_concat dim=0" << endl;
260
261 ss << "encoder_final_state_ = cat input=";
262 for (size_t i = 0; i < 2 * encoder_rnn_descr.size(); i++)
263 {
264 ss << "encoder_final_l" << i;
265 if (i < 2 * encoder_rnn_descr.size() - 1)
266 {
267 ss << ",";
268 }
269 }
270 ss << " dim=1" << endl;
271
272 // encoder_final_state = [batch_size, 4*enc_hidden_size]
273 if (cat_embd_descr.size() > 0)
274 {
275 ss << "encoder_final_state = cat input=encoder_final_state_,categories_encoded dim=1" << std::endl << std::endl;
276 ss << "decoder_h0_ = forward module=interm_fc_h0 input=encoder_final_state" << endl;
277 ss << "decoder_c0_ = forward module=interm_fc_c0 input=encoder_final_state" << endl;
278 }
279 else
280 {
281 ss << "decoder_h0_ = forward module=interm_fc_h0 input=encoder_final_state_" << endl;
282 ss << "decoder_c0_ = forward module=interm_fc_c0 input=encoder_final_state_" << endl;
283 }
284
285 ss << "decoder_h0 = unsqueeze input=decoder_h0_ dim=0" << endl;
286 ss << "decoder_c0 = unsqueeze input=decoder_c0_ dim=0" << endl;
287
288 ss << std::endl;
289
290 // Decoder forward pass
291 for (size_t i = 0; i < decoder_embd_descr.size(); i++)
292 {
293 if (1 == decoder_embd_descr[i].m_type)
294 {
295 ss << "decoder_input_" << decoder_embd_descr[i].m_name
296 << " = forward module=embd_"
297 << decoder_embd_descr[i].m_name
298 << " input=" << decoder_embd_descr[i].m_name << std::endl;
299 }
300 }
301
302 ss << std::endl;
303
304 ss << "decoder_input_brut = cat input=";
305 for (size_t i = 0; i < decoder_embd_descr.size(); i++)
306 {
307 if (0 == decoder_embd_descr[i].m_type)
308 {
309 ss << decoder_embd_descr[i].m_name;
310 }
311 else if (1 == decoder_embd_descr[i].m_type)
312 {
313 ss << "decoder_input_" << decoder_embd_descr[i].m_name;
314 }
315 if (i < decoder_embd_descr.size() - 1)
316 {
317 ss << ",";
318 }
319 }
320 ss << " dim=2" << std::endl;
321
322 ss << std::endl;
323
324 ss << "decoder_input = forward module=decoder_embd_dropout input=decoder_input_brut" << std::endl;
325
326 ss << std::endl;
327
328 last_output_name = "decoder_input";
329 for (size_t i = 0; i < decoder_rnn_descr.size(); i++)
330 {
331 ss << "decoder_rnn_out_" << i
332 << ",decoder_rnn_h_" << i
333 << ",decoder_rnn_c_" << i
334 << " = forward module=decoder_lstm_" << i
335 << " input=" << last_output_name
336 << ",decoder_h0,decoder_c0" // order: input,h0,c0
337 << std::endl;
338 std::stringstream t;
339 t << "decoder_rnn_out_" << i;
340 last_output_name = t.str();
341 }
342
343 ss << std::endl;
344
345 for (size_t i = 0; i < output_names.size(); ++i)
346 {
347 ss << output_names[i] << "_raw = forward module=fc_" << output_names[i]
348 << " input=" << last_output_name << std::endl;
349
350 ss << output_names[i] << " = log_softmax input=" << output_names[i] << "_raw" << std::endl;
351 ss << endl;
352 }
353
354 return ss.str();
355}
356
357} // namespace nets
358} // namespace deeplima
359
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)
STL namespace.