LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
deep_biaffine_attention_label_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 DeepBiaffineAttentionLabelDecoderImpl::forward(torch::Tensor input)
16{
17 using torch::indexing::Slice;
18
19 const int64_t batch_size = input.size(1);
20
21 // [seq, batch, input_dim] -> [batch, seq, input_dim]
22 torch::Tensor input_t = input.transpose(0, 1);
23
24 // Per-token dependent and head representations (eq. (4)-(5) of Dozat&Manning).
25 torch::Tensor h_dep = elu(mlp_dep(input_t)); // [batch, seq, hidden]
26 torch::Tensor h_head = elu(mlp_head(input_t)); // [batch, seq, hidden]
27
29 {
30 // Prepend a learned <ROOT> row on the head axis so head index 0 is the root.
31 torch::Tensor roots = torch::tile(root, {batch_size, 1, 1}); // [batch, 1, hidden]
32 h_head = torch::cat({roots, h_head}, 1); // [batch, seq+1, hidden]
33 }
34
35 // Affine augmentation: append a constant 1 feature to fold the linear and bias
36 // terms into a single per-label matrix multiply.
37 auto append_ones = [](const torch::Tensor& t) {
38 torch::Tensor ones = torch::ones({t.size(0), t.size(1), 1}, t.options());
39 return torch::cat({t, ones}, 2);
40 };
41 torch::Tensor aug_dep = append_ones(h_dep); // [batch, dep, hidden+1]
42 torch::Tensor aug_head = append_ones(h_head); // [batch, head, hidden+1]
43
44 // s[b, i, j, l] = aug_dep[b,i,:]^T U[l] aug_head[b,j,:]
45 // -> [batch, dep, head, num_labels]
46 torch::Tensor logits = torch::einsum("bxi,lij,byj->bxyl", {aug_dep, U, aug_head});
47
48 return logits;
49}
50
51} // torch_modules
52} // nets
53} // deeplima