6#ifndef DEEPLIMA_LIBS_NN_STANZA_MODELS_DEPPARSE_MODEL_H
7#define DEEPLIMA_LIBS_NN_STANZA_MODELS_DEPPARSE_MODEL_H
9#include <torch/torch.h>
10#include <torch/nn/modules/dropout.h>
11#include <torch/nn/modules/embedding.h>
12#include <torch/nn/modules/linear.h>
13#include <torch/nn/modules/rnn.h>
14#include <torch/nn/utils/rnn.h>
16#include <ATen/Functions.h>
24namespace torch_modules
40 torch::Tensor
forward(torch::Tensor x, torch::Tensor replacement=torch::Tensor())
45 auto masksize = std::vector<int64_t>(x.numel());
46 for (int64_t y = 0; y < x.numel(); y++) masksize[y] = y;
49 auto dropmask = torch::rand(masksize) <
dropprob;
51 auto res = x.masked_fill(dropmask, 0);
52 if (replacement.numel() == 0)
53 res = res + dropmask.to(torch::kFloat) * replacement;
60 return std::string(
"p=") + std::to_string(
dropprob);
74 bool pad=
false,
float rec_dropout=0):
81 drop(torch::nn::DropoutOptions().p(
dropout).inplace(true)),
82 rec_drop(torch::nn::DropoutOptions().p(rec_dropout).inplace(true)),
90 cells->push_back(torch::nn::LSTMCell(torch::nn::LSTMCellOptions(in_size,
hidden_size).bias(bias)));
96 std::pair<torch::Tensor, std::pair<torch::Tensor, torch::Tensor>>
rnn_loop(torch::Tensor x,
97 torch::Tensor batch_sizes,
98 torch::nn::LSTMCellImpl* cell,
99 std::vector<torch::Tensor> inits,
102 auto batch_size = batch_sizes[0].item();
103 std::vector<std::vector<torch::Tensor>> states;
104 for (
auto init: inits)
106 auto l = init.split(std::vector<int64_t>(batch_size.to<int64_t>(), 1));
109 auto h_drop_mask = x.new_ones({batch_size.to<int64_t>(),
hidden_size});
110 h_drop_mask =
rec_drop(h_drop_mask);
111 std::vector<torch::Tensor> resh;
116 for (
int i = 0; i < batch_sizes.numel(); i++)
118 auto bs = *(batch_sizes.data_ptr<int64_t>()+i);
121 std::vector<torch::Tensor> slice0, slice1;
122 for (int64_t p = 0; p < bs; p++)
124 slice0.push_back(states[0][p]);
125 slice1.push_back(states[1][p]);
127 auto s1 = cell->forward(x.index({torch::indexing::Slice(st, st+bs)}),
129 torch::cat(slice0, 0) * h_drop_mask.index({torch::indexing::Slice(0, bs)}),
130 torch::cat(slice1, 0)
132 resh.push_back(std::get<0>(s1));
133 for (int64_t j = 0; j < bs; j++)
135 states[0][j] = std::get<0>(s1).index({j}).unsqueeze(0);
136 states[1][j] = std::get<1>(s1).index({j}).unsqueeze(0);
144 for (int64_t i = batch_sizes.size(0)-1; i > 0; i--)
146 auto bs = *(batch_sizes.data_ptr<int64_t>()+i);
148 std::vector<torch::Tensor> slice0, slice1;
149 for (int64_t p = 0; p < bs; p++)
151 slice0.push_back(states[0][p]);
152 slice1.push_back(states[1][p]);
154 auto s1 = cell->forward(x.index({torch::indexing::Slice(en-bs, en)}),
155 std::make_tuple<torch::Tensor, torch::Tensor>(
156 torch::cat(slice0, 0) * h_drop_mask.index({torch::indexing::Slice(0, bs)}),
157 torch::cat(slice1, 0)));
158 resh.push_back(std::get<0>(s1));
159 for (int64_t j = 0; j < bs; j++)
161 states[0][j] = std::get<0>(s1).index({j}).unsqueeze(0);
162 states[1][j] = std::get<1>(s1).index({j}).unsqueeze(0);
166 std::reverse(resh.begin(), resh.end());
169 return std::make_pair(torch::cat(resh, 0), std::make_pair(torch::cat(states[0], 0), torch::cat(states[0], 0)));
172 std::tuple<torch::nn::utils::rnn::PackedSequence, std::tuple<torch::Tensor, torch::Tensor>>
forward(
173 torch::nn::utils::rnn::PackedSequence input,
174 torch::optional<std::tuple<torch::Tensor, torch::Tensor>> hx = {})
176 std::pair<std::vector<torch::Tensor>, std::vector<torch::Tensor>> all_states;
177 auto inputdata = input.data();
178 auto batch_sizes = input.batch_sizes();
179 for (int64_t l = 0; l < num_layers; l++)
181 std::vector<torch::Tensor> new_input;
183 if (dropout > 0 && l > 0)
184 inputdata = drop(inputdata);
185 for (int64_t d = 0; d < num_directions; d++)
187 auto idx = l * num_directions + d;
188 auto cell = cells[idx]->as<torch::nn::LSTMCell>();
190 std::vector<torch::Tensor> inits;
192 inits = {std::get<0>(*hx).index({idx}), std::get<1>(*hx).index({idx})};
195 input.data().new_zeros({input.batch_sizes().index({0}).item().to<int64_t>(), hidden_size},
196 torch::TensorOptions().requires_grad(
false)),
197 input.data().new_zeros({input.batch_sizes().index({0}).item().to<int64_t>(), hidden_size},
198 torch::TensorOptions().requires_grad(
false))
200 auto loop_result = rnn_loop(inputdata, batch_sizes, cell, inits, (d == 1));
201 auto out = std::get<0>(loop_result);
202 auto states = std::get<1>(loop_result);
203 new_input.push_back(out);
204 std::get<0>(all_states).push_back(std::get<0>(states).unsqueeze(0));
205 std::get<1>(all_states).push_back(std::get<1>(states).unsqueeze(0));
207 if (num_directions > 1)
209 inputdata = torch::cat(new_input, 1);
211 inputdata = new_input[0];
213 input = torch::nn::utils::rnn::PackedSequence(inputdata, batch_sizes);
215 return {input, {torch::cat(std::get<0>(all_states), 0), torch::cat(std::get<1>(all_states))} };
238 bool bias=
true,
bool batch_first=
false,
239 float dropout=0,
bool bidirectional=
false,
240 bool pad=
false,
float rec_dropout=0):
242 batch_first(batch_first),
245 if (rec_dropout == 0)
248 lstm = std::make_shared<torch::nn::LSTMImpl>(
249 torch::nn::LSTMOptions(input_size, hidden_size).num_layers(num_layers).batch_first(batch_first)
250 .bidirectional(bidirectional).bias(bias).dropout(dropout));
254 lstm = std::make_shared<LSTMwRecDropoutImpl>(input_size, hidden_size, num_layers, bias, batch_first,
255 dropout, bidirectional, rec_dropout);
259 std::tuple<torch::Tensor, std::tuple<torch::Tensor, torch::Tensor>>
forward(torch::Tensor input, torch::Tensor lengths,
260 torch::optional<std::tuple<torch::Tensor, torch::Tensor>> hx = {})
264 auto res = lstm.forward<std::tuple<torch::Tensor, std::tuple<torch::Tensor, torch::Tensor>>>(input, hx);
270 std::tuple<torch::nn::utils::rnn::PackedSequence, std::tuple<torch::Tensor, torch::Tensor>>
forward_with_packed_input(
const torch::nn::utils::rnn::PackedSequence &packed_input, torch::Tensor lengths,
271 torch::optional<std::tuple<torch::Tensor, torch::Tensor>> hx = {})
275 auto effective = lstm.ptr<torch::nn::LSTMImpl>();
278 auto res = effective->forward_with_packed_input(packed_input, hx);
286 auto effective_rec = lstm.ptr<LSTMwRecDropoutImpl>();
288 auto res = effective_rec->forward(packed_input, hx);
310 int64_t num_layers=1,
bool bias=
true,
bool batch_first=
false,
311 float dropout=0,
bool bidirectional=
false,
float rec_dropout=0,
312 std::function<torch::Tensor(
const torch::Tensor&)> highway_func=
nullptr,
315 input_size(input_size),
316 hidden_size(hidden_size),
317 num_layers(num_layers),
319 batch_first(batch_first),
322 bidirectional(bidirectional),
323 num_directions(bidirectional ? 2 : 1),
324 highway_func(highway_func),
329 drop(torch::nn::DropoutOptions().p(dropout).inplace(true))
331 auto in_size = input_size;
332 for (int64_t l = 0; l < num_layers; l++)
334 lstm->push_back(PackedLSTM(in_size, hidden_size, 1, bias,
335 batch_first, 0, bidirectional, rec_dropout));
336 highway->push_back(torch::nn::Linear(in_size, hidden_size * num_directions));
337 gate->push_back(torch::nn::Linear(in_size, hidden_size * num_directions));
340 in_size = hidden_size * num_directions;
344 std::tuple<torch::Tensor, std::tuple<torch::Tensor, torch::Tensor>>
forward(
345 torch::Tensor input, torch::Tensor seqlens,
346 std::tuple<torch::Tensor, torch::Tensor> hx={torch::Tensor(),torch::Tensor()})
348 highway_func = highway_func.target<torch::Tensor(
const torch::Tensor&)>() ==
nullptr ? [](
const torch::Tensor& t) {
return t; } : highway_func;
350 std::vector<torch::Tensor> hs;
351 std::vector<torch::Tensor> cs;
352 auto packed_sequence = torch::nn::utils::rnn::pack_padded_sequence(input, seqlens, batch_first);
354 for (int64_t l = 0 ; l < num_layers; l++)
357 packed_sequence = torch::nn::utils::rnn::PackedSequence(drop(packed_sequence.data()),
358 packed_sequence.batch_sizes(),
359 packed_sequence.sorted_indices(),
360 packed_sequence.unsorted_indices());
361 auto layer_hx = std::get<0>(hx).numel() > 0 ?
362 std::make_tuple(std::get<0>(hx).index({torch::indexing::Slice(l * num_directions,(l+1)*num_directions)}),
363 std::get<1>(hx).index({torch::indexing::Slice(l * num_directions, (l+1)*num_directions)})) :
364 std::make_tuple(torch::Tensor(), torch::Tensor());
365 auto X = (lstm[l]->as<torch::nn::LSTMImpl>()) ?
366 lstm[l]->as<torch::nn::LSTMImpl>()->forward_with_packed_input(packed_sequence, layer_hx) :
367 lstm[l]->as<LSTMwRecDropoutImpl>()->forward(packed_sequence, layer_hx);
368 auto h = std::get<0>(X);
369 auto t = std::get<1>(X);
370 auto ht = std::get<0>(t);
371 auto ct = std::get<1>(t);
375 packed_sequence = torch::nn::utils::rnn::PackedSequence(
376 (h.data() + torch::sigmoid(gate[l]->as<torch::nn::Linear>()->forward(packed_sequence.data()))
377 * highway_func(highway[l]->as<torch::nn::Linear>()->forward(packed_sequence.data()))),
378 packed_sequence.batch_sizes(),
379 packed_sequence.sorted_indices(),
380 packed_sequence.unsorted_indices());
385 return {input, {torch::cat(hs, 0), torch::cat(cs, 0)}};
421 this->output_size = output_size;
423 this->register_parameter(
"weight", torch::zeros({input1_size, input2_size, output_size}),
false);
424 register_parameter(
"bias", bias ? torch::zeros({output_size}) : torch::zeros({}),
false);
427 torch::Tensor
forward(torch::Tensor input1, torch::Tensor input2)
429 auto input1_size = input1.sizes();
430 auto input2_size = input2.sizes();
435 auto intermediate = at::mm(input1.view({-1, input1_size[-1]}),
436 named_parameters()[
"weight"].view({-1, this->input2_size * this->output_size}));
438 input2 = input2.transpose(1, 2);
440 auto output = intermediate.view({input1_size[0], input1_size[1] * this->output_size, input2_size[2]}).bmm(input2);
442 output = output.view({input1_size[0], input1_size[1], this->output_size, input2_size[1]}).transpose(2, 3);
457 W_bilin(input1_size + 1, input2_size + 1, output_size)
460 W_bilin->weight.data().zero_();
461 W_bilin->bias.data().zero_();
464 torch::Tensor
forward(torch::Tensor input1, torch::Tensor input2)
469 return W_bilin(input1, input2);
481 W_bilin(input1_size + 1, input2_size + 1, output_size)
488 torch::Tensor
forward(torch::Tensor input1, torch::Tensor input2)
493 return W_bilin(input1, input2);
507 int64_t hidden_size, int64_t output_size,
508 float dropout_value=0,
bool pairwise=
true,
509 std::function<torch::Tensor(
const torch::Tensor&)> hidden_func=at::relu):
512 dropout(dropout_value),
513 W1(input1_size, hidden_size),
514 W2(input2_size, hidden_size),
515 m_hidden_func(hidden_func)
518 scorer = std::make_shared<PairwiseBiaffineScorerImpl>(hidden_size, hidden_size, output_size);
520 scorer = std::make_shared<BiaffineScorerImpl>(hidden_size, hidden_size, output_size);
523 torch::Tensor
forward(torch::Tensor input1, torch::Tensor input2)
525 return scorer.forward(dropout(m_hidden_func(W1(input1))), dropout(m_hidden_func(W2(input2))));
530 torch::nn::Linear
W1, W2;
540 int64_t num_layers,
float dropout,
float rec_dropout,
541 std::shared_ptr<std::map<std::string, std::vector<std::string>>> vocab,
542 std::shared_ptr<std::vector<std::vector<std::string>>> feats_vocabs,
543 int64_t deep_biaff_hidden_dim,
546 int64_t word_dropout);
548 std::pair<float, std::vector<torch::Tensor>> forward(
549 torch::Tensor word, torch::Tensor word_mask, torch::Tensor wordchars, torch::Tensor wordchars_mask,
550 torch::Tensor upos, torch::Tensor xpos, torch::Tensor ufeats, torch::Tensor pretrained,
551 torch::Tensor lemma, torch::Tensor head, torch::Tensor deprels, torch::Tensor word_orig_idx,
552 torch::Tensor sentlens, torch::Tensor wordlens);
555 torch::nn::utils::rnn::PackedSequence pack(torch::Tensor x, torch::Tensor sentlens);
557 std::shared_ptr<std::map<std::string, std::vector<std::string>>> m_vocab;
558 std::shared_ptr<std::vector<std::vector<std::string>>> m_feats_vocabs;
559 std::vector<std::string> m_unsaved_modules;
560 torch::nn::Embedding word_emb, lemma_emb, upos_emb, xpos_emb;
561 torch::nn::ModuleList ufeats_emb;
562 HighwayLSTM parserlstm;
563 torch::Tensor drop_replacement, parserlstm_h_init, parserlstm_c_init;
564 DeepBiaffineScorer unlabeled;
565 DeepBiaffineScorer deprel;
566 std::unique_ptr<DeepBiaffineScorer> linearization;
567 std::unique_ptr<DeepBiaffineScorer> distance;
568 torch::nn::CrossEntropyLoss crit;
569 torch::nn::Dropout drop;
570 WordDropout worddrop;
571 int64_t word_emb_dim;
torch::Tensor forward(torch::Tensor input1, torch::Tensor input2)
BiaffineScorerImpl(int64_t input1_size, int64_t input2_size, int64_t output_size)
BiaffineScorerImpl()=default
torch::nn::Bilinear W_bilin
torch::Tensor forward(torch::Tensor input1, torch::Tensor input2)
torch::nn::AnyModule scorer
DeepBiaffineScorerImpl(int64_t input1_size, int64_t input2_size, int64_t hidden_size, int64_t output_size, float dropout_value=0, bool pairwise=true, std::function< torch::Tensor(const torch::Tensor &)> hidden_func=at::relu)
torch::nn::Dropout dropout
std::function< torch::Tensor(const torch::Tensor &)> m_hidden_func
Highway LSTM network, does NOT use the HLSTMCell above A Highway LSTM network, as used in the origina...
torch::nn::ModuleList highway
std::function< torch::Tensor(const torch::Tensor &)> highway_func
HighwayLSTMImpl(int64_t input_size, int64_t hidden_size, int64_t num_layers=1, bool bias=true, bool batch_first=false, float dropout=0, bool bidirectional=false, float rec_dropout=0, std::function< torch::Tensor(const torch::Tensor &)> highway_func=nullptr, bool pad=false)
torch::nn::ModuleList lstm
HighwayLSTMImpl()=default
torch::nn::ModuleList gate
std::tuple< torch::Tensor, std::tuple< torch::Tensor, torch::Tensor > > forward(torch::Tensor input, torch::Tensor seqlens, std::tuple< torch::Tensor, torch::Tensor > hx={torch::Tensor(), torch::Tensor()})
An LSTM implementation that supports recurrent dropout.
torch::nn::Dropout rec_drop
LSTMwRecDropoutImpl(int64_t input_size, int64_t hidden_size, int64_t num_layers, bool bias=true, bool batch_first=false, float dropout=0, bool bidirectional=false, bool pad=false, float rec_dropout=0)
LSTMwRecDropoutImpl()=default
torch::nn::ModuleList cells
std::pair< torch::Tensor, std::pair< torch::Tensor, torch::Tensor > > rnn_loop(torch::Tensor x, torch::Tensor batch_sizes, torch::nn::LSTMCellImpl *cell, std::vector< torch::Tensor > inits, bool reverse=false)
RNN loop for one layer in one direction with recurrent dropout Assumes input is PackedSequence,...
std::tuple< torch::nn::utils::rnn::PackedSequence, std::tuple< torch::Tensor, torch::Tensor > > forward(torch::nn::utils::rnn::PackedSequence input, torch::optional< std::tuple< torch::Tensor, torch::Tensor > > hx={})
torch::nn::AnyModule lstm
std::tuple< torch::nn::utils::rnn::PackedSequence, std::tuple< torch::Tensor, torch::Tensor > > forward_with_packed_input(const torch::nn::utils::rnn::PackedSequence &packed_input, torch::Tensor lengths, torch::optional< std::tuple< torch::Tensor, torch::Tensor > > hx={})
std::tuple< torch::Tensor, std::tuple< torch::Tensor, torch::Tensor > > forward(torch::Tensor input, torch::Tensor lengths, torch::optional< std::tuple< torch::Tensor, torch::Tensor > > hx={})
PackedLSTMImpl(int64_t input_size, int64_t hidden_size, int64_t num_layers, bool bias=true, bool batch_first=false, float dropout=0, bool bidirectional=false, bool pad=false, float rec_dropout=0)
PairwiseBiaffineScorerImpl()=default
PairwiseBiaffineScorerImpl(int64_t input1_size, int64_t input2_size, int64_t output_size)
torch::Tensor forward(torch::Tensor input1, torch::Tensor input2)
A bilinear module that deals with broadcasting for efficient memory usage.
PairwiseBilinearImpl(int64_t input1_size, int64_t input2_size, int output_size, bool bias=true)
torch::Tensor forward(torch::Tensor input1, torch::Tensor input2)
PairwiseBilinearImpl()=default
StanzaDepparseParserImpl()=default
A word dropout layer that's designed for embedded inputs (e.g., any inputs to an LSTM layer).
WordDropoutImpl()=default
torch::Tensor forward(torch::Tensor x, torch::Tensor replacement=torch::Tensor())
WordDropoutImpl(int64_t dropprob)
TORCH_MODULE(DeepBiaffineAttentionDecoder)
TORCH_MODULE(BiRnnSeq2Seq)
PUGI__FN void reverse(I begin, I end)