LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
deep_biaffine_attention_decoder.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
8namespace deeplima
9{
10namespace nets
11{
12namespace torch_modules
13{
14
15torch::Tensor DeepBiaffineAttentionDecoderImpl::forward(torch::Tensor input)
16{
17 int64_t batch_size = input.size(1);
18
19 //std::cerr << "input.sizes() == " << input.sizes() << std::endl;
20 torch::Tensor input_t = input.transpose(0, 1); // to [ batch x len x input_dim ]
21 //std::cerr << "input_t.sizes() == " << input_t.sizes() << std::endl;
22
23 torch::Tensor arc_head;
25 {
26 torch::Tensor roots2 = torch::tile(root2, { batch_size, 1, 1 });
27 //arc_head = torch::cat({ roots2, elu(mlp_head(input_t)) }, 1);
28 arc_head = elu(torch::cat({ roots2, mlp_head(input_t) }, 1));
29 }
30 else
31 {
32 // h_j^(arc-head) = MLP^(arc-head)(r_j) in (5) from [Dozat&Manning,2017]
33 // and thus
34 // H^(arc-head) in (6) from [Dozat&Manning,2017]
35 arc_head = elu(mlp_head(input_t));
36 }
37 // arc_head = [ batch x len x hidden_dim ]
38 //std::cerr << "arc_head.sizes() == " << arc_head.sizes() << std::endl;
39
40 torch::Tensor arc_dep;
42 {
43 torch::Tensor roots = torch::tile(root, { batch_size, 1, 1 });
44 //arc_dep = torch::cat({ roots, elu(mlp_dep(input_t)) }, 1);
45 arc_dep = elu(torch::cat({ roots, mlp_dep(input_t) }, 1));
46 }
47 else
48 {
49 // h_i^(arc-dep) = MLP^(arc-dep)(r_i) in (4) from [Dozat&Manning,2017]
50 arc_dep = elu(mlp_dep(input_t));
51 }
52
53 //std::cerr << "arc_dep.sizes() == " << arc_dep.sizes() << std::endl;
54 //std::cerr << "U1.sizes() == " << U1.sizes() << std::endl;
55 torch::Tensor W = torch::matmul(arc_head, U1);
56 //std::cerr << "W.sizes() == " << W.sizes() << std::endl;
57 torch::Tensor Wx = torch::matmul(W, arc_dep.transpose(1, 2)); //
58 //std::cerr << "Wx.sizes() == " << Wx.sizes() << std::endl;
59
60 //std::cerr << "u2.sizes() == " << u2.sizes() << std::endl;
61 torch::Tensor b = torch::matmul(arc_head, torch::tile(u2, { 1, arc_head.size(1) })); // H^(arc-head)u^(2) in (6) from [Dozat&Manning,2017]
62
63 //std::cerr << "b.sizes() == " << b.sizes() << std::endl;
64 torch::Tensor r = torch::add(Wx, b);
65 //std::cerr << "r.sizes() == " << r.sizes() << std::endl;
66
68 {
69 r = r.index({ torch::indexing::Slice(),
70 torch::indexing::Slice(1, r.size(1)),
71 torch::indexing::Slice()});
72 }
73 //std::cerr << "r.sizes() == " << r.sizes() << std::endl;
74
75 return r;
76}
77
78} // torch_modules
79} // nets
80} // deeplima