LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
deep_biaffine_attention_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_DECODER_H
7#define DEEPLIMA_LIBS_NN_DEEP_BIAFFINE_ATTENTION_DECODER_H
8
9#include <torch/torch.h>
10
11namespace deeplima
12{
13namespace nets
14{
15namespace torch_modules
16{
17
18class DeepBiaffineAttentionDecoderImpl : public torch::nn::Module
19{
20public:
21 DeepBiaffineAttentionDecoderImpl(int64_t input_dim, int64_t hidden_arc_dim, bool input_includes_root=false)
22 : m_input_includes_root(input_includes_root),
23 m_hidden_arc_dim(hidden_arc_dim),
24 mlp_head(register_module("mlp_head", torch::nn::Linear(input_dim, hidden_arc_dim))),
25 mlp_dep(register_module("mlp_dep", torch::nn::Linear(input_dim, hidden_arc_dim))),
26 U1(register_parameter("U1", torch::randn({hidden_arc_dim, hidden_arc_dim}))),
27 u2(register_parameter("u2", torch::randn({hidden_arc_dim, 1}))),
28 root(register_parameter("root", torch::randn({1, 1, hidden_arc_dim}))),
29 root2(register_parameter("root2", torch::randn({1, 1, hidden_arc_dim})))
30 {
31 // std::cerr << U1.dtype() << std::endl;
32 }
33
34 torch::Tensor forward(torch::Tensor input);
35
36 // see https://nlp.stanford.edu/pubs/dozat2017deep.pdf , page 3
39 torch::nn::Linear mlp_head;
40 torch::nn::Linear mlp_dep;
41 torch::nn::ELU elu; // this with the linear above makes the MLP from [Dozat&Manning,2017]
42 torch::Tensor U1, u2;
43 torch::Tensor root, root2; // used if input_includes_root == false
44};
45
46TORCH_MODULE(DeepBiaffineAttentionDecoder);
47
48} // torch_modules
49} // nets
50} // deeplima
51
52#endif // DEEPLIMA_LIBS_NN_DEEP_BIAFFINE_ATTENTION_DECODER_H
DeepBiaffineAttentionDecoderImpl(int64_t input_dim, int64_t hidden_arc_dim, bool input_includes_root=false)
TORCH_MODULE(DeepBiaffineAttentionDecoder)