6#ifndef DEEPLIMA_SRC_INFERENCE_EIGEN_EMBD_DICT_H
7#define DEEPLIMA_SRC_INFERENCE_EIGEN_EMBD_DICT_H
9#include <unordered_map>
10#include <eigen3/Eigen/Dense>
22template <
class K,
class M,
class I=Eigen::Index>
34 void init(std::shared_ptr<
Dict<T>> dict,
const M& tensor,
bool transpose=
true)
38 m_embd = tensor.transpose();
45 for (
const auto& it : dict->get_v2i() )
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);
55 virtual I
dim()
const override
64 virtual void get(
const K key, M& target, I timepoint, I pos)
const override
69 inline void get_direct(
const I idx, M& target, I timepoint, I pos)
const
71 target.block(pos, timepoint,
m_dim, 1) =
m_embd.col(idx);
85 inline void get_static(
const K key, M& target, I timepoint, I pos)
const
88 target.block(pos, timepoint,
m_dim, 1) =
m_embd.col(idx);
111 return (
m_index.end() == i) ? 0 : i->second;
121template <
class M,
class I>
134 m_embd = tensor.transpose();
141 for (
const auto& it : dict->get_v2i() )
143 assert(it.second >= 0);
144 assert(it.second < std::numeric_limits<I>::max());
150 virtual I
dim()
const override
159 virtual void get(
const value_t& key, M& target, I timepoint, I pos)
const override
167 target.block(pos, timepoint,
m_dim, 1) =
m_embd.col(idx);
178 std::vector<value_t> vk;
180 for (
const auto& it :
m_index )
182 vk.push_back(it.first);
184 std::sort(vk.begin(), vk.end());
186 std::vector<K> vi(
m_index.size());
187 for (
size_t i = 0; i < vk.size(); ++i)
192 std::shared_ptr<Dict<K>> rv = std::make_shared<Dict<K>>(vi);
204 return (
m_index.end() == i) ? 0 : i->second;
std::vector< T >::size_type key_t
I lookup(const value_t &key) const
std::shared_ptr< Dict< K > > get_int_dict() const
virtual I dim() const override
const M & get_tensor() const
void init(std::shared_ptr< Dict< value_t > > dict, const M &tensor, bool transpose=true)
virtual void get(const value_t &key, M &target, I timepoint, I pos) const override
void get_static(const value_t &key, M &target, I timepoint, I pos) const
std::unordered_map< value_t, I > m_index
const M & get_tensor() const
I lookup(const K key) const
void set_tensor(const M &new_embd)
virtual I dim() const override
K decode(const I idx) const
void init(std::shared_ptr< Dict< T > > dict, const M &tensor, bool transpose=true)
void get_static(const K key, M &target, I timepoint, I pos) const
std::unordered_map< I, K > m_reverse_index
Eigen::Ref< const M > get_ref_by_key(const K key) const
std::unordered_map< K, I > m_index
Eigen::Ref< const M > get_ref_by_idx(const I idx) const
virtual void get(const K key, M &target, I timepoint, I pos) const override
void get_direct(const I idx, M &target, I timepoint, I pos) const
DictsHolderImpl< EmbdStrFloat > EmbdStrFloatHolder
DictsHolderImpl< EmbdUInt64Float > EmbdUInt64FloatHolder
EmbdDict< uint64_t, Eigen::MatrixXf, Eigen::Index > EmbdUInt64Float
EmbdDict< std::string, Eigen::MatrixXf, Eigen::Index > EmbdStrFloat