LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
graph_dp_model.cpp
Go to the documentation of this file.
1// Copyright 2021 CEA LIST
2// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
6#include <algorithm>
7
10
11#include "convert_from_torch.h"
12
13using namespace std;
14using namespace torch;
15using namespace deeplima::convert_from_torch;
16using namespace deeplima::eigen_impl;
17
18namespace deeplima
19{
20namespace graph_dp
21{
22namespace eigen_impl
23{
24
25void convert_classes(const DictsHolder& src, vector<vector<string>>& classes);
26
28{
30 torch::load(src, fn, torch::Device(torch::kCPU));
31
32 // dicts and embeddings
34 m_embd_fn.push_back(src.get_embd_fn(0));
35 assert(m_embd_fn[0].size() > 0);
36
39
40 // torch modules
41 Parent::m_lstm.reserve(src.get_layers_lstm().size());
42 for (size_t i = 0; i < src.get_layers_lstm().size(); i++)
43 {
44 const std::string name = src.get_module_name(i, "lstm");
45 Parent::m_lstm_idx[name] = i;
46
47 const nn::LSTM& m = src.get_layers_lstm()[i];
48 Parent::m_lstm.emplace_back(typename Parent::params_bilstm_spec_t());
49 typename Parent::params_bilstm_spec_t& layer = Parent::m_lstm.back();
50
52 }
53
54 Parent::m_multi_bilstm.emplace_back(std::make_shared<typename Parent::params_multilayer_bilstm_spec_t>(Parent::m_lstm));
55 Parent::m_multi_bilstm_idx["encoder"] = 0;
56
57 Parent::m_linear.reserve(src.get_layers_linear().size());
58 for (size_t i = 0; i < src.get_layers_linear().size(); i++)
59 {
60 const std::string name = src.get_module_name(i, "linear");
61 Parent::m_linear_idx[name] = i;
62
63 const nn::Linear& m = src.get_layers_linear()[i];
66
68 }
69
71
72 for (size_t i = 0; i < src.get_layers_deep_biaffine_attn_decoder().size(); i++)
73 {
74 const std::string name = src.get_module_name(i, "deep_biaffine_attention_decoder");
76
77 const deeplima::nets::torch_modules::DeepBiaffineAttentionDecoder& m
79
81 auto& layer = *m_deep_biaffine_attn_decoder.back().get();
82
84 }
85
86 // Label (deprel) decoder, present only in labeled models.
87 for (size_t i = 0; i < src.get_layers_deep_biaffine_attn_label_decoder().size(); i++)
88 {
89 const deeplima::nets::torch_modules::DeepBiaffineAttentionLabelDecoder& m
91
93 auto& layer = *m_deep_biaffine_attn_label_decoder.back().get();
94
96 }
98
99 // The arc head is the only generic "task" the model declares, so the output
100 // buffer would be sized to a single column. The deprel (rel) is produced by a
101 // separate label decoder rather than the generic task pool, so when one is
102 // present we add a second output column for it. predict() fills column 0 with
103 // heads and column 1 with deprel ids.
105 && std::find(Parent::m_output_str_dicts_names.begin(),
108 {
109 Parent::m_output_str_dicts_names.push_back("rel");
110 }
111
112 // temp: create exec plan
115
118
119 Parent::m_wb.resize(2);
120
121 // tags
122 // std::cerr << "TAGS:" << std::endl;
123 // for ( const auto& it : src.get_tags() )
124 // {
125 // std::cerr << "\t" << it.first << " = " << it.second << std::endl;
126 // }
127 // std::cerr << std::endl;
128}
129
130void convert_classes(const DictsHolder& src, vector<vector<string>>& classes)
131{
132 classes.resize(src.size());
133 for (size_t i = 0; i < classes.size(); ++i)
134 {
135 shared_ptr<StringDict> d = dynamic_pointer_cast<StringDict, DictBase>(src[i]);
136 classes[i].reserve(d->size());
137 for (size_t j = 0; j < d->size(); ++j)
138 {
139 classes[i].push_back(d->get_value(j));
140 }
141 }
142}
143
144} // namespace eigen_impl
145} // namespace graph_dp
146} // namespace deeplima
147
std::map< std::string, size_t > m_multi_bilstm_idx
std::vector< std::vector< std::shared_ptr< Op_Base::workbench_t > > > m_wb
std::vector< std::shared_ptr< Op_Base > > m_ops
std::vector< std::string > m_output_str_dicts_names
virtual void convert_dicts_and_embeddings(const nets::BiRnnClassifierImpl &src)
std::map< std::string, size_t > m_linear_idx
std::vector< std::string > m_input_str_dicts_names
std::vector< std::shared_ptr< params_multilayer_bilstm_spec_t > > m_multi_bilstm
std::map< std::string, size_t > m_lstm_idx
std::vector< params_linear_t< Eigen::MatrixXf, Eigen::VectorXf > > m_linear
std::vector< params_bilstm_spec_t > m_lstm
std::vector< std::shared_ptr< param_base_t > > m_params
std::vector< std::shared_ptr< deeplima::eigen_impl::params_deep_biaffine_attn_decoder_t< Eigen::MatrixXf, Eigen::VectorXf > > > m_deep_biaffine_attn_decoder
std::vector< std::shared_ptr< deeplima::eigen_impl::params_deep_biaffine_attn_label_decoder_t< Eigen::MatrixXf, Eigen::VectorXf > > > m_deep_biaffine_attn_label_decoder
const std::vector< torch::nn::Linear > & get_layers_linear() const
const std::vector< deeplima::nets::torch_modules::DeepBiaffineAttentionDecoder > & get_layers_deep_biaffine_attn_decoder() const
const std::vector< deeplima::nets::torch_modules::DeepBiaffineAttentionLabelDecoder > & get_layers_deep_biaffine_attn_label_decoder() const
const std::string get_module_name(size_t idx, const std::string &type) const
const std::vector< torch::nn::LSTM > & get_layers_lstm() const
void convert_module_from_torch(const torch::nn::LSTM &src, eigen_impl::params_bilstm_t< M, V > &dst)
void convert_classes(const DictsHolder &src, vector< vector< string > > &classes)
STL namespace.