LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
embd_dict.h
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#ifndef DEEPLIMA_SRC_INFERENCE_EIGEN_EMBD_DICT_H
7#define DEEPLIMA_SRC_INFERENCE_EIGEN_EMBD_DICT_H
8
9#include <unordered_map>
10#include <eigen3/Eigen/Dense>
11
13#include "static_graph/dict.h"
14
16
17namespace deeplima
18{
19
20// K - dict key
21// M - Eigen Matrix (Xf or Xd)
22template <class K, class M, class I=Eigen::Index>
23class EmbdDict : public FeatureVectorizerToMatrix<M, K, I>
24{
25public:
26 typedef K value_t;
27 typedef M tensor_t;
28 typedef I index_t;
29
30 EmbdDict() = default;
31 virtual ~EmbdDict() {}
32
33 template <class T>
34 void init(std::shared_ptr<Dict<T>> dict, const M& tensor, bool transpose=true)
35 {
36 if (transpose)
37 {
38 m_embd = tensor.transpose();
39 }
40 else
41 {
42 m_embd = tensor;
43 }
44 m_dim = m_embd.rows();
45 for ( const auto& it : dict->get_v2i() )
46 {
47 assert(it.second >= 0);
48 assert(it.second < std::numeric_limits<I>::max());
49 assert(m_embd.cols() >=0 && it.second < (unsigned long)(m_embd.cols()));
50 m_index[static_cast<K>(it.first)] = I(it.second);
51 m_reverse_index[static_cast<I>(it.second)] = static_cast<K>(it.first);
52 }
53 }
54
55 virtual I dim() const override
56 {
57 return m_dim;
58 }
59
60 // key - the value we are going to convert
61 // target - target matrix
62 // timepoint - column number
63 // pos - feature position (first row number)
64 virtual void get(const K key, M& target, I timepoint, I pos) const override
65 {
66 get_static(key, target, timepoint, pos);
67 }
68
69 inline void get_direct(const I idx, M& target, I timepoint, I pos) const
70 {
71 target.block(pos, timepoint, m_dim, 1) = m_embd.col(idx);
72 }
73
74 inline Eigen::Ref<const M> get_ref_by_idx(const I idx) const
75 {
76 return m_embd.col(idx);
77 }
78
79 inline Eigen::Ref<const M> get_ref_by_key(const K key) const
80 {
81 I idx = lookup(key);
82 return m_embd.col(idx);
83 }
84
85 inline void get_static(const K key, M& target, I timepoint, I pos) const
86 {
87 I idx = lookup(key);
88 target.block(pos, timepoint, m_dim, 1) = m_embd.col(idx);
89 }
90
91 const M& get_tensor() const
92 {
93 return m_embd;
94 }
95
96 void set_tensor(const M& new_embd)
97 {
98 m_embd = new_embd;
99 m_dim = m_embd.rows();
100 }
101
102 inline K decode(const I idx) const
103 {
104 auto it = m_reverse_index.find(idx);
105 return it->second;
106 }
107
108 inline I lookup(const K key) const
109 {
110 auto i = m_index.find(key);
111 return (m_index.end() == i) ? 0 : i->second;
112 }
113
114protected:
117 std::unordered_map<K, I> m_index;
118 std::unordered_map<I, K> m_reverse_index;
119};
120
121template <class M, class I>
122class EmbdDict<std::string, M, I> : public FeatureVectorizerToMatrix<M, const std::string&, I>
123{
124public:
125 typedef std::string value_t;
126
129
130 void init(std::shared_ptr<Dict<value_t>> dict, const M& tensor, bool transpose=true)
131 {
132 if (transpose)
133 {
134 m_embd = tensor.transpose();
135 }
136 else
137 {
138 m_embd = tensor;
139 }
140 m_dim = m_embd.rows();
141 for ( const auto& it : dict->get_v2i() )
142 {
143 assert(it.second >= 0);
144 assert(it.second < std::numeric_limits<I>::max());
145 assert(m_embd.cols() >= 0 && it.second < Dict<value_t>::key_t(m_embd.cols()));
146 m_index[value_t(it.first)] = I(it.second);
147 }
148 }
149
150 virtual I dim() const override
151 {
152 return m_dim;
153 }
154
155 // key - the value we are going to convert
156 // target - target matrix
157 // timepoint - column number
158 // pos - feature position (first row number)
159 virtual void get(const value_t& key, M& target, I timepoint, I pos) const override
160 {
161 get_static(key, target, timepoint, pos);
162 }
163
164 inline void get_static(const value_t& key, M& target, I timepoint, I pos) const
165 {
166 I idx = lookup(key);
167 target.block(pos, timepoint, m_dim, 1) = m_embd.col(idx);
168 }
169
170 const M& get_tensor() const
171 {
172 return m_embd;
173 }
174
175 template <class K>
176 std::shared_ptr<Dict<K>> get_int_dict() const
177 {
178 std::vector<value_t> vk;
179 vk.reserve(m_index.size());
180 for ( const auto& it : m_index )
181 {
182 vk.push_back(it.first);
183 }
184 std::sort(vk.begin(), vk.end());
185
186 std::vector<K> vi(m_index.size());
187 for (size_t i = 0; i < vk.size(); ++i)
188 {
189 vi[i] = lookup(vk[i]);
190 }
191
192 std::shared_ptr<Dict<K>> rv = std::make_shared<Dict<K>>(vi);
193 return rv;
194 }
195
196protected:
199 std::unordered_map<value_t, I> m_index;
200
201 inline I lookup(const value_t& key) const
202 {
203 auto i = m_index.find(key);
204 return (m_index.end() == i) ? 0 : i->second;
205 }
206};
207
212
213} // namespace deeplima
214
215#endif
std::vector< T >::size_type key_t
Definition dict.h:35
I lookup(const value_t &key) const
Definition embd_dict.h:201
std::shared_ptr< Dict< K > > get_int_dict() const
Definition embd_dict.h:176
virtual I dim() const override
Definition embd_dict.h:150
void init(std::shared_ptr< Dict< value_t > > dict, const M &tensor, bool transpose=true)
Definition embd_dict.h:130
virtual void get(const value_t &key, M &target, I timepoint, I pos) const override
Definition embd_dict.h:159
void get_static(const value_t &key, M &target, I timepoint, I pos) const
Definition embd_dict.h:164
std::unordered_map< value_t, I > m_index
Definition embd_dict.h:199
const M & get_tensor() const
Definition embd_dict.h:91
I lookup(const K key) const
Definition embd_dict.h:108
void set_tensor(const M &new_embd)
Definition embd_dict.h:96
virtual I dim() const override
Definition embd_dict.h:55
K decode(const I idx) const
Definition embd_dict.h:102
virtual ~EmbdDict()
Definition embd_dict.h:31
void init(std::shared_ptr< Dict< T > > dict, const M &tensor, bool transpose=true)
Definition embd_dict.h:34
void get_static(const K key, M &target, I timepoint, I pos) const
Definition embd_dict.h:85
std::unordered_map< I, K > m_reverse_index
Definition embd_dict.h:118
Eigen::Ref< const M > get_ref_by_key(const K key) const
Definition embd_dict.h:79
std::unordered_map< K, I > m_index
Definition embd_dict.h:117
Eigen::Ref< const M > get_ref_by_idx(const I idx) const
Definition embd_dict.h:74
virtual void get(const K key, M &target, I timepoint, I pos) const override
Definition embd_dict.h:64
void get_direct(const I idx, M &target, I timepoint, I pos) const
Definition embd_dict.h:69
DictsHolderImpl< EmbdStrFloat > EmbdStrFloatHolder
Definition embd_dict.h:211
DictsHolderImpl< EmbdUInt64Float > EmbdUInt64FloatHolder
Definition embd_dict.h:210
EmbdDict< uint64_t, Eigen::MatrixXf, Eigen::Index > EmbdUInt64Float
Definition embd_dict.h:208
EmbdDict< std::string, Eigen::MatrixXf, Eigen::Index > EmbdStrFloat
Definition embd_dict.h:209
STL namespace.