LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_and_deep_biaffine_attention_script_generator.cpp
Go to the documentation of this file.
1// Copyright 2022 CEA LIST
2// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
6#include <sstream>
7#include <cassert>
8
10
11using namespace std;
12using namespace deeplima::nets;
13
14namespace deeplima
15{
16namespace graph_dp
17{
18namespace train
19{
20
21string BiRnnAndDeepBiaffineAttentionImpl::generate_script(const vector<embd_descr_t>& embd_descr,
22 const vector<rnn_descr_t>& rnn_descr,
23 const vector<deep_biaffine_attention_descr_t>& decoder_descr,
24 const vector<std::string>& /*output_names*/,
25 bool input_includes_root,
26 int64_t num_labels/*,
27 const vector<uint32_t>& classes*/)
28{
29 stringstream ss;
30
31 ss << "input_dropout = def Dropout prob=0.3" << std::endl;
32 //ss << "decoder_dropout = def Dropout prob=0.1" << std::endl;
33
34 size_t rnn_input_size = 0;
35 for (size_t i = 0; i < embd_descr.size(); i++)
36 {
37 if (1 == embd_descr[i].m_type)
38 {
39 ss << "embd_" << embd_descr[i].m_name << " = def Embedding dict=" << i
40 << " dim=" << embd_descr[i].m_dim << std::endl;
41 }
42 rnn_input_size += embd_descr[i].m_dim;
43 }
44
45 ss << std::endl;
46
47 size_t input_size = rnn_input_size;
48 for (size_t i = 0; i < rnn_descr.size(); i++)
49 {
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;
53 }
54
55 ss << std::endl;
56
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")
61 << endl;
62
63 if (num_labels > 0)
64 {
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")
69 << endl;
70 }
71
72 ss << std::endl;
73
74 /*for (size_t i = 0; i < output_names.size(); ++i)
75 {
76 ss << "fc_" << output_names[i] << " = def Linear input_size=" << input_size
77 << " output_size=" << classes[i] << std::endl;
78
79 ss << std::endl;
80 }*/
81
82 for (size_t i = 0; i < embd_descr.size(); i++)
83 {
84 ss << embd_descr[i].m_name << " = def Arg" << std::endl;
85 }
86
87 ss << std::endl;
88
89 for (size_t i = 0; i < embd_descr.size(); i++)
90 {
91 if (1 == embd_descr[i].m_type)
92 {
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;
96 }
97 }
98
99 ss << std::endl;
100
101 ss << "input_brut = cat input=";
102 for (size_t i = 0; i < embd_descr.size(); i++)
103 {
104 if (0 == embd_descr[i].m_type)
105 {
106 ss << embd_descr[i].m_name;
107 }
108 else if (1 == embd_descr[i].m_type)
109 {
110 ss << "input_" << embd_descr[i].m_name;
111 }
112 if (i < embd_descr.size() - 1)
113 {
114 ss << ",";
115 }
116 }
117 ss << " dim=2" << std::endl;
118
119 ss << std::endl;
120
121 ss << "input = forward module=input_dropout input=input_brut" << std::endl;
122
123 ss << std::endl;
124
125 std::string last_output_name = "input";
126 for (size_t i = 0; i < rnn_descr.size(); i++)
127 {
128 ss << "rnn_out_" << i << " = forward module=lstm_" << i
129 << " input=" << last_output_name << std::endl;
130 std::stringstream t;
131 t << "rnn_out_" << i;
132 last_output_name = t.str();
133 }
134
135 ss << std::endl;
136
137 //ss << "decoder_input = forward module=decoder_dropout input=" << last_output_name << endl;
138
139 ss << "arc_raw = forward module=decoder_0 input=" << last_output_name << endl;
140 ss << "arc = log_softmax input=arc_raw dim=2" << endl;
141
142 if (num_labels > 0)
143 {
144 // Per (dependent, head) label logits: [batch, dep, head, num_labels].
145 // The caller selects each dependent's label scores at its head (gold while
146 // training, predicted at inference) and applies log_softmax over the labels.
147 ss << "rel_logits = forward module=label_decoder_0 input=" << last_output_name << endl;
148 }
149
150 /*for (size_t i = 0; i < output_names.size(); ++i)
151 {
152 ss << output_names[i] << "_raw = forward module=fc_" << output_names[i]
153 << " input=" << last_output_name << std::endl;
154
155 ss << output_names[i] << " = log_softmax input=" << output_names[i] << "_raw" << std::endl;
156 ss << endl;
157 }*/
158
159 return ss.str();
160}
161
162} // namespace train
163} // namespace graph_dp
164} // namespace deeplima
165
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)
STL namespace.