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,
35 : torch::nn::Module(),
37 m_feats_vocabs(feats_vocabs),
39 word_emb(EmbeddingOptions(0,0)),
40 lemma_emb(EmbeddingOptions(0,0)),
41 upos_emb(EmbeddingOptions(0,0)),
42 xpos_emb(EmbeddingOptions(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),
51 crit(CrossEntropyLossOptions().ignore_index(-1).reduction(torch::kSum)),
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)
63 int64_t input_size = 0;
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;
73 upos_emb = nn::Embedding(EmbeddingOptions((*m_vocab)[
"upos"].size(), tag_emb_dim).padding_idx(0));
75 auto V = (*m_vocab)[
"xpos"];
77 xpos_emb = nn::Embedding(EmbeddingOptions((*m_vocab)[
"xpos"].size(), tag_emb_dim).padding_idx(0));
85 for (
const auto& vocab: *m_feats_vocabs)
86 ufeats_emb->push_back(nn::Embedding(EmbeddingOptions(vocab.size(), tag_emb_dim).padding_idx(0)));
88 input_size += tag_emb_dim * 2;
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);
110 linearization = std::make_unique<DeepBiaffineScorer>(2 * hidden_dim, 2 * hidden_dim, deep_biaff_hidden_dim, 1, dropout);
112 distance = std::make_unique<DeepBiaffineScorer>(2 * hidden_dim, 2 * hidden_dim, deep_biaff_hidden_dim, 1, dropout);
122 torch::Tensor word_mask,
123 torch::Tensor wordchars,
124 torch::Tensor wordchars_mask,
127 torch::Tensor ufeats,
128 torch::Tensor pretrained,
131 torch::Tensor deprels,
132 torch::Tensor word_orig_idx,
133 torch::Tensor sentlens,
134 torch::Tensor wordlens)
136 std::vector<std::tuple<PackedSequence, PackedSequence>> inputs;
146 if (word_emb_dim > 0)
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});
177 std::vector<torch::Tensor> lstm_inputs_vec;
182 auto lstm_inputs = torch::cat(lstm_inputs_vec, 1);
184 lstm_inputs = worddrop(lstm_inputs, drop_replacement);
190 auto [lstm_outputs,
none] = parserlstm(lstm_inputs, sentlens,
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()));
196 auto unlabeled_scores = unlabeled(drop(lstm_outputs), drop(lstm_outputs)).squeeze(3);
197 auto deprel_scores = deprel(drop(lstm_outputs), drop(lstm_outputs));
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());
220 std::vector<torch::Tensor> preds;
224 unlabeled_scores = unlabeled_scores.index({None, Slice(1), None});
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>();
229 deprel_scores = deprel_scores.index({None, Slice(1)});
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>();
254 loss /= wordchars.size(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());
263 return {loss, preds};
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)