LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
fastText_wrp.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_FASTTEXT_WRP_H
7#define DEEPLIMA_FASTTEXT_WRP_H
8
9#include "fastText/src/fasttext.h"
10#include <eigen3/Eigen/Dense>
11
12namespace deeplima
13{
14
15template <class Idx=uint64_t>
17{
18public:
19 virtual ~FeatureVectorizerBase() = default;
20 virtual Idx dim() const = 0;
21};
22
23template <class Matrix, class Type, class Idx=uint64_t>
25{
26public:
27 virtual ~FeatureVectorizerToMatrix() = default;
28 virtual void get(const Type value, Matrix& target, Idx time, Idx pos) const = 0;
29
30};
31
32template <class Matrix, class Idx=uint64_t>
33class DirectDict : public deeplima::FeatureVectorizerToMatrix<Matrix, uint64_t, Idx>
34{
35public:
36 DirectDict(size_t ) {}
37
38 virtual Idx dim() const override
39 {
40 return 1;
41 }
42
43 virtual void get(const uint64_t value, Matrix& target, Idx time, Idx feat) const override
44 {
45 target.set(time, feat, value);
46 }
47};
48
49template <>
50class DirectDict<Eigen::MatrixXf, Eigen::Index> : public deeplima::FeatureVectorizerToMatrix<Eigen::MatrixXf, uint64_t, Eigen::Index>
51{
52public:
53 DirectDict(size_t )
54 {
55 }
56
57 virtual Eigen::Index dim() const override
58 {
59 return 1;
60 }
61
62 virtual void get(const uint64_t value, Eigen::MatrixXf& target, Eigen::Index time, Eigen::Index feat) const override
63 {
64 target.block<1,1>(feat, time) << value;
65 //target.set(time, feat, value);
66 }
67};
68
69class FastTextExt : public fasttext::FastText
70{
71
72};
73
74template <class Matrix, class Idx=uint64_t>
75class FastTextVectorizer : public FeatureVectorizerToMatrix<Matrix, const std::string&, Idx>
76{
77protected:
78 fasttext::FastText m_fasttext;
79 Idx m_dim;
80 std::shared_ptr<fasttext::Vector> m_vec;
81public:
82
83 FastTextVectorizer(const std::string& fn = "")
84 : m_dim(0),
85 m_vec(nullptr)
86 {
87 if (!fn.empty())
88 {
89 load(fn);
90 }
91 }
92
93 virtual ~FastTextVectorizer() = default;
94
95 virtual void load(const std::string& fn)
96 {
97 if (fn.empty())
98 {
99 throw std::invalid_argument("empty file name in FastTextVectorizer::load()");
100 }
101 m_fasttext.loadModel(fn, false);
102
103 m_dim = m_fasttext.getDimension();
104 assert(m_dim > 0);
105 m_vec = std::make_shared<fasttext::Vector>(m_dim);
106 assert(nullptr != m_vec);
107 m_vec->zero();
108 }
109
110 virtual Idx dim() const override
111 {
112 return m_dim;
113 }
114
115 virtual void get(const std::string& value, Matrix& target, Idx time, Idx pos) const override
116 {
117 assert(m_dim > 0);
118 assert(nullptr != m_vec);
119
120 m_fasttext.getWordVector(*m_vec, value);
121 for (Idx i = 0; i < m_dim; i++)
122 {
123 target.set(time, pos + i, (*m_vec)[i]);
124 }
125 }
126};
127
128template <>
129class FastTextVectorizer<Eigen::MatrixXf, Eigen::Index>
130 : public FeatureVectorizerToMatrix<Eigen::MatrixXf, const std::string&, Eigen::Index>
131{
132protected:
133 fasttext::FastText m_fasttext;
134 Eigen::Index m_dim;
135 std::shared_ptr<fasttext::Vector> m_vec;
136public:
137
138 FastTextVectorizer(const std::string& fn = "")
139 : m_dim(0),
140 m_vec(nullptr)
141 {
142 if (!fn.empty())
143 {
144 load(fn);
145 }
146 }
147
148 virtual ~FastTextVectorizer() = default;
149
150 virtual void load(const std::string& fn)
151 {
152 if (fn.empty())
153 {
154 throw std::invalid_argument("empty file name in FastTextVectorizer::load()");
155 }
156 m_fasttext.loadModel(fn);
157
158 m_dim = m_fasttext.getDimension();
159 assert(m_dim > 0);
160 m_vec = std::make_shared<fasttext::Vector>(m_dim);
161 assert(nullptr != m_vec);
162 m_vec->zero();
163 }
164
165 virtual Eigen::Index dim() const override
166 {
167 return m_dim;
168 }
169
170 virtual void get(const std::string& value, Eigen::MatrixXf& target, Eigen::Index time, Eigen::Index pos) const override
171 {
172 assert(m_dim > 0);
173 assert(nullptr != m_vec);
174
175 m_fasttext.getWordVector(*m_vec, value);
176 auto blk = target.block(pos, time, m_dim, 1);
177 for (Eigen::Index i = 0; i < m_dim; i++)
178 {
179 blk(i, 0) = (*m_vec)[i];
180 }
181 }
182
183 typedef std::function< void (const std::string& word) > word_callback_t;
185 {
186 std::shared_ptr<const fasttext::Dictionary> pd = m_fasttext.getDictionary();
187 for (int32_t i = 0; i < pd->nwords(); ++i)
188 {
189 fn(pd->getWord(i));
190 }
191 }
192};
193
194}
195
196#endif
virtual void get(const uint64_t value, Eigen::MatrixXf &target, Eigen::Index time, Eigen::Index feat) const override
virtual Eigen::Index dim() const override
virtual Idx dim() const override
virtual void get(const uint64_t value, Matrix &target, Idx time, Idx feat) const override
std::function< void(const std::string &word) > word_callback_t
virtual void get(const std::string &value, Eigen::MatrixXf &target, Eigen::Index time, Eigen::Index pos) const override
FastTextVectorizer(const std::string &fn="")
virtual Idx dim() const override
virtual void get(const std::string &value, Matrix &target, Idx time, Idx pos) const override
std::shared_ptr< fasttext::Vector > m_vec
virtual ~FastTextVectorizer()=default
fasttext::FastText m_fasttext
virtual void load(const std::string &fn)
virtual Idx dim() const =0
virtual ~FeatureVectorizerBase()=default
virtual ~FeatureVectorizerToMatrix()=default
virtual void get(const Type value, Matrix &target, Idx time, Idx pos) const =0