22 const vector<rnn_descr_t>& rnn_descr,
23 const vector<deep_biaffine_attention_descr_t>& decoder_descr,
24 const vector<std::string>& ,
25 bool input_includes_root,
31 ss <<
"input_dropout = def Dropout prob=0.3" << std::endl;
34 size_t rnn_input_size = 0;
35 for (
size_t i = 0; i < embd_descr.size(); i++)
37 if (1 == embd_descr[i].m_type)
39 ss <<
"embd_" << embd_descr[i].m_name <<
" = def Embedding dict=" << i
40 <<
" dim=" << embd_descr[i].m_dim << std::endl;
42 rnn_input_size += embd_descr[i].m_dim;
47 size_t input_size = rnn_input_size;
48 for (
size_t i = 0; i < rnn_descr.size(); i++)
50 ss <<
"lstm_" << i <<
" = def LSTM input_size=" << input_size
51 <<
" hidden_size=" << rnn_descr[i].m_dim <<
" num_layers=1 bidirectional=true" << std::endl;
52 input_size = rnn_descr[i].m_dim * 2;
57 assert(decoder_descr.size() == 1);
58 ss <<
"decoder_0 = def DeepBiaffineAttentionDecoder input_dim=" << input_size
59 <<
" hidden_arc_dim=" << decoder_descr[0].m_arc_dim
60 <<
" input_includes_root=" << (input_includes_root ?
"true" :
"false")
65 ss <<
"label_decoder_0 = def DeepBiaffineAttentionLabelDecoder input_dim=" << input_size
66 <<
" hidden_dim=" << decoder_descr[0].m_arc_dim
67 <<
" num_labels=" << num_labels
68 <<
" input_includes_root=" << (input_includes_root ?
"true" :
"false")
82 for (
size_t i = 0; i < embd_descr.size(); i++)
84 ss << embd_descr[i].m_name <<
" = def Arg" << std::endl;
89 for (
size_t i = 0; i < embd_descr.size(); i++)
91 if (1 == embd_descr[i].m_type)
93 ss <<
"input_" << embd_descr[i].m_name <<
" = forward module=embd_"
94 << embd_descr[i].m_name
95 <<
" input=" << embd_descr[i].m_name << std::endl;
101 ss <<
"input_brut = cat input=";
102 for (
size_t i = 0; i < embd_descr.size(); i++)
104 if (0 == embd_descr[i].m_type)
106 ss << embd_descr[i].m_name;
108 else if (1 == embd_descr[i].m_type)
110 ss <<
"input_" << embd_descr[i].m_name;
112 if (i < embd_descr.size() - 1)
117 ss <<
" dim=2" << std::endl;
121 ss <<
"input = forward module=input_dropout input=input_brut" << std::endl;
125 std::string last_output_name =
"input";
126 for (
size_t i = 0; i < rnn_descr.size(); i++)
128 ss <<
"rnn_out_" << i <<
" = forward module=lstm_" << i
129 <<
" input=" << last_output_name << std::endl;
131 t <<
"rnn_out_" << i;
132 last_output_name = t.str();
139 ss <<
"arc_raw = forward module=decoder_0 input=" << last_output_name << endl;
140 ss <<
"arc = log_softmax input=arc_raw dim=2" << endl;
147 ss <<
"rel_logits = forward module=label_decoder_0 input=" << last_output_name << endl;
static std::string generate_script(const std::vector< nets::embd_descr_t > &embd_descr, const std::vector< nets::rnn_descr_t > &rnn_descr, const std::vector< nets::deep_biaffine_attention_descr_t > &decoder_descr, const std::vector< std::string > &output_names, bool input_includes_root=false, int64_t num_labels=0)