LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
stanza_models_depparse_model.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
7#include <torch/nn/functional/activation.h>
8#include <torch/nn/modules/embedding.h>
9
10using namespace torch;
11using namespace torch::indexing;
12using namespace torch::nn;
13using namespace torch::nn::functional;
14using namespace torch::nn::utils::rnn;
15
16namespace deeplima
17{
18namespace nets
19{
20namespace torch_modules
21{
22
23// def add_unsaved_module(name, module):
24// unsaved_modules += [name]
25// setattr(self, name, module)
26
27StanzaDepparseParserImpl::StanzaDepparseParserImpl(int64_t word_emb_dim, int64_t tag_emb_dim, int64_t hidden_dim,
28 int64_t num_layers, float dropout, float rec_dropout,
29 std::shared_ptr<std::map<std::string, std::vector<std::string>>> vocab,
30 std::shared_ptr<std::vector<std::vector<std::string>>> feats_vocabs,
31 int64_t deep_biaff_hidden_dim,
32 bool linearize,
33 bool dist,
34 int64_t word_dropout)
35 : torch::nn::Module(),
36 m_vocab(vocab),
37 m_feats_vocabs(feats_vocabs),
38 m_unsaved_modules(),
39 word_emb(EmbeddingOptions(0,0)),
40 lemma_emb(EmbeddingOptions(0,0)),
41 upos_emb(EmbeddingOptions(0,0)),
42 xpos_emb(EmbeddingOptions(0,0)),
43 ufeats_emb(),
44 parserlstm(0, 0),
45 drop_replacement(), parserlstm_h_init(), parserlstm_c_init(),
46 unlabeled(2 * hidden_dim, 2 * hidden_dim, deep_biaff_hidden_dim, 1, dropout),
47 deprel(2 * hidden_dim, 2 * hidden_dim, deep_biaff_hidden_dim, (*m_vocab)["deprel"].size(), dropout),
48 linearization(),
49 distance(),
50 // criterion
51 crit(CrossEntropyLossOptions().ignore_index(-1).reduction(torch::kSum)), // ignore padding
52 drop(Dropout(dropout)),
53 worddrop(word_dropout),
54 word_emb_dim(word_emb_dim),
55 num_layers(num_layers),
56 hidden_dim(hidden_dim)
57{
58 // // self.args = args
59 // // self.share_hid = share_hid
60 // // self.unsaved_modules = []
61 //
62 // input layers
63 int64_t input_size = 0;
64 if (word_emb_dim > 0)
65 {
66 // frequent word embeddings
67 word_emb = nn::Embedding(EmbeddingOptions((*m_vocab)["word"].size(), word_emb_dim).padding_idx(0));
68 lemma_emb = nn::Embedding(EmbeddingOptions((*m_vocab)["lemma"].size(), word_emb_dim).padding_idx(0));
69 input_size += word_emb_dim * 2;
70 }
71 if (tag_emb_dim > 0)
72 {
73 upos_emb = nn::Embedding(EmbeddingOptions((*m_vocab)["upos"].size(), tag_emb_dim).padding_idx(0));
74
75 auto V = (*m_vocab)["xpos"];
76 // if not isinstance((*m_vocab)["xpos"], CompositeVocab):
77 xpos_emb = nn::Embedding(EmbeddingOptions((*m_vocab)["xpos"].size(), tag_emb_dim).padding_idx(0));
78 // else:
79 // xpos_emb = nn.ModuleList()
80 //
81 // for l in vocab["xpos"].lens():
82 // xpos_emb.append(nn::Embedding(EmbeddingOptions(l, tag_emb_dim).padding_idx(0));
83 //
84
85 for (const auto& vocab: *m_feats_vocabs)
86 ufeats_emb->push_back(nn::Embedding(EmbeddingOptions(vocab.size(), tag_emb_dim).padding_idx(0)));
87
88 input_size += tag_emb_dim * 2;
89 }
90 // if self.args["char"] and self.args["char_emb_dim"] > 0:
91 // charmodel = CharacterModel(args, vocab)
92 // trans_char = nn::Linear(self.args["char_hidden_dim"], self.args["transformed_dim"], bias=False)
93 // input_size += self.args["transformed_dim"]
94 //
95 // if self.args["pretrain"]:
96 // # pretrained embeddings, by default this won't be saved into model file
97 // add_unsaved_module("pretrained_emb", nn::Embedding.from_pretrained(torch.from_numpy(emb_matrix), freeze=True))
98 // trans_pretrained = nn.Linear(emb_matrix.shape[1], self.args["transformed_dim"], bias=False)
99 // input_size += self.args["transformed_dim"]
100
101 // recurrent layers
102 parserlstm = HighwayLSTM(input_size, hidden_dim, num_layers, true, true, dropout, true,
103 rec_dropout, torch::tanh);
104 parserlstm->register_parameter("drop_replacement", torch::randn(input_size) / sqrt(input_size), false);
105 parserlstm->register_parameter("parserlstm_h_init", torch::zeros({2 * num_layers, 1, hidden_dim}), false);
106 parserlstm->register_parameter("parserlstm_c_init", torch::zeros({2 * num_layers, 1, hidden_dim}), false);
107
108 // classifiers
109 if (linearize)
110 linearization = std::make_unique<DeepBiaffineScorer>(2 * hidden_dim, 2 * hidden_dim, deep_biaff_hidden_dim, 1, dropout);
111 if (dist)
112 distance = std::make_unique<DeepBiaffineScorer>(2 * hidden_dim, 2 * hidden_dim, deep_biaff_hidden_dim, 1, dropout);
113}
114
115PackedSequence StanzaDepparseParserImpl::pack(torch::Tensor x, torch::Tensor sentlens)
116{
117 return pack_padded_sequence(x, sentlens, /*batch_first=T*/true);
118}
119
120std::pair<float, std::vector<torch::Tensor>> StanzaDepparseParserImpl::forward(
121 torch::Tensor word,
122 torch::Tensor word_mask,
123 torch::Tensor wordchars,
124 torch::Tensor wordchars_mask,
125 torch::Tensor upos,
126 torch::Tensor xpos,
127 torch::Tensor ufeats,
128 torch::Tensor pretrained,
129 torch::Tensor lemma,
130 torch::Tensor head,
131 torch::Tensor deprels,
132 torch::Tensor word_orig_idx,
133 torch::Tensor sentlens,
134 torch::Tensor wordlens)
135{
136 std::vector<std::tuple<PackedSequence, PackedSequence>> inputs;
137 // if self.args["pretrain"]:
138 // pretrained_emb = pretrained_emb(pretrained)
139 // pretrained_emb = trans_pretrained(pretrained_emb)
140 // pretrained_emb = pack(pretrained_emb)
141 // inputs += [pretrained_emb]
142 //
143 // #def pad(x):
144 // # return pad_packed_sequence(PackedSequence(x, pretrained_emb.batch_sizes), batch_first=True)[0]
145 //
146 if (word_emb_dim > 0)
147 {
148 auto word_embed = word_emb(word);
149 auto packed_word_embed = pack(word_embed, sentlens);
150 auto lemma_embed = lemma_emb(lemma);
151 auto packed_lemma_embed = pack(lemma_embed, sentlens);
152 inputs.push_back({packed_word_embed, packed_lemma_embed});
153 }
154 //
155 // if tag_emb_dim > 0:
156 // pos_emb = upos_emb(upos)
157 //
158 // if isinstance((*m_vocab)["xpos"], CompositeVocab):
159 // for i in range(len((*m_vocab)["xpos"])):
160 // pos_emb += xpos_emb[i](xpos[:, :, i])
161 // else:
162 // pos_emb += xpos_emb(xpos)
163 // pos_emb = pack(pos_emb)
164 //
165 // feats_emb = 0
166 // for i in range(len((*m_vocab)["feats"])):
167 // feats_emb += ufeats_emb[i](ufeats[:, :, i])
168 // feats_emb = pack(feats_emb)
169 //
170 // inputs += [pos_emb, feats_emb]
171 //
172 // if self.args["char"] and self.args["char_emb_dim"] > 0:
173 // char_reps = charmodel(wordchars, wordchars_mask, word_orig_idx, sentlens, wordlens)
174 // char_reps = PackedSequence(trans_char(drop(char_reps.data)), char_reps.batch_sizes)
175 // inputs += [char_reps]
176 //
177 std::vector<torch::Tensor> lstm_inputs_vec;
178 for (auto x: inputs)
179 {
180 // lstm_inputs_vec.push_back(x.data());
181 }
182 auto lstm_inputs = torch::cat(lstm_inputs_vec, 1);
183 //
184 lstm_inputs = worddrop(lstm_inputs, drop_replacement);
185 // lstm_inputs = drop(lstm_inputs);
186 //
187 // lstm_inputs = PackedSequence(lstm_inputs, inputs.index({0}).batch_sizes());
188 //
189// torch::Tensor input, torch::Tensor seqlens, torch::Tensor hx=torch::Tensor()
190 auto [lstm_outputs, none] = parserlstm(lstm_inputs, sentlens,
191 std::make_tuple(
192 parserlstm_h_init.expand({2 * num_layers, word.size(0), hidden_dim}).contiguous(),
193 parserlstm_c_init.expand({2 * num_layers, word.size(0), hidden_dim}).contiguous()));
194 // lstm_outputs, _ = pad_packed_sequence(lstm_outputs, batch_first=True)
195 //
196 auto unlabeled_scores = unlabeled(drop(lstm_outputs), drop(lstm_outputs)).squeeze(3);
197 auto deprel_scores = deprel(drop(lstm_outputs), drop(lstm_outputs));
198 //
199 // #goldmask = head.new_zeros(*head.size(), head.size(-1)+1, dtype=torch.uint8)
200 // #goldmask.scatter_(2, head.unsqueeze(2), 1)
201 //
202 // if self.args["linearization"] or self.args["distance"]:
203 // head_offset = torch.arange(word.size(1), device=head.device).view(1, 1, -1).expand(word.size(0), -1, -1) - torch.arange(word.size(1), device=head.device).view(1, -1, 1).expand(word.size(0), -1, -1)
204 //
205 // if self.args["linearization"]:
206 // lin_scores = linearization(drop(lstm_outputs), drop(lstm_outputs)).squeeze(3)
207 // unlabeled_scores += F.logsigmoid(lin_scores * torch.sign(head_offset).float()).detach()
208 //
209 // if self.args["distance"]:
210 // dist_scores = distance(drop(lstm_outputs), drop(lstm_outputs)).squeeze(3)
211 // dist_pred = 1 + F.softplus(dist_scores)
212 // dist_target = torch.abs(head_offset)
213 // dist_kld = -torch.log((dist_target.float() - dist_pred)**2/2 + 1)
214 // unlabeled_scores += dist_kld.detach()
215 //
216 auto diag = torch::eye(head.size(-1)+1, torch::dtype(torch::kUInt8).device(head.device())).unsqueeze(0);
217 unlabeled_scores.masked_fill_(diag, -std::numeric_limits<float>::infinity());
218
219 float loss;
220 std::vector<torch::Tensor> preds;
221
222 if (is_training())
223 {
224 unlabeled_scores = unlabeled_scores.index({None, Slice(1), None}); // exclude attachment for the root symbol
225 unlabeled_scores = unlabeled_scores.masked_fill(word_mask.unsqueeze(1), -std::numeric_limits<float>::infinity());
226 auto unlabeled_target = head.masked_fill(word_mask.index({None, Slice(1)}), -1);
227 loss = crit(unlabeled_scores.contiguous().view({-1, unlabeled_scores.size(2)}), unlabeled_target.view(-1)).item<float>();
228
229 deprel_scores = deprel_scores.index({None, Slice(1)}); // exclude attachment for the root symbol
230 // #deprel_scores = deprel_scores.masked_select(goldmask.unsqueeze(3)).view(-1, len((*m_vocab)["deprel"]))
231 deprel_scores = torch::gather(deprel_scores, 2,
232 head.unsqueeze(2).unsqueeze(3).expand({-1, -1, -1,
233 (*m_vocab)["deprel"].size()})).view({-1, (*m_vocab)["deprel"].size()});
234 auto deprel_target = deprels.masked_fill(word_mask.index({None, Slice(1)}), -1);
235 loss += crit(deprel_scores.contiguous(), deprel_target.view(-1)).item<float>();
236
237 // if (self.args["linearization"])
238 // {
239 // // #lin_scores = lin_scores[:, 1:].masked_select(goldmask)
240 // lin_scores = torch::gather(lin_scores[:, 1:], 2, head.unsqueeze(2)).view(-1)
241 // lin_scores = torch::cat([-lin_scores.unsqueeze(1)/2, lin_scores.unsqueeze(1)/2], 1)
242 // // #lin_target = (head_offset[:, 1:] > 0).long().masked_select(goldmask)
243 // lin_target = torch::gather((head_offset[:, 1:] > 0).long(), 2, head.unsqueeze(2))
244 // loss += crit(lin_scores.contiguous(), lin_target.view(-1))
245 // }
246 //
247 // if (self.args["distance"])
248 // {
249 // // #dist_kld = dist_kld[:, 1:].masked_select(goldmask)
250 // dist_kld = torch::gather(dist_kld[:, 1:], 2, head.unsqueeze(2))
251 // loss -= dist_kld.sum()
252 // }
253
254 loss /= wordchars.size(0); // # number of words
255 }
256 else
257 {
258 loss = 0;
259 auto X = log_softmax(unlabeled_scores, 2).detach().cpu();
260 preds.push_back(log_softmax(unlabeled_scores, 2).detach().cpu());
261 preds.push_back(std::get<1>(torch::max(deprel_scores, 3)).detach().cpu());
262 }
263 return {loss, preds};
264}
265
266} // torch_modules
267} // nets
268} // deeplima
std::pair< float, std::vector< torch::Tensor > > forward(torch::Tensor word, torch::Tensor word_mask, torch::Tensor wordchars, torch::Tensor wordchars_mask, torch::Tensor upos, torch::Tensor xpos, torch::Tensor ufeats, torch::Tensor pretrained, torch::Tensor lemma, torch::Tensor head, torch::Tensor deprels, torch::Tensor word_orig_idx, torch::Tensor sentlens, torch::Tensor wordlens)