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)
25 ss <<
"input_dropout = def Dropout prob=" << input_dropout_prob << std::endl;
27 size_t rnn_input_size = 0;
28 for (
size_t i = 0; i < embd_descr.size(); i++)
30 if (1 == embd_descr[i].m_type)
32 ss <<
"embd_" << embd_descr[i].m_name <<
" = def Embedding dict=" << i
33 <<
" dim=" << embd_descr[i].m_dim << std::endl;
35 rnn_input_size += embd_descr[i].m_dim;
40 size_t input_size = rnn_input_size;
41 for (
size_t i = 0; i < rnn_descr.size(); i++)
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;
50 for (
size_t i = 0; i < output_names.size(); ++i)
52 ss <<
"fc_" << output_names[i] <<
" = def Linear input_size=" << input_size
53 <<
" output_size=" << classes[i] << std::endl;
58 for (
size_t i = 0; i < embd_descr.size(); i++)
60 ss << embd_descr[i].m_name <<
" = def Arg" << std::endl;
65 for (
size_t i = 0; i < embd_descr.size(); i++)
67 if (1 == embd_descr[i].m_type)
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;
77 ss <<
"input_brut = cat input=";
78 for (
size_t i = 0; i < embd_descr.size(); i++)
80 if (0 == embd_descr[i].m_type)
82 ss << embd_descr[i].m_name;
84 else if (1 == embd_descr[i].m_type)
86 ss <<
"input_" << embd_descr[i].m_name;
88 if (i < embd_descr.size() - 1)
93 ss <<
" dim=2" << std::endl;
97 ss <<
"input = forward module=input_dropout input=input_brut" << std::endl;
101 std::string last_output_name =
"input";
102 for (
size_t i = 0; i < rnn_descr.size(); i++)
104 ss <<
"rnn_out_" << i <<
" = forward module=lstm_" << i
105 <<
" input=" << last_output_name << std::endl;
107 t <<
"rnn_out_" << i;
108 last_output_name = t.str();
113 for (
size_t i = 0; i < output_names.size(); ++i)
115 ss << output_names[i] <<
"_raw = forward module=fc_" << output_names[i]
116 <<
" input=" << last_output_name << std::endl;
118 ss << output_names[i] <<
" = log_softmax input=" << output_names[i] <<
"_raw" << std::endl;