static_embeddings 0.1.1
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- checksums.yaml +7 -0
- data/CHANGELOG.md +136 -0
- data/LICENSE.txt +21 -0
- data/README.md +439 -0
- data/docs/ARCHITECTURE.md +278 -0
- data/docs/MODEL_AUDIT.md +61 -0
- data/docs/PERFORMANCE.md +60 -0
- data/exe/static_embeddings +6 -0
- data/ext/static_embeddings/extconf.rb +44 -0
- data/ext/static_embeddings/se_embed.c +192 -0
- data/ext/static_embeddings/se_format.c +762 -0
- data/ext/static_embeddings/se_internal.h +266 -0
- data/ext/static_embeddings/se_tokenizer.c +506 -0
- data/ext/static_embeddings/se_unicode.c +120 -0
- data/ext/static_embeddings/static_embeddings.c +2070 -0
- data/lib/static_embeddings/cli.rb +178 -0
- data/lib/static_embeddings/converter.rb +284 -0
- data/lib/static_embeddings/errors.rb +6 -0
- data/lib/static_embeddings/format.rb +289 -0
- data/lib/static_embeddings/model.rb +48 -0
- data/lib/static_embeddings/paths.rb +29 -0
- data/lib/static_embeddings/reference.rb +191 -0
- data/lib/static_embeddings/safetensors.rb +87 -0
- data/lib/static_embeddings/unicode_tables.rb +127 -0
- data/lib/static_embeddings/version.rb +3 -0
- data/lib/static_embeddings.rb +118 -0
- data/static_embeddings.gemspec +45 -0
- data/tools/benchmark.rb +38 -0
- data/tools/build_demo_model.rb +16 -0
- data/tools/check_model2vec_parity.rb +107 -0
- data/tools/make_fixture_model.rb +165 -0
- metadata +134 -0
|
@@ -0,0 +1,506 @@
|
|
|
1
|
+
#include "se_internal.h"
|
|
2
|
+
|
|
3
|
+
#include <stdint.h>
|
|
4
|
+
#include <stdlib.h>
|
|
5
|
+
#include <string.h>
|
|
6
|
+
|
|
7
|
+
#define SE_SIZE_MAX ((size_t)-1)
|
|
8
|
+
#define SE_CANCEL_CHECK_MASK 0x3ffu
|
|
9
|
+
|
|
10
|
+
void se_scratch_init(se_scratch_t *s) {
|
|
11
|
+
memset(s, 0, sizeof(*s));
|
|
12
|
+
}
|
|
13
|
+
|
|
14
|
+
void se_scratch_free(se_scratch_t *s) {
|
|
15
|
+
free(s->cps);
|
|
16
|
+
free(s->cps2);
|
|
17
|
+
free(s->bytes);
|
|
18
|
+
free(s->ids);
|
|
19
|
+
free(s->acc);
|
|
20
|
+
memset(s, 0, sizeof(*s));
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
static int grow_u32(uint32_t **buf, size_t *cap, size_t need) {
|
|
24
|
+
if (*cap >= need)
|
|
25
|
+
return 1;
|
|
26
|
+
|
|
27
|
+
size_t next = *cap ? *cap : 256;
|
|
28
|
+
while (next < need) {
|
|
29
|
+
if (next > SE_SIZE_MAX / 2)
|
|
30
|
+
return 0;
|
|
31
|
+
next *= 2;
|
|
32
|
+
}
|
|
33
|
+
|
|
34
|
+
if (next > SE_SIZE_MAX / sizeof(uint32_t))
|
|
35
|
+
return 0;
|
|
36
|
+
|
|
37
|
+
uint32_t *p = (uint32_t *)realloc(*buf, next * sizeof(uint32_t));
|
|
38
|
+
if (!p)
|
|
39
|
+
return 0;
|
|
40
|
+
|
|
41
|
+
*buf = p;
|
|
42
|
+
*cap = next;
|
|
43
|
+
return 1;
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
static int grow_bytes(uint8_t **buf, size_t *cap, size_t need) {
|
|
47
|
+
if (*cap >= need)
|
|
48
|
+
return 1;
|
|
49
|
+
|
|
50
|
+
size_t next = *cap ? *cap : 256;
|
|
51
|
+
while (next < need) {
|
|
52
|
+
if (next > SE_SIZE_MAX / 2)
|
|
53
|
+
return 0;
|
|
54
|
+
next *= 2;
|
|
55
|
+
}
|
|
56
|
+
|
|
57
|
+
uint8_t *p = (uint8_t *)realloc(*buf, next);
|
|
58
|
+
if (!p)
|
|
59
|
+
return 0;
|
|
60
|
+
|
|
61
|
+
*buf = p;
|
|
62
|
+
*cap = next;
|
|
63
|
+
return 1;
|
|
64
|
+
}
|
|
65
|
+
|
|
66
|
+
static int grow_float(float **buf, size_t *cap, size_t need) {
|
|
67
|
+
if (*cap >= need)
|
|
68
|
+
return 1;
|
|
69
|
+
|
|
70
|
+
if (need > SE_SIZE_MAX / sizeof(float))
|
|
71
|
+
return 0;
|
|
72
|
+
|
|
73
|
+
float *p = (float *)realloc(*buf, need * sizeof(float));
|
|
74
|
+
if (!p)
|
|
75
|
+
return 0;
|
|
76
|
+
|
|
77
|
+
*buf = p;
|
|
78
|
+
*cap = need;
|
|
79
|
+
return 1;
|
|
80
|
+
}
|
|
81
|
+
|
|
82
|
+
static int push_id(se_scratch_t *sc, size_t *n_ids, uint32_t id) {
|
|
83
|
+
if (*n_ids >= UINT32_MAX)
|
|
84
|
+
return 0;
|
|
85
|
+
if (!grow_u32(&sc->ids, &sc->ids_cap, *n_ids + 1))
|
|
86
|
+
return 0;
|
|
87
|
+
sc->ids[(*n_ids)++] = id;
|
|
88
|
+
return 1;
|
|
89
|
+
}
|
|
90
|
+
|
|
91
|
+
int se_scratch_reserve(se_scratch_t *s, uint32_t dim) {
|
|
92
|
+
if (!grow_u32(&s->cps, &s->cps_cap, 256))
|
|
93
|
+
return 0;
|
|
94
|
+
if (!grow_u32(&s->cps2, &s->cps2_cap, 256))
|
|
95
|
+
return 0;
|
|
96
|
+
if (!grow_u32(&s->ids, &s->ids_cap, 256))
|
|
97
|
+
return 0;
|
|
98
|
+
if (!grow_bytes(&s->bytes, &s->bytes_cap, 256))
|
|
99
|
+
return 0;
|
|
100
|
+
if (!grow_float(&s->acc, &s->acc_cap, dim))
|
|
101
|
+
return 0;
|
|
102
|
+
|
|
103
|
+
return 1;
|
|
104
|
+
}
|
|
105
|
+
|
|
106
|
+
static int is_whitespace(const se_model_t *m, uint32_t cp) {
|
|
107
|
+
if (cp < 0x80)
|
|
108
|
+
return se_is_ascii_whitespace(cp);
|
|
109
|
+
return se_range_contains(m->norm.whitespace, m->norm.whitespace_count, cp);
|
|
110
|
+
}
|
|
111
|
+
|
|
112
|
+
static int is_control(const se_model_t *m, uint32_t cp) {
|
|
113
|
+
if (cp < 0x80)
|
|
114
|
+
return cp < 0x20 && cp != '\t' && cp != '\n' && cp != '\r';
|
|
115
|
+
return se_range_contains(m->norm.control, m->norm.control_count, cp);
|
|
116
|
+
}
|
|
117
|
+
|
|
118
|
+
static int is_punct(const se_model_t *m, uint32_t cp) {
|
|
119
|
+
if (cp < 0x80)
|
|
120
|
+
return se_is_ascii_punct(cp);
|
|
121
|
+
return se_range_contains(m->norm.punct, m->norm.punct_count, cp);
|
|
122
|
+
}
|
|
123
|
+
|
|
124
|
+
static int decode_one(const uint8_t *src, size_t len, size_t *i, uint32_t *cp_out) {
|
|
125
|
+
uint8_t b0 = src[*i];
|
|
126
|
+
uint32_t cp;
|
|
127
|
+
size_t need;
|
|
128
|
+
|
|
129
|
+
if (b0 < 0x80) {
|
|
130
|
+
*cp_out = b0;
|
|
131
|
+
(*i)++;
|
|
132
|
+
return 1;
|
|
133
|
+
}
|
|
134
|
+
if ((b0 & 0xE0) == 0xC0) {
|
|
135
|
+
cp = b0 & 0x1Fu;
|
|
136
|
+
need = 1;
|
|
137
|
+
} else if ((b0 & 0xF0) == 0xE0) {
|
|
138
|
+
cp = b0 & 0x0Fu;
|
|
139
|
+
need = 2;
|
|
140
|
+
} else if ((b0 & 0xF8) == 0xF0) {
|
|
141
|
+
cp = b0 & 0x07u;
|
|
142
|
+
need = 3;
|
|
143
|
+
} else {
|
|
144
|
+
return 0;
|
|
145
|
+
}
|
|
146
|
+
|
|
147
|
+
if (*i + need >= len)
|
|
148
|
+
return 0;
|
|
149
|
+
|
|
150
|
+
for (size_t k = 1; k <= need; k++) {
|
|
151
|
+
uint8_t bk = src[*i + k];
|
|
152
|
+
if ((bk & 0xC0) != 0x80)
|
|
153
|
+
return 0;
|
|
154
|
+
cp = (cp << 6) | (uint32_t)(bk & 0x3Fu);
|
|
155
|
+
}
|
|
156
|
+
|
|
157
|
+
if ((need == 1 && cp < 0x80) || (need == 2 && cp < 0x800) || (need == 3 && cp < 0x10000))
|
|
158
|
+
return 0;
|
|
159
|
+
if (cp > 0x10FFFF || (cp >= 0xD800 && cp <= 0xDFFF))
|
|
160
|
+
return 0;
|
|
161
|
+
|
|
162
|
+
*cp_out = cp;
|
|
163
|
+
*i += need + 1;
|
|
164
|
+
return 1;
|
|
165
|
+
}
|
|
166
|
+
|
|
167
|
+
static int normalization_stable(const se_model_t *m, uint32_t cp) {
|
|
168
|
+
if (cp < 0x80)
|
|
169
|
+
return 1;
|
|
170
|
+
if (m->meta.do_lower_case && se_map_lookup(m->norm.lower, m->norm.lower_count, cp))
|
|
171
|
+
return 0;
|
|
172
|
+
if (m->meta.strip_accents) {
|
|
173
|
+
if (se_map_lookup(m->norm.nfd, m->norm.nfd_count, cp))
|
|
174
|
+
return 0;
|
|
175
|
+
if (se_range_contains(m->norm.mn, m->norm.mn_count, cp))
|
|
176
|
+
return 0;
|
|
177
|
+
}
|
|
178
|
+
return 1;
|
|
179
|
+
}
|
|
180
|
+
|
|
181
|
+
static int cp_is_cjk_segment(const se_model_t *m, uint32_t cp) {
|
|
182
|
+
return m->meta.handle_chinese_chars && cp >= 0x3400 && se_is_cjk(cp);
|
|
183
|
+
}
|
|
184
|
+
|
|
185
|
+
size_t se_prefix_boundary_len(const se_model_t *model, const uint8_t *input, size_t input_len,
|
|
186
|
+
size_t target, size_t backscan) {
|
|
187
|
+
if (target >= input_len)
|
|
188
|
+
return input_len;
|
|
189
|
+
if (target == 0)
|
|
190
|
+
return 0;
|
|
191
|
+
|
|
192
|
+
size_t floor = target > backscan ? target - backscan : 0;
|
|
193
|
+
size_t pos = target;
|
|
194
|
+
|
|
195
|
+
while (pos > floor) {
|
|
196
|
+
pos--;
|
|
197
|
+
while (pos > floor && (input[pos] & 0xC0) == 0x80)
|
|
198
|
+
pos--;
|
|
199
|
+
if ((input[pos] & 0xC0) == 0x80)
|
|
200
|
+
return 0;
|
|
201
|
+
|
|
202
|
+
uint32_t cp = 0;
|
|
203
|
+
size_t scan = pos;
|
|
204
|
+
if (!decode_one(input, input_len, &scan, &cp))
|
|
205
|
+
return 0;
|
|
206
|
+
size_t after = scan;
|
|
207
|
+
|
|
208
|
+
if (!normalization_stable(model, cp))
|
|
209
|
+
continue;
|
|
210
|
+
|
|
211
|
+
if (cp_is_cjk_segment(model, cp)) {
|
|
212
|
+
if (after <= target)
|
|
213
|
+
return after;
|
|
214
|
+
if (pos > 0)
|
|
215
|
+
return pos;
|
|
216
|
+
return 0;
|
|
217
|
+
}
|
|
218
|
+
|
|
219
|
+
if ((is_whitespace(model, cp) || is_punct(model, cp)) && after <= target)
|
|
220
|
+
return after;
|
|
221
|
+
}
|
|
222
|
+
|
|
223
|
+
return 0;
|
|
224
|
+
}
|
|
225
|
+
|
|
226
|
+
typedef struct {
|
|
227
|
+
const se_model_t *model;
|
|
228
|
+
se_scratch_t *scratch;
|
|
229
|
+
uint32_t max_tokens;
|
|
230
|
+
size_t n_ids;
|
|
231
|
+
size_t n_unk;
|
|
232
|
+
size_t segment_len;
|
|
233
|
+
se_token_stats_t *stats;
|
|
234
|
+
se_error_t *err;
|
|
235
|
+
volatile sig_atomic_t *cancelled;
|
|
236
|
+
} token_state_t;
|
|
237
|
+
|
|
238
|
+
static int token_cancelled(const token_state_t *st) {
|
|
239
|
+
return st->cancelled && *st->cancelled;
|
|
240
|
+
}
|
|
241
|
+
|
|
242
|
+
static se_status_t oom(token_state_t *st, const char *where) {
|
|
243
|
+
se_error_set(st->err, SE_ERR_OOM, "out of memory while %s", where);
|
|
244
|
+
return SE_ERR_OOM;
|
|
245
|
+
}
|
|
246
|
+
|
|
247
|
+
static int append_segment_cp(token_state_t *st, uint32_t cp) {
|
|
248
|
+
if (!grow_u32(&st->scratch->cps, &st->scratch->cps_cap, st->segment_len + 1))
|
|
249
|
+
return 0;
|
|
250
|
+
st->scratch->cps[st->segment_len++] = cp;
|
|
251
|
+
return 1;
|
|
252
|
+
}
|
|
253
|
+
|
|
254
|
+
static int cap_after_append(token_state_t *st) {
|
|
255
|
+
if (st->max_tokens == 0 || st->n_ids <= (size_t)st->max_tokens)
|
|
256
|
+
return 0;
|
|
257
|
+
|
|
258
|
+
size_t dropped_unk = 0;
|
|
259
|
+
for (size_t k = st->max_tokens; k < st->n_ids; k++) {
|
|
260
|
+
if (st->scratch->ids[k] == st->model->meta.unk_id)
|
|
261
|
+
dropped_unk++;
|
|
262
|
+
}
|
|
263
|
+
st->n_unk -= dropped_unk;
|
|
264
|
+
st->n_ids = st->max_tokens;
|
|
265
|
+
st->stats->truncated = 1;
|
|
266
|
+
return 1;
|
|
267
|
+
}
|
|
268
|
+
|
|
269
|
+
typedef enum { WORDPIECE_OK = 1, WORDPIECE_OOM = 0, WORDPIECE_INVALID = -1 } wordpiece_status_t;
|
|
270
|
+
|
|
271
|
+
static wordpiece_status_t wordpiece(const se_model_t *m, se_scratch_t *sc, const uint32_t *word,
|
|
272
|
+
size_t word_len, size_t *n_ids, size_t *n_unk) {
|
|
273
|
+
const se_meta_t *meta = &m->meta;
|
|
274
|
+
|
|
275
|
+
if (word_len == 0)
|
|
276
|
+
return 1;
|
|
277
|
+
|
|
278
|
+
if (word_len > meta->max_input_chars_per_word) {
|
|
279
|
+
if (!push_id(sc, n_ids, meta->unk_id))
|
|
280
|
+
return 0;
|
|
281
|
+
(*n_unk)++;
|
|
282
|
+
return 1;
|
|
283
|
+
}
|
|
284
|
+
|
|
285
|
+
if (word_len > (SE_SIZE_MAX / 4) - 1)
|
|
286
|
+
return 0;
|
|
287
|
+
if (!grow_bytes(&sc->bytes, &sc->bytes_cap, word_len * 4 + 4))
|
|
288
|
+
return 0;
|
|
289
|
+
if (!grow_u32(&sc->cps2, &sc->cps2_cap, word_len + 1))
|
|
290
|
+
return 0;
|
|
291
|
+
|
|
292
|
+
size_t blen = 0;
|
|
293
|
+
for (size_t k = 0; k < word_len; k++) {
|
|
294
|
+
if (blen > UINT32_MAX)
|
|
295
|
+
return 0;
|
|
296
|
+
sc->cps2[k] = (uint32_t)blen;
|
|
297
|
+
uint32_t cp = word[k];
|
|
298
|
+
if (cp < 0x80) {
|
|
299
|
+
sc->bytes[blen++] = (uint8_t)cp;
|
|
300
|
+
} else {
|
|
301
|
+
blen += se_utf8_encode(cp, sc->bytes + blen);
|
|
302
|
+
}
|
|
303
|
+
}
|
|
304
|
+
if (blen > UINT32_MAX)
|
|
305
|
+
return 0;
|
|
306
|
+
sc->cps2[word_len] = (uint32_t)blen;
|
|
307
|
+
|
|
308
|
+
if (word_len <= meta->max_token_chars) {
|
|
309
|
+
uint32_t exact_id = 0;
|
|
310
|
+
if (se_vocab_lookup(m, sc->bytes, blen, &exact_id))
|
|
311
|
+
return push_id(sc, n_ids, exact_id);
|
|
312
|
+
}
|
|
313
|
+
|
|
314
|
+
size_t start = 0;
|
|
315
|
+
size_t emitted_before = *n_ids;
|
|
316
|
+
|
|
317
|
+
while (start < word_len) {
|
|
318
|
+
size_t max_end = start + meta->max_token_chars;
|
|
319
|
+
if (max_end > word_len)
|
|
320
|
+
max_end = word_len;
|
|
321
|
+
|
|
322
|
+
size_t off = sc->cps2[start];
|
|
323
|
+
size_t max_off = sc->cps2[max_end];
|
|
324
|
+
size_t matched_len = 0;
|
|
325
|
+
uint32_t found_id = 0;
|
|
326
|
+
const se_trie_t *trie = start > 0 ? &m->cont_trie : &m->root_trie;
|
|
327
|
+
|
|
328
|
+
int found =
|
|
329
|
+
se_trie_longest_match(trie, sc->bytes + off, max_off - off, &found_id, &matched_len);
|
|
330
|
+
|
|
331
|
+
if (!found) {
|
|
332
|
+
*n_ids = emitted_before;
|
|
333
|
+
if (!push_id(sc, n_ids, meta->unk_id))
|
|
334
|
+
return 0;
|
|
335
|
+
(*n_unk)++;
|
|
336
|
+
return 1;
|
|
337
|
+
}
|
|
338
|
+
|
|
339
|
+
size_t target = off + matched_len;
|
|
340
|
+
size_t end = start + 1;
|
|
341
|
+
while (end <= max_end && sc->cps2[end] < target)
|
|
342
|
+
end++;
|
|
343
|
+
if (end > max_end || sc->cps2[end] != target)
|
|
344
|
+
return WORDPIECE_INVALID;
|
|
345
|
+
|
|
346
|
+
if (!push_id(sc, n_ids, found_id))
|
|
347
|
+
return 0;
|
|
348
|
+
start = end;
|
|
349
|
+
}
|
|
350
|
+
return 1;
|
|
351
|
+
}
|
|
352
|
+
|
|
353
|
+
static se_status_t append_wordpiece(token_state_t *st, const uint32_t *word, size_t word_len,
|
|
354
|
+
int *stop) {
|
|
355
|
+
wordpiece_status_t wp =
|
|
356
|
+
wordpiece(st->model, st->scratch, word, word_len, &st->n_ids, &st->n_unk);
|
|
357
|
+
if (wp == WORDPIECE_OOM)
|
|
358
|
+
return oom(st, "tokenizing");
|
|
359
|
+
if (wp == WORDPIECE_INVALID) {
|
|
360
|
+
se_error_set(st->err, SE_ERR_INTERNAL, "invalid WordPiece byte boundary");
|
|
361
|
+
return SE_ERR_INTERNAL;
|
|
362
|
+
}
|
|
363
|
+
if (cap_after_append(st))
|
|
364
|
+
*stop = 1;
|
|
365
|
+
return SE_OK;
|
|
366
|
+
}
|
|
367
|
+
|
|
368
|
+
static se_status_t flush_segment(token_state_t *st, int *stop) {
|
|
369
|
+
if (st->segment_len == 0)
|
|
370
|
+
return SE_OK;
|
|
371
|
+
se_status_t rc = append_wordpiece(st, st->scratch->cps, st->segment_len, stop);
|
|
372
|
+
st->segment_len = 0;
|
|
373
|
+
return rc;
|
|
374
|
+
}
|
|
375
|
+
|
|
376
|
+
static se_status_t feed_token_cp(token_state_t *st, uint32_t cp, int *stop) {
|
|
377
|
+
if (is_whitespace(st->model, cp))
|
|
378
|
+
return flush_segment(st, stop);
|
|
379
|
+
|
|
380
|
+
if (is_punct(st->model, cp)) {
|
|
381
|
+
se_status_t rc = flush_segment(st, stop);
|
|
382
|
+
if (rc != SE_OK || *stop)
|
|
383
|
+
return rc;
|
|
384
|
+
return append_wordpiece(st, &cp, 1, stop);
|
|
385
|
+
}
|
|
386
|
+
|
|
387
|
+
if (!append_segment_cp(st, cp))
|
|
388
|
+
return oom(st, "building token segment");
|
|
389
|
+
return SE_OK;
|
|
390
|
+
}
|
|
391
|
+
|
|
392
|
+
static se_status_t emit_lowered(token_state_t *st, uint32_t cp, int *stop) {
|
|
393
|
+
if (st->model->meta.do_lower_case) {
|
|
394
|
+
if (cp >= 'A' && cp <= 'Z')
|
|
395
|
+
cp += 32;
|
|
396
|
+
else if (cp >= 0x80) {
|
|
397
|
+
const se_map_entry_t *e =
|
|
398
|
+
se_map_lookup(st->model->norm.lower, st->model->norm.lower_count, cp);
|
|
399
|
+
if (e) {
|
|
400
|
+
for (uint32_t k = 0; k < e->len; k++) {
|
|
401
|
+
se_status_t rc = feed_token_cp(st, e->out[k], stop);
|
|
402
|
+
if (rc != SE_OK || *stop)
|
|
403
|
+
return rc;
|
|
404
|
+
}
|
|
405
|
+
return SE_OK;
|
|
406
|
+
}
|
|
407
|
+
}
|
|
408
|
+
}
|
|
409
|
+
return feed_token_cp(st, cp, stop);
|
|
410
|
+
}
|
|
411
|
+
|
|
412
|
+
static se_status_t emit_stripped(token_state_t *st, uint32_t cp, int *stop) {
|
|
413
|
+
if (!st->model->meta.strip_accents || cp < 0x80)
|
|
414
|
+
return emit_lowered(st, cp, stop);
|
|
415
|
+
|
|
416
|
+
const se_map_entry_t *e = se_map_lookup(st->model->norm.nfd, st->model->norm.nfd_count, cp);
|
|
417
|
+
if (!e) {
|
|
418
|
+
if (se_range_contains(st->model->norm.mn, st->model->norm.mn_count, cp))
|
|
419
|
+
return SE_OK;
|
|
420
|
+
return emit_lowered(st, cp, stop);
|
|
421
|
+
}
|
|
422
|
+
|
|
423
|
+
for (uint32_t k = 0; k < e->len; k++) {
|
|
424
|
+
uint32_t d = e->out[k];
|
|
425
|
+
if (se_range_contains(st->model->norm.mn, st->model->norm.mn_count, d))
|
|
426
|
+
continue;
|
|
427
|
+
se_status_t rc = emit_lowered(st, d, stop);
|
|
428
|
+
if (rc != SE_OK || *stop)
|
|
429
|
+
return rc;
|
|
430
|
+
}
|
|
431
|
+
return SE_OK;
|
|
432
|
+
}
|
|
433
|
+
|
|
434
|
+
static se_status_t emit_cleaned(token_state_t *st, uint32_t cp, int *stop) {
|
|
435
|
+
if (st->model->meta.clean_text) {
|
|
436
|
+
if (cp == 0 || cp == 0xFFFD || is_control(st->model, cp))
|
|
437
|
+
return SE_OK;
|
|
438
|
+
if (is_whitespace(st->model, cp))
|
|
439
|
+
cp = ' ';
|
|
440
|
+
}
|
|
441
|
+
|
|
442
|
+
if (st->model->meta.handle_chinese_chars && cp >= 0x3400 && se_is_cjk(cp)) {
|
|
443
|
+
se_status_t rc = emit_stripped(st, ' ', stop);
|
|
444
|
+
if (rc != SE_OK || *stop)
|
|
445
|
+
return rc;
|
|
446
|
+
rc = emit_stripped(st, cp, stop);
|
|
447
|
+
if (rc != SE_OK || *stop)
|
|
448
|
+
return rc;
|
|
449
|
+
return emit_stripped(st, ' ', stop);
|
|
450
|
+
}
|
|
451
|
+
|
|
452
|
+
return emit_stripped(st, cp, stop);
|
|
453
|
+
}
|
|
454
|
+
|
|
455
|
+
se_status_t se_tokenize(const se_model_t *model, se_scratch_t *sc, const uint8_t *input,
|
|
456
|
+
size_t input_len, uint32_t max_tokens, se_token_stats_t *stats,
|
|
457
|
+
se_error_t *err, volatile sig_atomic_t *cancelled) {
|
|
458
|
+
memset(stats, 0, sizeof(*stats));
|
|
459
|
+
|
|
460
|
+
token_state_t st;
|
|
461
|
+
memset(&st, 0, sizeof(st));
|
|
462
|
+
st.model = model;
|
|
463
|
+
st.scratch = sc;
|
|
464
|
+
st.max_tokens = max_tokens;
|
|
465
|
+
st.stats = stats;
|
|
466
|
+
st.err = err;
|
|
467
|
+
st.cancelled = cancelled;
|
|
468
|
+
|
|
469
|
+
size_t i = 0;
|
|
470
|
+
size_t iterations = 0;
|
|
471
|
+
while (i < input_len) {
|
|
472
|
+
if (((iterations++ & SE_CANCEL_CHECK_MASK) == 0) && token_cancelled(&st)) {
|
|
473
|
+
se_error_set(err, SE_ERR_INTERNAL, "operation cancelled");
|
|
474
|
+
return SE_ERR_INTERNAL;
|
|
475
|
+
}
|
|
476
|
+
|
|
477
|
+
uint32_t cp;
|
|
478
|
+
if (!decode_one(input, input_len, &i, &cp)) {
|
|
479
|
+
se_error_set(err, SE_ERR_INVALID_UTF8, "input is not valid UTF-8");
|
|
480
|
+
return SE_ERR_INVALID_UTF8;
|
|
481
|
+
}
|
|
482
|
+
|
|
483
|
+
int stop = 0;
|
|
484
|
+
se_status_t rc = emit_cleaned(&st, cp, &stop);
|
|
485
|
+
if (rc != SE_OK)
|
|
486
|
+
return rc;
|
|
487
|
+
if (stop)
|
|
488
|
+
break;
|
|
489
|
+
}
|
|
490
|
+
|
|
491
|
+
if (!stats->truncated) {
|
|
492
|
+
int stop = 0;
|
|
493
|
+
se_status_t rc = flush_segment(&st, &stop);
|
|
494
|
+
if (rc != SE_OK)
|
|
495
|
+
return rc;
|
|
496
|
+
}
|
|
497
|
+
|
|
498
|
+
if (st.n_ids > UINT32_MAX || st.n_unk > UINT32_MAX) {
|
|
499
|
+
se_error_set(err, SE_ERR_OOM, "input produced too many tokens");
|
|
500
|
+
return SE_ERR_OOM;
|
|
501
|
+
}
|
|
502
|
+
|
|
503
|
+
stats->token_count = (uint32_t)st.n_ids;
|
|
504
|
+
stats->unk_count = (uint32_t)st.n_unk;
|
|
505
|
+
return SE_OK;
|
|
506
|
+
}
|
|
@@ -0,0 +1,120 @@
|
|
|
1
|
+
#include "se_internal.h"
|
|
2
|
+
|
|
3
|
+
#include <string.h>
|
|
4
|
+
|
|
5
|
+
size_t se_utf8_decode(const uint8_t *src, size_t len, uint32_t *out, size_t out_cap, int *ok) {
|
|
6
|
+
size_t i = 0, n = 0;
|
|
7
|
+
*ok = 1;
|
|
8
|
+
|
|
9
|
+
while (i < len) {
|
|
10
|
+
if (n >= out_cap) {
|
|
11
|
+
*ok = 0;
|
|
12
|
+
return n;
|
|
13
|
+
}
|
|
14
|
+
uint8_t b0 = src[i];
|
|
15
|
+
uint32_t cp;
|
|
16
|
+
size_t need;
|
|
17
|
+
|
|
18
|
+
if (b0 < 0x80) {
|
|
19
|
+
cp = b0;
|
|
20
|
+
need = 0;
|
|
21
|
+
} else if ((b0 & 0xE0) == 0xC0) {
|
|
22
|
+
cp = b0 & 0x1Fu;
|
|
23
|
+
need = 1;
|
|
24
|
+
} else if ((b0 & 0xF0) == 0xE0) {
|
|
25
|
+
cp = b0 & 0x0Fu;
|
|
26
|
+
need = 2;
|
|
27
|
+
} else if ((b0 & 0xF8) == 0xF0) {
|
|
28
|
+
cp = b0 & 0x07u;
|
|
29
|
+
need = 3;
|
|
30
|
+
} else {
|
|
31
|
+
*ok = 0;
|
|
32
|
+
return n;
|
|
33
|
+
}
|
|
34
|
+
|
|
35
|
+
if (i + need >= len && need > 0) {
|
|
36
|
+
*ok = 0;
|
|
37
|
+
return n;
|
|
38
|
+
}
|
|
39
|
+
|
|
40
|
+
for (size_t k = 1; k <= need; k++) {
|
|
41
|
+
uint8_t bk = src[i + k];
|
|
42
|
+
if ((bk & 0xC0) != 0x80) {
|
|
43
|
+
*ok = 0;
|
|
44
|
+
return n;
|
|
45
|
+
}
|
|
46
|
+
cp = (cp << 6) | (uint32_t)(bk & 0x3Fu);
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
if ((need == 1 && cp < 0x80) || (need == 2 && cp < 0x800) || (need == 3 && cp < 0x10000)) {
|
|
50
|
+
*ok = 0;
|
|
51
|
+
return n;
|
|
52
|
+
}
|
|
53
|
+
if (cp > 0x10FFFF || (cp >= 0xD800 && cp <= 0xDFFF)) {
|
|
54
|
+
*ok = 0;
|
|
55
|
+
return n;
|
|
56
|
+
}
|
|
57
|
+
|
|
58
|
+
out[n++] = cp;
|
|
59
|
+
i += need + 1;
|
|
60
|
+
}
|
|
61
|
+
return n;
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
size_t se_utf8_encode(uint32_t cp, uint8_t *dst) {
|
|
65
|
+
if (cp < 0x80) {
|
|
66
|
+
dst[0] = (uint8_t)cp;
|
|
67
|
+
return 1;
|
|
68
|
+
}
|
|
69
|
+
if (cp < 0x800) {
|
|
70
|
+
dst[0] = (uint8_t)(0xC0 | (cp >> 6));
|
|
71
|
+
dst[1] = (uint8_t)(0x80 | (cp & 0x3F));
|
|
72
|
+
return 2;
|
|
73
|
+
}
|
|
74
|
+
if (cp < 0x10000) {
|
|
75
|
+
dst[0] = (uint8_t)(0xE0 | (cp >> 12));
|
|
76
|
+
dst[1] = (uint8_t)(0x80 | ((cp >> 6) & 0x3F));
|
|
77
|
+
dst[2] = (uint8_t)(0x80 | (cp & 0x3F));
|
|
78
|
+
return 3;
|
|
79
|
+
}
|
|
80
|
+
dst[0] = (uint8_t)(0xF0 | (cp >> 18));
|
|
81
|
+
dst[1] = (uint8_t)(0x80 | ((cp >> 12) & 0x3F));
|
|
82
|
+
dst[2] = (uint8_t)(0x80 | ((cp >> 6) & 0x3F));
|
|
83
|
+
dst[3] = (uint8_t)(0x80 | (cp & 0x3F));
|
|
84
|
+
return 4;
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
int se_range_contains(const se_range_t *ranges, uint32_t count, uint32_t cp) {
|
|
88
|
+
uint32_t lo = 0, hi = count;
|
|
89
|
+
while (lo < hi) {
|
|
90
|
+
uint32_t mid = lo + (hi - lo) / 2;
|
|
91
|
+
if (cp < ranges[mid].lo)
|
|
92
|
+
hi = mid;
|
|
93
|
+
else if (cp > ranges[mid].hi)
|
|
94
|
+
lo = mid + 1;
|
|
95
|
+
else
|
|
96
|
+
return 1;
|
|
97
|
+
}
|
|
98
|
+
return 0;
|
|
99
|
+
}
|
|
100
|
+
|
|
101
|
+
const se_map_entry_t *se_map_lookup(const se_map_entry_t *entries, uint32_t count, uint32_t cp) {
|
|
102
|
+
uint32_t lo = 0, hi = count;
|
|
103
|
+
while (lo < hi) {
|
|
104
|
+
uint32_t mid = lo + (hi - lo) / 2;
|
|
105
|
+
if (cp < entries[mid].cp)
|
|
106
|
+
hi = mid;
|
|
107
|
+
else if (cp > entries[mid].cp)
|
|
108
|
+
lo = mid + 1;
|
|
109
|
+
else
|
|
110
|
+
return &entries[mid];
|
|
111
|
+
}
|
|
112
|
+
return NULL;
|
|
113
|
+
}
|
|
114
|
+
|
|
115
|
+
int se_is_cjk(uint32_t cp) {
|
|
116
|
+
return (cp >= 0x4E00 && cp <= 0x9FFF) || (cp >= 0x3400 && cp <= 0x4DBF) ||
|
|
117
|
+
(cp >= 0x20000 && cp <= 0x2A6DF) || (cp >= 0x2A700 && cp <= 0x2B73F) ||
|
|
118
|
+
(cp >= 0x2B740 && cp <= 0x2B81F) || (cp >= 0x2B820 && cp <= 0x2CEAF) ||
|
|
119
|
+
(cp >= 0xF900 && cp <= 0xFAFF) || (cp >= 0x2F800 && cp <= 0x2FA1F);
|
|
120
|
+
}
|