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)
29 ss <<
"encoder_embd_dropout = def Dropout prob=0.7" << std::endl;
31 size_t encoder_rnn_input_size = 0;
32 for (
size_t i = 0; i < encoder_embd_descr.size(); i++)
34 if (1 == encoder_embd_descr[i].m_type)
36 ss <<
"embd_" << encoder_embd_descr[i].m_name
37 <<
" = def Embedding dict=" << i
38 <<
" dim=" << encoder_embd_descr[i].m_dim << std::endl;
40 encoder_rnn_input_size += encoder_embd_descr[i].m_dim;
45 size_t encoder_input_size = encoder_rnn_input_size;
46 for (
size_t i = 0; i < encoder_rnn_descr.size(); i++)
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;
58 ss <<
"decoder_embd_dropout = def Dropout prob=0.7" << std::endl;
60 size_t decoder_rnn_input_size = 0;
61 for (
size_t i = 0; i < decoder_embd_descr.size(); i++)
63 if (1 == decoder_embd_descr[i].m_type)
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;
69 decoder_rnn_input_size += decoder_embd_descr[i].m_dim;
74 size_t decoder_input_size = decoder_rnn_input_size;
75 for (
size_t i = 0; i < decoder_rnn_descr.size(); i++)
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;
86 vector<string> output_names = {
"output" };
88 size_t input_size = decoder_input_size;
89 for (
size_t i = 0; i < output_names.size(); ++i)
91 ss <<
"fc_" << output_names[i] <<
" = def Linear input_size=" << input_size
92 <<
" output_size=" << n_output_classes << std::endl;
99 size_t cat_embd_input_dim = 0;
100 for (
size_t i = 0; i < cat_embd_descr.size(); i++)
102 if (1 == cat_embd_descr[i].m_type)
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;
108 cat_embd_input_dim += cat_embd_descr[i].m_dim;
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;
120 if (cat_embd_descr.size() > 0)
122 ss <<
"cat_embd_dropout = def Dropout prob=0.7" << std::endl << std::endl;
124 ss <<
"fc_cat2decoder = def Linear input_size=" << cat_embd_input_dim
125 <<
" output_size=" << cat_embd_input_dim << std::endl << std::endl;
129 ss <<
"cat_embd_2_encoder_dropout = def Dropout prob=0.7" << std::endl << std::endl;
131 ss <<
"fc_cat2encoder = def Linear input_size=" << cat_embd_input_dim
132 <<
" output_size=" << encoder_input_size << std::endl << std::endl;
136 for (
size_t i = 0; i < encoder_embd_descr.size(); i++)
138 ss << encoder_embd_descr[i].m_name <<
" = def Arg" << std::endl;
141 for (
size_t i = 0; i < decoder_embd_descr.size(); i++)
143 ss << decoder_embd_descr[i].m_name <<
" = def Arg" << std::endl;
146 for (
size_t i = 0; i < cat_embd_descr.size(); i++)
148 ss << cat_embd_descr[i].m_name <<
" = def Arg" << std::endl;
154 if (cat_embd_descr.size() > 0)
156 for (
size_t i = 0; i < cat_embd_descr.size(); i++)
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;
163 ss <<
"categories_embd_brut = cat input=";
164 for (
size_t i = 0; i < cat_embd_descr.size(); i++)
166 ss <<
"encoder_input_categories_" << cat_embd_descr[i].m_name;
167 if (i < cat_embd_descr.size() - 1)
172 ss <<
" dim=1" << std::endl;
175 ss <<
"categories_embd = forward module=cat_embd_dropout input=categories_embd_brut" << std::endl;
177 ss <<
"categories_encoded = forward module=fc_cat2decoder input=categories_embd" << endl;
180 ss <<
"categories_enc_embd = forward module=cat_embd_2_encoder_dropout input=categories_embd_brut" << 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;
189 for (
size_t i = 0; i < encoder_embd_descr.size(); i++)
191 if (1 == encoder_embd_descr[i].m_type)
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;
202 ss <<
"encoder_input_brut = cat input=";
203 for (
size_t i = 0; i < encoder_embd_descr.size(); i++)
205 if (0 == encoder_embd_descr[i].m_type)
207 ss << encoder_embd_descr[i].m_name;
209 else if (1 == encoder_embd_descr[i].m_type)
211 ss <<
"encoder_input_" << encoder_embd_descr[i].m_name;
213 if (i < encoder_embd_descr.size() - 1)
218 ss <<
" dim=2" << std::endl;
222 ss <<
"encoder_input = forward module=encoder_embd_dropout input=encoder_input_brut" << std::endl;
226 std::string last_output_name =
"encoder_input";
227 for (
size_t i = 0; i < encoder_rnn_descr.size(); i++)
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;
235 t <<
"encoder_rnn_out_" << i;
236 last_output_name = t.str();
247 ss <<
"encoder_final_concat = cat input=encoder_rnn_h_0,encoder_rnn_c_0 dim=2";
251 for (
size_t i = 0; i < 2 * encoder_rnn_descr.size(); i++)
253 ss <<
"encoder_final_l" << i;
254 if (i < 2 * encoder_rnn_descr.size() - 1)
259 ss <<
" = unbind input=encoder_final_concat dim=0" << endl;
261 ss <<
"encoder_final_state_ = cat input=";
262 for (
size_t i = 0; i < 2 * encoder_rnn_descr.size(); i++)
264 ss <<
"encoder_final_l" << i;
265 if (i < 2 * encoder_rnn_descr.size() - 1)
270 ss <<
" dim=1" << endl;
273 if (cat_embd_descr.size() > 0)
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;
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;
285 ss <<
"decoder_h0 = unsqueeze input=decoder_h0_ dim=0" << endl;
286 ss <<
"decoder_c0 = unsqueeze input=decoder_c0_ dim=0" << endl;
291 for (
size_t i = 0; i < decoder_embd_descr.size(); i++)
293 if (1 == decoder_embd_descr[i].m_type)
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;
304 ss <<
"decoder_input_brut = cat input=";
305 for (
size_t i = 0; i < decoder_embd_descr.size(); i++)
307 if (0 == decoder_embd_descr[i].m_type)
309 ss << decoder_embd_descr[i].m_name;
311 else if (1 == decoder_embd_descr[i].m_type)
313 ss <<
"decoder_input_" << decoder_embd_descr[i].m_name;
315 if (i < decoder_embd_descr.size() - 1)
320 ss <<
" dim=2" << std::endl;
324 ss <<
"decoder_input = forward module=decoder_embd_dropout input=decoder_input_brut" << std::endl;
328 last_output_name =
"decoder_input";
329 for (
size_t i = 0; i < decoder_rnn_descr.size(); i++)
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"
339 t <<
"decoder_rnn_out_" << i;
340 last_output_name = t.str();
345 for (
size_t i = 0; i < output_names.size(); ++i)
347 ss << output_names[i] <<
"_raw = forward module=fc_" << output_names[i]
348 <<
" input=" << last_output_name << std::endl;
350 ss << output_names[i] <<
" = log_softmax input=" << output_names[i] <<
"_raw" << std::endl;