LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
char_dict_builder.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 SRC_INCLUDE_SEGMENTATION_TRAIN_CHAR_DICT_BUILDER_H
7#define SRC_INCLUDE_SEGMENTATION_TRAIN_CHAR_DICT_BUILDER_H
8
9#include <string>
10#include <vector>
11#include <unordered_map>
12#include <memory>
13
15#include "static_graph/dict.h"
17
18namespace deeplima
19{
20namespace segmentation
21{
22namespace train
23{
24
25template <class InputEncoder>
27{
28public:
29 DictionaryBuilderImpl(const std::vector<impl::ngram_descr_t>& ngram_descr)
30 : m_input_encoder(ngram_descr),
31 m_unk(~0)
32 {
33 assert(sizeof(uint64_t) >= sizeof(m_input_encoder.get_feat(0)));
34 }
35
36 inline uint64_t get_num_features() const
37 {
38 assert(m_input_encoder.size() < std::numeric_limits<int64_t>::max());
39 return m_input_encoder.size();
40 }
41
42 DictsHolder process(const std::string& text,
43 uint32_t min_ipm,
44 uint64_t& char_counter)
45 {
46 std::vector<std::unordered_map<uint64_t, uint64_t>> temp_dicts;
47 temp_dicts.resize(m_input_encoder.size());
48
49 char_counter = 0;
50 read_string(text, char_counter, temp_dicts);
51
52 DictsHolder dicts;
53 dicts.resize(temp_dicts.size());
54 for (size_t i = 0; i < temp_dicts.size(); i++)
55 {
56 if (m_input_encoder.allow_unk(i))
57 {
58 dicts[i] = std::make_shared<UInt64Dict>(m_unk,
59 temp_dicts[i].begin(), temp_dicts[i].end(),
60 [char_counter, min_ipm](uint64_t c){
61 return ipm(c, char_counter) > min_ipm;
62 });
63 }
64 else
65 {
66 dicts[i] = std::make_shared<UInt64Dict>(temp_dicts[i].begin(), temp_dicts[i].end());
67 }
68 }
69
70 // std::cerr << "dicts[0].size() == " << dicts[0]->size() << std::endl;
71 //std::cout << dicts[0].to_string();
72
73 return dicts;
74 }
75
76 void process(const std::wstring& text, uint32_t min_ipm)
77 {
78 std::vector<std::unordered_map<uint64_t, uint64_t>> temp_dicts;
79 temp_dicts.resize(m_input_encoder.size());
80
81 std::unordered_map<uint64_t, uint64_t> char_dict;
82 for (size_t i = 0; i < text.size(); i++)
83 {
84 uint64_t v = (uint64_t)(text[i]);
85 auto it = char_dict.find(v);
86 if (char_dict.end() == it)
87 {
88 char_dict[v] = 1;
89 }
90 else
91 {
92 it->second++;
93 }
94 }
95
96 uint64_t total = text.size();
97 DictsHolder dicts;
98 dicts.resize(temp_dicts.size());
99 dicts[0] = std::make_shared<UInt64Dict>(m_unk, char_dict,
100 [total, min_ipm](uint64_t c){
101 return ipm(c, total) > min_ipm;
102 });
103
104 // std::cerr << "dicts[0].size() == " << dicts[0]->size() << std::endl;
105 }
106
107protected:
108
109 inline static float ipm(uint64_t count, uint64_t total)
110 {
111 return float(count * 1000000) / total;
112 }
113
114 void read_string(const std::string& text,
115 uint64_t& char_counter,
116 std::vector<std::unordered_map<uint64_t, uint64_t>>& temp_dicts)
117 {
118 int32_t pos = 0;
119 char_counter = 0;
120
121 while (! m_input_encoder.ready_to_generate())
122 {
123 m_input_encoder.warmup((const uint8_t*)text.data(), &pos, text.size());
124 }
125
126 while (size_t(pos) < text.size())
127 {
128 if (m_input_encoder.parse((const uint8_t*)text.data(), &pos, text.size()) > 0)
129 {
130 handle_timepoint(char_counter, temp_dicts);
131 }
132 }
133
134 char final_spaces[] = " ";
135 for (size_t i = 0; i < m_input_encoder.get_lookahead(); i++)
136 {
137 int32_t pos = 0;
138 if (m_input_encoder.parse((uint8_t*)final_spaces, &pos, 1) > 0)
139 {
140 handle_timepoint(char_counter, temp_dicts);
141 }
142 else
143 {
144 throw std::runtime_error("Something wrong.");
145 }
146 }
147 }
148
149 inline void handle_timepoint(uint64_t& char_counter,
150 std::vector<std::unordered_map<uint64_t, uint64_t>>& temp_dicts)
151 {
152 for (size_t i = 0; i < m_input_encoder.size(); i++)
153 {
154 uint64_t v = m_input_encoder.get_feat(i);
155 assert(v != m_unk);
156 auto it = temp_dicts[i].find(v);
157 if (temp_dicts[i].end() == it)
158 {
159 temp_dicts[i][v] = 1;
160 }
161 else
162 {
163 it->second++;
164 }
165 }
166
167 if (char_counter == std::numeric_limits<uint64_t>::max())
168 {
169 throw std::overflow_error("Too much characters in training set.");
170 }
171
172 char_counter++;
173 }
174
175 InputEncoder m_input_encoder;
176 uint64_t m_unk;
177};
178
179} // namespace train
180} // namespace segmentation
181} // namespace deeplima
182
183#endif
DictsHolder process(const std::string &text, uint32_t min_ipm, uint64_t &char_counter)
DictionaryBuilderImpl(const std::vector< impl::ngram_descr_t > &ngram_descr)
void process(const std::wstring &text, uint32_t min_ipm)
void handle_timepoint(uint64_t &char_counter, std::vector< std::unordered_map< uint64_t, uint64_t > > &temp_dicts)
static float ipm(uint64_t count, uint64_t total)
void read_string(const std::string &text, uint64_t &char_counter, std::vector< std::unordered_map< uint64_t, uint64_t > > &temp_dicts)