LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
segmentation_decoder.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_APPS_SEGMENTATION_DECODER_H
7#define DEEPLIMA_APPS_SEGMENTATION_DECODER_H
8
9#include <iostream>
10#include <functional>
11#include <stdexcept>
12#include <string>
13#include <vector>
14
15#include <unicode/uchar.h>
16
18#include "deeplima/token_type.h"
19
20namespace deeplima
21{
22namespace segmentation
23{
24
26{
27 uint16_t m_offset; // offset from previous token end
28 uint16_t m_len; // length of this token in bytes
29 const char* m_pch;
31
32 // Multiword-token (MWT) expansion metadata, set by MwtExpander. On the first
33 // sub-word of an expanded surface token (e.g. "de" of "du"->de+le) m_mwt_len
34 // holds the number of sub-words and m_mwt_surface_* the surface form bytes;
35 // both are zero/null on normal tokens and on the non-first sub-words.
36 uint8_t m_mwt_len;
37 const char* m_mwt_surface_pch;
39
43
44 inline void clear()
45 {
46 m_offset = m_len = 0;
47 m_pch = nullptr;
49 m_mwt_len = 0;
50 m_mwt_surface_pch = nullptr;
52 }
53
54 inline bool empty() const
55 {
56 return (0 == m_offset) && (0 == m_len) && (nullptr == m_pch);
57 }
58
59 inline bool too_long() const
60 {
61 return m_len & (1 << (sizeof(uint16_t) * 8 - 1));
62 }
63};
64
65typedef std::function < void (const std::vector<token_pos>& tokens,
66 uint32_t len) > segmentation_callback_t;
67
68enum segm_tag_t : uint8_t
69{
70 // tokenization-only tags
71 X = 0x00,
72 B = 0x01,
73 I = 0x02,
74 E = 0x03,
75 S = 0x04,
77
78 // + sentence segmentation tags
79 E_EOS = 0x05,
80 S_EOS = 0x06,
82
83 // + multiword-token tags: marks a surface token to be expanded into sub-words
84 // (e.g. French "du"->de+le). Appended after the base tags so models trained
85 // without MWT (<= max_segm_tag classes) keep identical tag values. Each is the
86 // base end-tag value + max_tok_tag (E_MWT=E+max_tok_tag, etc.).
87 E_MWT = 0x07,
88 S_MWT = 0x08,
89 E_EOS_MWT = 0x09,
90 S_EOS_MWT = 0x0A,
91 max_segm_mwt_tag = 0x0B
92};
93
94namespace impl
95{
96
97template <uint8_t M = 4>
99{
100public:
101
103 : m_buff_len(0)
104 {
105 memset(m_buff, 0, M);
106 }
107
108 inline bool read(const char** pch, const char* end, uint8_t len)
109 {
110 assert(*pch < end);
111 assert(len <= M);
112 assert(m_buff_len <= M);
113
114 size_t available_bytes = end - *pch;
115
116 if (available_bytes < len)
117 {
118 assert(m_buff_len + available_bytes <= M);
119 memcpy(m_buff + m_buff_len, *pch, available_bytes);
120 m_buff_len += available_bytes;
121 *pch += available_bytes;
122 assert(m_buff_len <= M);
123 return false;
124 }
125
126 return true;
127 }
128
129 inline bool not_empty() const
130 {
131 return m_buff_len > 0;
132 }
133
134 inline bool start_reading(const char** pch,
135 const char* end,
136 uint8_t len, // length of the current character in bytes
137 std::vector<uint8_t>& dst, // destination buffer
138 token_pos& tmp,
139 bool inside)
140 {
141 assert(*pch < end);
142 assert(len <= M);
143 assert(m_buff_len <= M);
144
145 if (m_buff_len + (end - *pch) < len)
146 {
147 // full character isn't still available
148 return read(pch, end, len);
149 }
150
151 // TODO: Don't trust ML model.
152 // This code skips non-space characters
153 // in case of ML errors (wrong X tag).
154 if (inside)
155 {
156 if (dst.size() < size_t(tmp.m_len + len + 1))
157 {
158 dst.resize(tmp.m_len + len + 1);
159 }
160
161 memcpy(dst.data() + tmp.m_len, m_buff, m_buff_len);
162 memcpy(dst.data() + tmp.m_len + m_buff_len, *pch, len - m_buff_len);
163
164 tmp.m_len += len;
165 tmp.m_pch = (const char*)dst.data();
166 }
167 else
168 {
169 tmp.m_offset += len;
170 }
171
172 m_buff_len = 0;
173
174 return true;
175 }
176
177protected:
178 char m_buff[M];
179 uint8_t m_buff_len;
180};
181
183{
184public:
185
186 SegmentationDecoder(std::shared_ptr< StdMatrix<uint8_t> > out, const std::vector<uint8_t>& len)
187 : m_out(out),
188 m_len(len)
189 {
190 init();
191 }
192
193 void init()
194 {
195 m_tokens.resize(1024);
196 }
197
199 {
200 m_callback = fn;
201 }
202
203 inline void save_current_token(size_t& pos, size_t& temp_token_len, const char* start)
204 {
205 if (m_tokens[pos].m_len > 0)
206 {
207 if (m_tokens[pos].m_pch == (const char*)m_temp_text.data())
208 {
209 assert(0 == pos);
210 if (m_temp_text.size() < size_t(m_tokens[pos].m_len + 1))
211 {
212 m_temp_text.resize(m_tokens[pos].m_len + 1);
213 }
214 memcpy(m_temp_text.data() + temp_token_len, start, m_tokens[pos].m_len - temp_token_len);
215 m_tokens[pos].m_pch = (const char*)(m_temp_text.data());
216 }
217 pos++;
218 if (pos >= m_tokens.size())
219 {
220 m_tokens.resize(m_tokens.size() + 1024);
221 }
222 m_tokens[pos].clear();
223 }
224 }
225
226 inline void consome_character(size_t& pos, uint64_t from, const char* pch)
227 {
228 if (nullptr == m_tokens[pos].m_pch)
229 {
230 assert(0 == m_tokens[pos].m_len);
231 m_tokens[pos].m_pch = pch;
232 }
233 m_tokens[pos].m_len += m_len[from];
234 }
235
236 inline uint64_t decode(const char** pch, uint32_t max, uint64_t from, uint64_t to)
237 {
238 assert(nullptr != pch);
239 assert(nullptr != *pch);
240 const char* start = *pch;
241 const char* end = *pch + max;
242
243 size_t pos = 0;
244 size_t temp_token_len = m_tokens[0].m_len;
245
246 if (not_empty())
247 {
248 uint8_t bytes_stored = m_buff_len;
249 if (start_reading(pch, end, m_len[from], m_temp_text, m_tokens[0], m_out->get(from, 0)))
250 {
251 *pch += m_len[from] - bytes_stored;
252 start += m_len[from] - bytes_stored;
253 if (m_out->get(from, 0) != segm_tag_t::X)
254 {
255 temp_token_len += m_len[from];
256 }
257 from++;
258 }
259 else
260 {
261 // can't read full character (we need more data)
262 return from;
263 }
264 }
265
266 while (from < to && *pch < end)
267 {
268 if (! read(pch, end, m_len[from]))
269 {
270 break;
271 }
272
273//#ifndef NDEBUG
274// std::cerr << "[" << from << "]==" << (int)(m_out->get(from, 0)) << std::endl;
275//#endif
276 int8_t gen_cat = 0;
277 UChar uch;
278 int32_t zero = 0;
279 // TODO: move U8_NEXT from here (we don't handle UTF-8 while decoding)
280 U8_NEXT(*pch, zero, m_len[from], uch);
281 gen_cat = u_charType(uch);
282
283 auto tag = m_out->get(from, 0);
284
285 switch (tag)
286 {
287 case segm_tag_t::X:
288 save_current_token(pos, temp_token_len, start);
289
290 if (gen_cat == U_SPACE_SEPARATOR || gen_cat == U_LINE_SEPARATOR
291 || gen_cat == U_PARAGRAPH_SEPARATOR || gen_cat == U_CONTROL_CHAR
292 || gen_cat == U_FORMAT_CHAR)
293 {
294 m_tokens[pos].m_offset += m_len[from];
295 }
296 else
297 {
298 m_tokens[pos].m_pch = *pch;
299 m_tokens[pos].m_len += m_len[from];
300 save_current_token(pos, temp_token_len, start);
301 }
302 break;
303
304 case segm_tag_t::B:
305 save_current_token(pos, temp_token_len, start);
306
307 if (gen_cat == U_SPACE_SEPARATOR || gen_cat == U_PARAGRAPH_SEPARATOR
308 || gen_cat == U_LINE_SEPARATOR || gen_cat == U_CONTROL_CHAR
309 || gen_cat == U_FORMAT_CHAR)
310 {
311 m_tokens[pos].m_offset += m_len[from];
312 break;
313 }
314 else
315 {
316 assert(0 == m_tokens[pos].m_len);
317 m_tokens[pos].m_pch = *pch;
318 }
319 [[fallthrough]];
320
321 case segm_tag_t::I:
322 if (0 == m_tokens[pos].m_len
323 && (gen_cat == U_SPACE_SEPARATOR || gen_cat == U_PARAGRAPH_SEPARATOR
324 || gen_cat == U_LINE_SEPARATOR || gen_cat == U_CONTROL_CHAR
325 || gen_cat == U_FORMAT_CHAR))
326 {
327 m_tokens[pos].m_offset += m_len[from];
328 }
329 else
330 {
331 consome_character(pos, from, *pch);
332 }
333 break;
334
335 // TODO insert the marker for case continuing [[case_]]
337 m_tokens[pos].m_flags = token_flags_t(m_tokens[pos].m_flags | token_flags_t::sentence_brk);
338 [[fallthrough]];
339
340 case segm_tag_t::E:
341 if (0 == m_tokens[pos].m_len
342 && (gen_cat == U_SPACE_SEPARATOR || gen_cat == U_PARAGRAPH_SEPARATOR
343 || gen_cat == U_LINE_SEPARATOR || gen_cat == U_CONTROL_CHAR
344 || gen_cat == U_FORMAT_CHAR))
345 {
346 m_tokens[pos].m_offset += m_len[from];
347 }
348 else
349 {
350 consome_character(pos, from, *pch);
351 save_current_token(pos, temp_token_len, start);
352 }
353 break;
354
356 save_current_token(pos, temp_token_len, start);
357
358 if (gen_cat == U_SPACE_SEPARATOR || gen_cat == U_PARAGRAPH_SEPARATOR
359 || gen_cat == U_LINE_SEPARATOR || gen_cat == U_CONTROL_CHAR
360 || gen_cat == U_FORMAT_CHAR)
361 {
362 m_tokens[pos].m_offset += m_len[from];
363 }
364 else
365 {
366 assert(0 == m_tokens[pos].m_len);
367 m_tokens[pos].m_pch = *pch;
368 m_tokens[pos].m_len += m_len[from];
369 m_tokens[pos].m_flags = token_flags_t(m_tokens[pos].m_flags | token_flags_t::sentence_brk);
370 save_current_token(pos, temp_token_len, start);
371 }
372 break;
373
374 case segm_tag_t::S:
375 save_current_token(pos, temp_token_len, start);
376
377 if (gen_cat == U_SPACE_SEPARATOR || gen_cat == U_PARAGRAPH_SEPARATOR
378 || gen_cat == U_LINE_SEPARATOR || gen_cat == U_CONTROL_CHAR
379 || gen_cat == U_FORMAT_CHAR)
380 {
381 m_tokens[pos].m_offset += m_len[from];
382 }
383 else
384 {
385 assert(0 == m_tokens[pos].m_len);
386 m_tokens[pos].m_pch = *pch;
387 m_tokens[pos].m_len += m_len[from];
388 save_current_token(pos, temp_token_len, start);
389 }
390 break;
391
392 // Multiword-token end tags: behave exactly like their base tag
393 // (E / E_EOS / S / S_EOS) but additionally mark the finished surface
394 // token with token_flags_t::multiword so the MwtExpander expands it.
396 m_tokens[pos].m_flags = token_flags_t(m_tokens[pos].m_flags | token_flags_t::sentence_brk);
397 [[fallthrough]];
398
400 if (0 == m_tokens[pos].m_len
401 && (gen_cat == U_SPACE_SEPARATOR || gen_cat == U_PARAGRAPH_SEPARATOR
402 || gen_cat == U_LINE_SEPARATOR || gen_cat == U_CONTROL_CHAR
403 || gen_cat == U_FORMAT_CHAR))
404 {
405 m_tokens[pos].m_offset += m_len[from];
406 }
407 else
408 {
409 m_tokens[pos].m_flags = token_flags_t(m_tokens[pos].m_flags | token_flags_t::multiword);
410 consome_character(pos, from, *pch);
411 save_current_token(pos, temp_token_len, start);
412 }
413 break;
414
416 save_current_token(pos, temp_token_len, start);
417
418 if (gen_cat == U_SPACE_SEPARATOR || gen_cat == U_PARAGRAPH_SEPARATOR
419 || gen_cat == U_LINE_SEPARATOR || gen_cat == U_CONTROL_CHAR
420 || gen_cat == U_FORMAT_CHAR)
421 {
422 m_tokens[pos].m_offset += m_len[from];
423 }
424 else
425 {
426 assert(0 == m_tokens[pos].m_len);
427 m_tokens[pos].m_pch = *pch;
428 m_tokens[pos].m_len += m_len[from];
429 m_tokens[pos].m_flags = token_flags_t(m_tokens[pos].m_flags
432 save_current_token(pos, temp_token_len, start);
433 }
434 break;
435
437 save_current_token(pos, temp_token_len, start);
438
439 if (gen_cat == U_SPACE_SEPARATOR || gen_cat == U_PARAGRAPH_SEPARATOR
440 || gen_cat == U_LINE_SEPARATOR || gen_cat == U_CONTROL_CHAR
441 || gen_cat == U_FORMAT_CHAR)
442 {
443 m_tokens[pos].m_offset += m_len[from];
444 }
445 else
446 {
447 assert(0 == m_tokens[pos].m_len);
448 m_tokens[pos].m_pch = *pch;
449 m_tokens[pos].m_len += m_len[from];
450 m_tokens[pos].m_flags = token_flags_t(m_tokens[pos].m_flags | token_flags_t::multiword);
451 save_current_token(pos, temp_token_len, start);
452 }
453 break;
454
455 default:
456 throw std::runtime_error("Unknown code in output.");
457 }
458
459 if (m_tokens[pos].too_long())
460 {
461 // This is a workaround to handle garbage in the input data
462 // very long (and meaningless) tokens are artificially splitted
463 // into several parts.
464 // TODO: the same type of handling is required for very long
465 // sequence of spaces: an empty token must be generated
466 // to avoid overflow of token_pos::m_offset field.
467 save_current_token(pos, temp_token_len, start);
468 }
469
470 *pch += m_len[from]; // next byte
471 from++; // next character
472 }
473
474 if (pos > 0)
475 {
476 m_callback(m_tokens, pos);
477 }
478 else
479 {
480 if (m_tokens[0].m_pch == (const char*)m_temp_text.data())
481 {
482 if (m_temp_text.size() < size_t(m_tokens[0].m_len + 1))
483 {
484 m_temp_text.resize(m_tokens[0].m_len + 1);
485 }
486 assert(0 == pos);
487 assert(m_tokens[pos].m_len >= temp_token_len);
488 memcpy(m_temp_text.data() + temp_token_len, start, m_tokens[pos].m_len - temp_token_len);
489 m_tokens[pos].m_pch = (const char*)(m_temp_text.data());
490 }
491 // std::cerr << std::endl;
492 }
493
494 if (!m_tokens[pos].empty())
495 {
496 if (pos > 0)
497 {
498 m_tokens[0] = m_tokens[pos];
499 }
500 if (m_temp_text.size() < size_t(m_tokens[0].m_len + 1))
501 {
502 m_temp_text.resize(m_tokens[0].m_len + 1);
503 }
504 if ((const char*)m_temp_text.data() != m_tokens[0].m_pch)
505 {
506 if (m_tokens[0].m_pch != nullptr)
507 {
508 memcpy(m_temp_text.data(), m_tokens[0].m_pch, m_tokens[0].m_len);
509 }
510 m_tokens[0].m_pch = (const char*)(m_temp_text.data());
511 }
512 }
513 else
514 {
515 m_tokens[0].clear();
516 }
517
518 return from;
519 }
520
521protected:
522 // input
523 std::shared_ptr< StdMatrix<uint8_t> > m_out;
524 const std::vector<uint8_t>& m_len;
525
526 // output
527 std::vector<token_pos> m_tokens;
528
529 // callback
531
532 // temp buffers
533 std::vector<uint8_t> m_temp_text;
534};
535
536} // namespace impl
537} // namespace segmentation
538} // namespace deeplima
539
540#endif
bool start_reading(const char **pch, const char *end, uint8_t len, std::vector< uint8_t > &dst, token_pos &tmp, bool inside)
bool read(const char **pch, const char *end, uint8_t len)
void consome_character(size_t &pos, uint64_t from, const char *pch)
SegmentationDecoder(std::shared_ptr< StdMatrix< uint8_t > > out, const std::vector< uint8_t > &len)
uint64_t decode(const char **pch, uint32_t max, uint64_t from, uint64_t to)
void save_current_token(size_t &pos, size_t &temp_token_len, const char *start)
std::shared_ptr< StdMatrix< uint8_t > > m_out
void register_handler(const segmentation_callback_t fn)
std::function< void(const std::vector< token_pos > &tokens, uint32_t len) > segmentation_callback_t
@ sentence_brk
Definition token_type.h:21