LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
birnn_seq_cls.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_RNN_SEQ_CLS_H
7#define DEEPLIMA_SRC_INFERENCE_RNN_SEQ_CLS_H
8
9#include <chrono>
10#include <cstdlib>
11#include <iterator>
12#include <vector>
13
15
17
18namespace deeplima
19{
20
21template <typename T>
22std::ostream& operator<< (std::ostream& out, const std::vector<T>& v) {
23 out << '[';
24 if ( !v.empty() ) {
25 std::copy (v.begin(), v.end(), std::ostream_iterator<T>(out, ", "));
26 }
27 out << "]";
28 return out;
29}
30
45template <class Model, class InputVectorizer/*=TorchMatrix<int64_t>*/, class Out>
46class RnnSequenceClassifier : public InputVectorizer,
47 public ThreadPool< RnnSequenceClassifier<Model, InputVectorizer, Out> >,
48 public Model
49{
50public:
52 : m_overlap(0),
53 m_num_slots(0),
54 m_slot_len(0),
55 m_slots(),
56 m_lengths(),
57 m_output(std::make_shared< StdMatrix<Out> >())
58 {}
59
60 // RnnSequenceClassifier(uint32_t max_feat,
61 // uint32_t overlap,
62 // uint32_t num_slots,
63 // uint32_t slot_len,
64 // uint32_t num_threads)
65 // : m_overlap(0),
66 // m_num_slots(0),
67 // m_slot_len(0),
68 // m_slots(),
69 // m_lengths(),
70 // m_output(std::make_shared< StdMatrix<Out> >())
71 // {
72 // init(max_feat, overlap, num_slots, slot_len, num_threads);
73 // }
74
76 {
77 // std::cerr << "-> ~RnnSequenceClassifier" << std::endl;
79 // std::cerr << "<- ~RnnSequenceClassifier" << std::endl;
80 }
81
82protected:
86
87 enum slot_flags_t : uint8_t
88 {
89 none = 0x00,
93 };
94
95
100 struct slot_t
101 {
102 // input positions
104 uint64_t m_input_end;
105
106 // output positions
108 uint64_t m_output_end;
109
111
113 std::atomic<uint8_t> m_lock_count;
114
117
118 std::vector<size_t> m_lengths;
119
121 : m_input_begin(0),
122 m_input_end(0),
124 m_output_end(0),
125 m_flags(none),
126 m_work_started(false),
127 m_lock_count(0),
128 m_prev(nullptr),
129 m_next(nullptr),
130 m_lengths()
131 { }
144 ~slot_t() = default;
146 {
151 m_flags = s.m_flags;
153 m_lock_count = 0;
154 m_prev = nullptr;
155 m_next = nullptr;
157 return *this;
158 }
159 };
160
161 inline int32_t prev_slot(uint32_t idx)
162 {
163 assert(idx < m_num_slots);
164 return (idx == 0) ? (m_num_slots - 1) : idx - 1;
165 }
166
167protected:
168 inline void clear_slot(uint32_t idx)
169 {
170 assert(idx < m_num_slots);
171 slot_t& slot = m_slots[idx];
172
173 slot.m_work_started = false;
174 //slot.m_done = false;
175 // workaround after first use of the first slot
176 if (0 == idx /*&& 0 == (slot.m_flags & left_overlap)*/)
177 {
178 assert(slot.m_output_begin >= slot.m_input_begin);
179 //slot.m_flags = slot_flags_t(left_overlap | right_overlap);
181 }
182 }
183
185 inline void start_job_impl(uint32_t idx)
186 {
187 assert(idx < m_num_slots);
188 slot_t& slot = m_slots[idx];
189 // std::cerr << "RnnSequenceClassifier::start_job_impl slot id=" << (idx+1)
190 // << ", slot begin=" << slot.m_input_begin
191 // << ", slot end=" << slot.m_input_end
192 // << ", m_lengths=" << m_lengths
193 // << std::endl;
194
195 if (! slot.m_work_started)
196 {
197 slot.m_work_started = true;
199 }
200 }
201
202 inline static void run_one_job(ThisClass* this_ptr, size_t worker_id, void* p)
203 {
204 slot_t& slot = *((slot_t*)p);
205
206 // std::cerr << "RnnSequenceClassifier::run_one_job worker " << worker_id
207 // << ". slot lock_count=" << int(slot.m_lock_count)
208 // << "; begin= " << slot.m_input_begin
209 // << "; end=" << slot.m_input_end
210 // << "; flags= " << int(slot.m_flags)
211 // << "; prev=" << (void*)slot.m_prev
212 // << "; next=" << (void*)slot.m_next
213 // << "; lengths=" << slot.m_lengths
214 // << ". lengths=" << this_ptr->m_lengths
215 // << std::endl;
216 // this_ptr->pretty_print();
217
218 this_ptr->predict(worker_id,
219 this_ptr->get_tensor(),
220 slot.m_input_begin, slot.m_input_end,
221 slot.m_output_begin, slot.m_output_end,
222 this_ptr->m_output,
223 slot.m_lengths,
224 // this_ptr->m_lengths[worker_id],
225 {"tokens"});
226
227 // std::cerr << "RnnSequenceClassifier::run_one_job after predict worker " << worker_id
228 // << ". slot lock_count=" << int(slot.m_lock_count)
229 // << "; begin= " << slot.m_input_begin
230 // << "; end=" << slot.m_input_end
231 // << "; flags= " << int(slot.m_flags)
232 // << "; prev=" << (void*)slot.m_prev
233 // << "; next=" << (void*)slot.m_next
234 // << "; output=" << (*(this_ptr->m_output))[0]
235 // << std::endl;
236 // this_ptr->pretty_print();
237
238 assert(slot.m_lock_count > 0);
239 slot.m_lock_count--;
240 if (slot.m_flags & left_overlap && slot.m_prev != nullptr)
241 {
242 assert(slot.m_prev->m_lock_count > 0);
243 // slot.m_prev->m_lock_count--;
244 slot.m_prev->m_lock_count--;
245 }
246 if (slot.m_flags & right_overlap && slot.m_next != nullptr)
247 {
248 assert(slot.m_next->m_lock_count > 0);
249 slot.m_next->m_lock_count--;
250 }
251 }
252
253public:
254 inline int32_t next_slot(uint32_t idx)
255 {
256 assert(idx < m_num_slots);
257 return (idx == m_num_slots - 1) ? 0 : idx + 1;
258 }
259
260
261 std::shared_ptr< StdMatrix<Out> > get_output()
262 {
263 return m_output;
264 }
265
269 virtual void reset()
270 {
271 // std::cerr << "RnnSequenceClassifier::reset()" << (void*)this << std::endl;
272 m_slots.clear();
273 m_slots.resize(m_num_slots);
274 for (size_t i = 0; i < m_num_slots; i++)
275 {
276 m_slots[i] = slot_t();
277 slot_t& slot = m_slots[i];
278
282 slot.m_input_end = slot.m_output_end + m_overlap;
283 if (m_overlap > 0)
284 {
286
287 uint32_t prev_idx = prev_slot(i);
288 slot.m_prev = &(m_slots[prev_idx]);
289
290 uint32_t next_idx = next_slot(i);
291 slot.m_next = &(m_slots[next_idx]);
292 }
293
294 if (0 == i)
295 {
296 slot.m_input_begin = slot.m_output_begin;
297 slot.m_flags = right_overlap;
298 }
299
300 if (m_num_slots - 1 == i)
301 {
302 slot.m_flags = left_overlap;
303 }
304 // std::cerr << "RnnSequenceClassifier::reset slot " << (i+1)
305 // << " input begin=" << slot.m_input_begin << ", end=" << slot.m_input_end
306 // << " output begin=" << slot.m_output_begin << ", end=" << slot.m_output_end
307 // << std::endl;
308 }
309 }
310
311
312 virtual void init(uint32_t max_feat,
313 uint32_t overlap,
314 uint32_t num_slots,
315 uint32_t slot_len,
316 uint32_t num_threads,
317 bool precomputed_input=false)
318 {
319 // RnnSequenceClassifier::init 7, 4, 18, 16384, 8, false
320 // RnnSequenceClassifier::init 1024, 16, 8, 1024, 1, true
321 // RnnSequenceClassifier::init 464, 0, 8, 1024, 1, false
322
323 // std::cerr << "RnnSequenceClassifier::init "<<(void*)this<<" max_feat=" << max_feat << ", overlap=" << overlap
324 // << ", num_slots=" << num_slots
325 // << ", slot_len=" << slot_len
326 // << ", num_threads=" << num_threads
327 // << ", precomputed_input=" << precomputed_input << std::endl;
328 m_num_slots = num_slots;
329 m_overlap = overlap;
330 m_slot_len = slot_len;
331
332 InputVectorizer::init(/*Model::get_dicts(),*/ m_num_slots * m_slot_len + m_overlap * 2, max_feat);
333 //InputVectorizer::set_dicts(Model::get_dicts());
334 for (size_t i = 0; i < num_threads; i++)
335 {
336 Model::init_new_worker(m_slot_len + m_overlap * 2, precomputed_input); // skip id - all workers are identical
337 }
339
340 reset(); // set up slots
341
342 m_lengths.resize(m_num_slots);
343
344 // Vector for calculation results
345 m_output->resize(Model::get_output_str_dicts_names().size());
346 assert(m_output->size() > 0);
347 for (auto& v : m_output->m_tensor)
348 {
349 v.resize(InputVectorizer::size());
350 assert(v.size() > 0);
351 }
352 }
353
354 void load(const std::string& fn)
355 {
356 Model::load(fn);
357 //InputVectorizer::set_dicts(Model::get_dicts());
358 }
359
360 void get_classes_from_fn(const std::string& fn, std::vector<std::string>& classes_names, std::vector<std::vector<std::string>>& classes){
361 Model::get_classes_from_fn(fn, classes_names, classes);
362 }
363
364 inline uint8_t get_output(uint64_t pos, uint8_t cls)
365 {
366 assert(cls < m_output->size());
367 uint32_t idx = get_slot_idx(pos);
368 assert(m_slots[idx].m_lock_count == 1);
369 assert(m_slots[idx].m_work_started);
370 return (*m_output)[cls][pos];
371 }
372
373 inline uint64_t get_slot_begin(uint32_t idx) const
374 {
375 assert(idx < m_num_slots);
376 return m_slots[idx].m_output_begin;
377 }
378
379 inline bool get_slot_started(uint32_t idx) const
380 {
381 assert(idx < m_num_slots);
382 return m_slots[idx].m_work_started;
383 }
384
385 inline uint64_t get_slot_end(uint32_t idx) const
386 {
387 assert(idx < m_num_slots);
388 return m_slots[idx].m_output_end;
389 }
390
391 inline uint8_t get_lock_count(uint32_t idx) const
392 {
393 assert(idx < m_num_slots);
394 return m_slots[idx].m_lock_count;
395 }
396
397 inline void increment_lock_count(uint32_t idx, uint8_t v = 1)
398 {
399 assert(idx < m_num_slots);
400 m_slots[idx].m_lock_count += v;
401 // std::cerr << "RnnSequenceClassifier::increment_lock_count by " << int(v)
402 // << " for slot " << int(idx+1)
403 // << ". it is now: " << int(m_slots[idx].m_lock_count) << std::endl;
404 // pretty_print();
405 }
406
407 inline void decrement_lock_count(uint32_t idx)
408 {
409 assert(idx < m_num_slots);
410 assert(m_slots[idx].m_lock_count > 0);
411
412 m_slots[idx].m_lock_count--;
413 if (0 == m_slots[idx].m_lock_count)
414 {
415 clear_slot(idx);
416 }
417 // std::cerr << "RnnSequenceClassifier::decrement_lock_count Lock for slot " << int(idx+1)
418 // << " set to " << int(m_slots[idx].m_lock_count) << std::endl;
419 // pretty_print();
420 }
421
422 inline uint64_t get_start_timepoint() const
423 {
424 // std::cerr << "RnnSequenceClassifier::get_start_timepoint return " << m_overlap << std::endl;
425 // pretty_print();
426 return m_overlap;
427 }
428
429 inline void increment_timepoint(uint64_t& timepoint)
430 {
431 timepoint++;
432 if (timepoint >= InputVectorizer::size() - m_overlap)
433 {
434 timepoint = get_start_timepoint();
435 }
436 }
437
438 inline uint32_t get_num_slots() const
439 {
440 return m_num_slots;
441 }
442
443 inline uint32_t get_slot_size() const
444 {
445 return m_slot_len;
446 }
447
448 inline int32_t get_slot_idx(uint64_t timepoint) const
449 {
450 assert(timepoint < m_num_slots * m_slot_len);
451 uint32_t rv = timepoint / m_slot_len;
452 assert(rv < m_num_slots);
453 return rv;
454 }
455
456 inline void set_slot_lengths(uint32_t idx, const std::vector<size_t>& lengths)
457 {
458 // std::cerr << "RnnSequenceClassifier::set_slot_lengths slot id=" << (idx+1) << ", lengths=" << lengths << std::endl;
459 assert(idx < m_num_slots);
460 assert(idx < m_lengths.size());
461
462 m_lengths[idx] = lengths;
463 auto& slot = m_slots[idx];
464 slot.m_lengths = lengths;
465 }
466
467 // used in graph-based dependency parser
468 inline void set_slot_begin(uint32_t idx, uint64_t slot_begin)
469 {
470 // std::cerr << "RnnSequenceClassifier::set_slot_begin slot id=" << (idx+1) << ", slot_begin=" << slot_begin << std::endl;
471 assert(idx < m_num_slots);
472
473 slot_t& slot = m_slots[idx];
474
475 slot.m_output_begin = slot_begin;
476 slot.m_input_begin = slot_begin;
477
478 // Restore this slot's full allocated capacity. set_slot_end() (called right
479 // after, on the DP path) overwrites m_output_end/m_input_end with the actual
480 // data end (count) so the dumper can read [m_output_begin, m_output_end).
481 // But m_output_end is *also* the capacity bound asserted in set_slot_end().
482 // Without restoring it here, reusing this buffer/slot for a later, longer
483 // token batch checks count against the stale (shrunken) end from the
484 // previous use -> assert in debug, silent buffer overrun (heap corruption)
485 // in release. The DP writes [slot_begin, count) with count <= m_slot_len.
486 slot.m_output_end = slot_begin + m_slot_len;
487 slot.m_input_end = slot_begin + m_slot_len;
488 }
489
490 inline void set_slot_end(uint32_t idx, uint64_t slot_end)
491 {
492 // std::cerr << "RnnSequenceClassifier::set_slot_end slot id=" << (idx+1) << ", slot_end=" << slot_end << std::endl;
493 assert(idx < m_num_slots);
494
495 slot_t& slot = m_slots[idx];
496 // std::cerr << "RnnSequenceClassifier::set_slot_end slot output begin=" << slot.m_output_begin
497 // << ", end=" << slot.m_output_end << std::endl;
498
499 assert(slot_end > slot.m_output_begin);
500 assert(slot_end <= slot.m_output_end);
501
502 slot.m_output_end = slot_end;
503 slot.m_input_end = slot_end;
504 }
505
506 inline void start_job(uint32_t idx, bool no_more_data=false)
507 {
508 assert(idx < m_num_slots);
509 slot_t& slot = m_slots[idx];
510 // std::cerr << "RnnSequenceClassifier::start_job " << int(idx+1) << ", " << no_more_data
511 // << "; lock count=" << int(slot.m_lock_count)
512 // << "; flags= " << int(slot.m_flags)
513 // << "; overlap= " << int(m_overlap)
514 // << std::endl;
515 if (no_more_data)
516 {
517 // Must not take into account flags if overlap is zero (parsing)
518 slot.m_flags = m_overlap == 0 ? none : (slot_flags_t)(slot.m_flags & (~right_overlap)) ;
519 // std::cerr << "RnnSequenceClassifier::start_job after reverse" << int(idx+1) << ", " << no_more_data
520 // << "; lock count=" << int(slot.m_lock_count)
521 // << "; flags= " << int(slot.m_flags)
522 // << std::endl;
523 }
524
525 // increment_lock_count(idx);
527 + ((slot.m_flags & left_overlap) > 0 ? 1 : 0)
528 + ((slot.m_flags & right_overlap) > 0 ? 1 : 0));
529
530 if (slot.m_flags & left_overlap)
531 {
532 uint32_t prev_idx = prev_slot(idx);
533 start_job_impl(prev_idx);
534 }
535
536 if (0 == m_overlap || (0 == (slot.m_flags & right_overlap)))
537 {
538 start_job_impl(idx);
539 }
540 }
541
542 inline void wait_for_slot(uint32_t idx)
543 {
544 assert(idx < m_num_slots);
545 const slot_t& slot = m_slots[idx];
546 assert(slot.m_work_started);
547 // std::cerr << "RnnSequenceClassifier::wait_for_slot " << (idx+1) << "/" << m_num_slots
548 // << "; lock count=" << int(slot.m_lock_count) << std::endl;
549 // pretty_print();
550 while (slot.m_lock_count > 1)
551 {
552 // std::cerr << "RnnSequenceClassifier::wait_for_slot in while lock_count=" << int(slot.m_lock_count) << std::endl;
553 // pretty_print();
555 return 1 == slot.m_lock_count;
556 }
557 );
558 }
559 }
560
561 void pretty_print() const
562 {
563 std::cerr << (void*)this << " " << "SLOTS: ";
564 for (size_t i = 0; i < m_num_slots; i++)
565 {
566 std::cerr << " | " << int(m_slots[i].m_lock_count);
567 }
568 std::cerr << " |" << std::endl;
569 }
570
571protected:
572 uint32_t m_overlap;
573 uint32_t m_num_slots;
574 uint32_t m_slot_len;
575
576 std::vector<slot_t> m_slots;
577 std::vector<std::vector<size_t>> m_lengths;
578 std::shared_ptr< StdMatrix<Out> > m_output; // external - classifier id, internal - time position
579
580};
581
582} // namespace deeplima
583
584#endif
Handles multithreading.
std::vector< std::vector< size_t > > m_lengths
bool get_slot_started(uint32_t idx) const
int32_t next_slot(uint32_t idx)
void start_job_impl(uint32_t idx)
Push the slot idx in the thread pool for starting the job on it.
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)
std::vector< slot_t > m_slots
uint8_t get_lock_count(uint32_t idx) const
int32_t prev_slot(uint32_t idx)
void get_classes_from_fn(const std::string &fn, std::vector< std::string > &classes_names, std::vector< std::vector< std::string > > &classes)
std::shared_ptr< StdMatrix< Out > > get_output()
void set_slot_begin(uint32_t idx, uint64_t slot_begin)
void set_slot_end(uint32_t idx, uint64_t slot_end)
void increment_lock_count(uint32_t idx, uint8_t v=1)
static void run_one_job(ThisClass *this_ptr, size_t worker_id, void *p)
void decrement_lock_count(uint32_t idx)
int32_t get_slot_idx(uint64_t timepoint) const
uint64_t get_slot_begin(uint32_t idx) const
RnnSequenceClassifier< Model, InputVectorizer, Out > ThisClass
void start_job(uint32_t idx, bool no_more_data=false)
uint8_t get_output(uint64_t pos, uint8_t cls)
void set_slot_lengths(uint32_t idx, const std::vector< size_t > &lengths)
ThreadPool< ThisClass > RnnSequenceClassifierThreadPool
virtual void reset()
Need to be called to be able to reuse this classifier on several sequences.
std::shared_ptr< StdMatrix< Out > > m_output
void init(size_t num_threads)
Definition thread_pool.h:35
void push(void *job)
Definition thread_pool.h:94
void wait_for_any_job_notification(const std::function< bool()> fn)
std::ostream & operator<<(std::ostream &out, const std::vector< T > &v)
STL namespace.
slot_t represents the part of the job processed by a classifier thread (tokenizer or tagger).
slot_t & operator=(const slot_t &s)