LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
tagging_impl.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_TAGGING_IMPL_H
7#define DEEPLIMA_TAGGING_IMPL_H
8
9#include "deeplima/ner.h"
10#include "deeplima/token_type.h"
14
16
17namespace deeplima
18{
19namespace tagging
20{
21
22namespace impl {
23
25{
27
28protected:
31
32public:
34 {
35 m_ptoken = p;
36 }
37
39 : m_stridx(stridx),
40 m_ptoken(nullptr)
41 { }
42
43 inline token_flags_t flags() const
44 {
45 assert(nullptr != m_ptoken);
46 return m_ptoken->m_flags;
47 }
48
49 inline bool eos() const
50 {
51 assert(nullptr != m_ptoken);
53 }
54
55 inline const std::string& form() const
56 {
57 assert(nullptr != m_ptoken);
58 const std::string& f = m_stridx.get_str(m_ptoken->m_form_idx);
59 return f;
60 }
61};
62
64{
65 const token_buffer_t<>& m_data;
66 mutable enriched_token_t m_token; // WARNING: only one iterator is supported
67
68public:
70
72 : m_data(data), m_token(stridx) { }
73
75 {
76 return m_data.size();
77 }
78
79 inline const enriched_token_t& operator[](size_t idx) const
80 {
81 m_token.set_token(m_data.data() + idx);
82 return m_token;
83 }
84};
85
86
92template <typename TaggingAuxScalar>
93class TaggingImpl: public EntityTaggingClassifier<TaggingAuxScalar>
94{
96 using Vectorizer = FeaturesVectorizerWithPrecomputing<
100public:
101
103 Classifier(),
104 m_fastText(std::make_shared<FastTextVectorizer<typename eigen_wrp::EigenMatrixXf::matrix_t, Eigen::Index>>()),
107 // m_current_slot_no(-1),
109 // m_curr_buff_idx(0)
110 {}
111
112 // TaggingImpl(
113 // size_t threads,
114 // size_t buffer_size_per_thread
115 // ) :
116 // Classifier(),
117 // m_fastText(std::make_shared<FastTextVectorizer<typename eigen_wrp::EigenMatrixXf::matrix_t, Eigen::Index>>()),
118 // m_current_timepoint(Classifier::get_start_timepoint()),
119 // m_current_slot_timepoints(0),
120 // // m_current_slot_no(-1),
121 // m_last_completed_slot(-1) //,
122 // // m_curr_buff_idx(0)
123 // {
124 // }
125
126 virtual ~TaggingImpl()
127 {
128 // std::cerr << "~TaggingImpl" << std::endl;
129 }
130
131 virtual void init(size_t threads, size_t num_buffers,
132 size_t buffer_size_per_thread, StringIndex& stridx)
133 {
134 m_fastText->get_words([&stridx](const std::string& word){ stridx.get_idx(word); });
135
136 m_vectorizer.init_features({
137 { Vectorizer::str_feature, "form", m_fastText }
138 });
139
140 m_vectorizer.set_model(this);
141
143 16, num_buffers, buffer_size_per_thread, threads,
144 m_vectorizer.is_precomputing());
145
147 }
148
149 virtual void reset()
150 {
151 // std::cerr << "TaggingImpl::reset" << std::endl;
156 }
157
158 virtual void load(const std::string& fn, const PathResolver& path_resolver)
159 {
161
162 if (this->get_input_str_dicts().size())
163 {
164 auto z = *(this->get_input_str_dicts().begin());
165 auto p = std::make_shared<EmbdStrFloat>(z);
166 }
167
168 auto fastText_fn = path_resolver.resolve("embd", Classifier::get_embd_fn(0), {"bin", "ftz"});
169 if (fastText_fn.empty())
170 {
171 throw std::runtime_error(std::string("Failed to resolve embedding file name with embd and ")+Classifier::get_embd_fn(0));
172 }
173 m_fastText->load(fastText_fn);
174 }
175
176 void precompute_inputs(const typename Vectorizer::dataset_t& buffer)
177 {
178 m_vectorizer.precompute(buffer);
179 }
180
181 typedef std::function < void (std::shared_ptr< StdMatrix<uint8_t> > classes,
182 size_t begin, size_t end, size_t slot_idx) > tagging_callback_t;
183
185 {
186 // std::cerr << "TaggingImpl::register_handler" << std::endl;
187 m_callback = fn;
188 }
189
190protected:
191
192 inline void increment_timepoint(uint64_t& timepoint)
193 {
194 assert(m_current_slot_timepoints > 0);
197 }
198
199 inline void send_results(int32_t slot_idx)
200 {
201 uint64_t from = Classifier::get_slot_begin(slot_idx);
202 const uint64_t to = Classifier::get_slot_end(slot_idx);
203
204 m_callback(Classifier::get_output(), from, to, slot_idx);
205
207 m_last_completed_slot = slot_idx;
208 }
209
210public:
211 inline void send_next_results()
212 {
213 int32_t slot_idx = m_last_completed_slot;
214 if (-1 == slot_idx)
215 {
216 slot_idx = 0;
217 }
218 else
219 {
220 slot_idx = Classifier::next_slot(slot_idx);
221 }
222
223 uint8_t lock_count = Classifier::get_lock_count(slot_idx);
224
225 while (lock_count > 1)
226 {
227 // Worker still uses this slot. Waiting...
228 // std::cerr << "TaggingImpl::send_next_results: waiting for slot " << slot_idx+1
229 // << " (lock_count==" << int(lock_count) << ")\n";
230 // Classifier::pretty_print();
232 lock_count = Classifier::get_lock_count(slot_idx);
233 }
234 if (1 == lock_count)
235 {
236 // Data is ready. We can return it to caller
237 send_results(slot_idx);
238 }
239 }
240
241 inline void send_all_results()
242 {
243 int32_t slot_idx = m_last_completed_slot;
244
245 while (true)
246 {
247 if (-1 == slot_idx)
248 {
249 slot_idx = 0;
250 }
251 else
252 {
253 slot_idx = Classifier::next_slot(slot_idx);
254 }
255
256 uint8_t lock_count = Classifier::get_lock_count(slot_idx);
257 if (0 == lock_count)
258 {
259 return;
260 }
261
262 while (lock_count > 1)
263 {
264 // Worker still uses this slot. Waiting...
265 // std::cerr << "TaggingImpl::send_all_results: Worker still uses this slot. Waiting... " << slot_idx+1
266 // << " (lock_count==" << int(lock_count) << ")\n";
267 // Classifier::pretty_print();
269 lock_count = Classifier::get_lock_count(slot_idx);
270 }
271 if (1 == lock_count)
272 {
273 send_results(slot_idx);
274 }
275 }
276 }
277
278protected:
280 {
281 auto slot_idx = m_last_completed_slot;
282 if (-1 == slot_idx)
283 {
284 slot_idx = 0;
285 }
286 else
287 {
288 slot_idx = Classifier::next_slot(slot_idx);
289 }
290
291 uint8_t lock_count = Classifier::get_lock_count(slot_idx);
292
293 if (1 == lock_count)
294 {
295 // Data is ready. We can return it to caller
296 send_results(slot_idx);
297 }
298 }
299
300 inline void acquire_slot(size_t slot_no)
301 {
302 // m_current_slot_no = Classifier::get_slot_idx(m_current_timepoint);
303 // std::cerr << "tagging acquiring_slot: " << slot_no << std::endl;
304 uint8_t lock_count = Classifier::get_lock_count(slot_no);
305
306 while (lock_count > 1)
307 {
308 // Worker still uses this slot. Waiting...
309 // std::cerr << "TaggingImpl::acquire_slot tagging handle_timepoint, waiting for slot " << slot_no
310 // << " lock_count=" << int(lock_count) << std::endl;
311 // Classifier::pretty_print();
313 lock_count = Classifier::get_lock_count(slot_no);
314 }
315 if (1 == lock_count)
316 {
317 // Data is ready. We can return it to caller
318 send_results(slot_no);
319 }
320
323 }
324
325public:
326 virtual void handle_token_buffer(size_t slot_no, const typename Vectorizer::dataset_t& buffer, int timepoints_to_analyze = -1)
327 {
328 // std::cerr << "TaggingImpl::handle_token_buffer " << slot_no << ", "
329 // << timepoints_to_analyze << std::endl;
331 acquire_slot(slot_no);
332 size_t offset = slot_no * buffer.size() + Classifier::get_start_timepoint();
333 size_t count = (timepoints_to_analyze > 0) ? timepoints_to_analyze : buffer.size();
334 for (size_t i = 0; i < count; i++)
335 {
336 m_vectorizer.vectorize_timepoint(eigen_wrp::EigenMatrixXf::get_tensor(), offset + i, buffer[i]);
337 }
338
339 Classifier::set_slot_end(slot_no, offset + count);
340 Classifier::start_job(slot_no, timepoints_to_analyze > 0);
341 // std::cerr << "Slot " << slot_no << " sent to inference engine (tagging)" << std::endl;
342 }
343
344 inline void no_more_data(size_t slot_no)
345 {
346 if (!Classifier::get_slot_started(slot_no))
347 {
348 while (Classifier::get_lock_count(slot_no) > 1)
349 {
351 }
352 Classifier::start_job(slot_no, true);
353 }
354 }
355
356protected:
357 Vectorizer m_vectorizer;
358 std::shared_ptr<FastTextVectorizer<typename eigen_wrp::EigenMatrixXf::matrix_t, Eigen::Index>> m_fastText;
359
361
364
365 // int32_t m_current_slot_no;
367
368 // size_t m_curr_buff_idx;
369};
370
371} // namespace impl
372} // namespace tagging
373} // namespace deeplima
374
375#endif
std::string resolve(const std::string &prefix, const std::string &path, const std::vector< std::string > &accepted_ext={}) const
bool get_slot_started(uint32_t idx) const
int32_t next_slot(uint32_t idx)
uint64_t get_slot_end(uint32_t idx) const
virtual void init(uint32_t max_feat, uint32_t overlap, uint32_t num_slots, uint32_t slot_len, uint32_t num_threads, bool precomputed_input=false)
void load(const std::string &fn)
void increment_timepoint(uint64_t &timepoint)
uint8_t get_lock_count(uint32_t idx) const
std::shared_ptr< StdMatrix< Out > > get_output()
void set_slot_end(uint32_t idx, uint64_t slot_end)
void increment_lock_count(uint32_t idx, uint8_t v=1)
void decrement_lock_count(uint32_t idx)
uint64_t get_slot_begin(uint32_t idx) const
void start_job(uint32_t idx, bool no_more_data=false)
virtual void reset()
Need to be called to be able to reuse this classifier on several sequences.
const S & get_str(const idx_t idx) const
Definition str_index.h:59
idx_t get_idx(const char *p, size_t len)
Definition str_index.h:24
A kind of RnnSequenceClassifier, used for named entities tagging (?), but also the parent of TaggingI...
Definition ner.h:129
Class implementing the tagger, used as member in TokenSequenceAnalyzer, the main tagger class Son of ...
virtual void register_handler(const tagging_callback_t fn)
virtual void load(const std::string &fn, const PathResolver &path_resolver)
void send_results(int32_t slot_idx)
virtual void init(size_t threads, size_t num_buffers, size_t buffer_size_per_thread, StringIndex &stridx)
virtual void handle_token_buffer(size_t slot_no, const typename Vectorizer::dataset_t &buffer, int timepoints_to_analyze=-1)
void increment_timepoint(uint64_t &timepoint)
virtual void reset()
Need to be called to be able to reuse this classifier on several sequences.
std::function< void(std::shared_ptr< StdMatrix< uint8_t > > classes, size_t begin, size_t end, size_t slot_idx) > tagging_callback_t
std::shared_ptr< FastTextVectorizer< typename eigen_wrp::EigenMatrixXf::matrix_t, Eigen::Index > > m_fastText
void precompute_inputs(const typename Vectorizer::dataset_t &buffer)
const enriched_token_t & operator[](size_t idx) const
token_buffer_t ::size_type size() const
enriched_token_buffer_t(const token_buffer_t<> &data, const StringIndex &stridx)
void set_token(const token_buffer_t<>::token_t *p)
const token_buffer_t ::token_t * m_ptoken
enriched_token_t(const StringIndex &stridx)
const std::string & form() const
@ sentence_brk
Definition token_type.h:21
STL namespace.