LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
char_seq_vectorizer.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_INCLUDE_SEGMENTATION_TRAIN_CHAR_SEQ_VECTORIZER_H
7#define DEEPLIMA_SRC_INCLUDE_SEGMENTATION_TRAIN_CHAR_SEQ_VECTORIZER_H
8
9#include <vector>
10#include <string>
11
13
14namespace deeplima
15{
16namespace segmentation
17{
18namespace train
19{
20
21template <class InputEncoder, class Matrix, class Adapter>
22class CharSeqVectorizerImpl : public Adapter
23{
24public:
25 CharSeqVectorizerImpl(const std::vector<impl::ngram_descr_t>& ngram_descr)
26 : Adapter(),
27 m_input_encoder(ngram_descr)
28 {
29 }
30
31 std::shared_ptr<Matrix> process(const std::string& text, int64_t len = -1)
32 {
33 assert(len > 0); // TODO: calculate length of text in characters in case len == -1
34 return read_string(text, len);
35 }
36
37protected:
38 std::shared_ptr<Matrix> read_string(const std::string& text, uint64_t len)
39 {
40 auto target = std::make_shared<Matrix>(len, m_input_encoder.size());
41
42 uint64_t current_timepoint = 0;
43 int32_t pos = 0;
44 m_input_encoder.reset();
45
46 while (! m_input_encoder.ready_to_generate())
47 {
48 m_input_encoder.warmup((const uint8_t*)text.data(), &pos, text.size());
49 }
50
51 while (size_t(pos) < text.size())
52 {
53 if (m_input_encoder.parse((const uint8_t*)text.data(), &pos, text.size()) > 0)
54 {
55 handle_timepoint(*target, current_timepoint);
56 }
57 }
58
59 char final_spaces[] = " ";
60 for (size_t i = 0; i < m_input_encoder.get_lookahead(); i++)
61 {
62 int32_t pos = 0;
63 if (m_input_encoder.parse((uint8_t*)final_spaces, &pos, 1) > 0)
64 {
65 handle_timepoint(*target, current_timepoint);
66 }
67 else
68 {
69 throw std::runtime_error("Something wrong.");
70 }
71 }
72
73 return target;
74 }
75
76 inline void handle_timepoint(Matrix& target, uint64_t& current_timepoint)
77 {
78 for (size_t i = 0; i < m_input_encoder.size(); i++)
79 {
80 uint64_t v = m_input_encoder.get_feat(i);
81 Adapter::set(target, current_timepoint, i, v);
82 }
83 current_timepoint++;
84 if (current_timepoint == std::numeric_limits<uint64_t>::max())
85 {
86 throw std::overflow_error("Too much characters in training set.");
87 }
88 }
89
90 InputEncoder m_input_encoder;
91};
92
93} // namespace train
94} // namespace segmentation
95} // namespace deeplima
96
97#endif
std::shared_ptr< Matrix > read_string(const std::string &text, uint64_t len)
std::shared_ptr< Matrix > process(const std::string &text, int64_t len=-1)
void handle_timepoint(Matrix &target, uint64_t &current_timepoint)
CharSeqVectorizerImpl(const std::vector< impl::ngram_descr_t > &ngram_descr)