LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
deep_biaffine_attention_label_decoder.h
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
6#ifndef DEEPLIMA_LIBS_NN_DEEP_BIAFFINE_ATTENTION_LABEL_DECODER_H
7#define DEEPLIMA_LIBS_NN_DEEP_BIAFFINE_ATTENTION_LABEL_DECODER_H
8
9#include <torch/torch.h>
10
11namespace deeplima
12{
13namespace nets
14{
15namespace torch_modules
16{
17
18// Biaffine label (deprel) scorer from Dozat & Manning, 2017
19// (https://nlp.stanford.edu/pubs/dozat2017deep.pdf), section 2.2 "Deep biaffine
20// scoring", equations (4)-(6) applied to the label scorer.
21//
22// For every ordered pair (dependent i, head j) it produces a score per relation
23// label. The caller selects, for each dependent, the scores at its head (the gold
24// head while training, the predicted head at inference).
25//
26// forward(input):
27// input: [seq_len, batch, input_dim] (the encoder states, time-major)
28// result: [batch, dep, head, num_labels] label logits.
29// When input_includes_root is false a learned <ROOT> row is prepended on the
30// head axis (so head index 0 is the artificial root and `head` == seq_len + 1),
31// and the root is removed from the dependent axis (`dep` == seq_len), matching
32// the convention of DeepBiaffineAttentionDecoder (the arc scorer).
33class DeepBiaffineAttentionLabelDecoderImpl : public torch::nn::Module
34{
35public:
37 int64_t hidden_dim,
38 int64_t num_labels,
39 bool input_includes_root = false)
41 m_hidden_dim(hidden_dim),
43 mlp_head(register_module("mlp_head", torch::nn::Linear(input_dim, hidden_dim))),
44 mlp_dep(register_module("mlp_dep", torch::nn::Linear(input_dim, hidden_dim))),
45 // Affine biaffine: the dependent and head representations are each augmented
46 // with a constant 1, so a single (hidden+1) x (hidden+1) matrix per label
47 // captures the bilinear term, both linear terms and the bias.
48 U(register_parameter("U", torch::randn({num_labels, hidden_dim + 1, hidden_dim + 1})
49 * (1.0 / std::sqrt(double(hidden_dim + 1))))),
50 root(register_parameter("root", torch::randn({1, 1, hidden_dim})))
51 {
52 }
53
54 int64_t num_labels() const { return m_num_labels; }
56
57 torch::Tensor forward(torch::Tensor input);
58
59 // Public (like DeepBiaffineAttentionDecoder) so the torch->eigen converter can
60 // read the parameters directly.
62 int64_t m_hidden_dim;
63 int64_t m_num_labels;
64 torch::nn::Linear mlp_head;
65 torch::nn::Linear mlp_dep;
66 torch::nn::ELU elu;
67 torch::Tensor U;
68 torch::Tensor root; // learned head-side <ROOT> row, used if input_includes_root == false
69};
70
71TORCH_MODULE(DeepBiaffineAttentionLabelDecoder);
72
73} // torch_modules
74} // nets
75} // deeplima
76
77#endif // DEEPLIMA_LIBS_NN_DEEP_BIAFFINE_ATTENTION_LABEL_DECODER_H
DeepBiaffineAttentionLabelDecoderImpl(int64_t input_dim, int64_t hidden_dim, int64_t num_labels, bool input_includes_root=false)
TORCH_MODULE(DeepBiaffineAttentionDecoder)