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,266 @@
|
|
|
1
|
+
#ifndef SE_INTERNAL_H
|
|
2
|
+
#define SE_INTERNAL_H
|
|
3
|
+
|
|
4
|
+
#include <stddef.h>
|
|
5
|
+
#include <stdint.h>
|
|
6
|
+
#include <signal.h>
|
|
7
|
+
|
|
8
|
+
#define SE_MAGIC "SEMBv1\0\0"
|
|
9
|
+
#define SE_MAGIC_LEN 8
|
|
10
|
+
#define SE_FORMAT_VERSION 2u
|
|
11
|
+
#define SE_HEADER_SIZE 320u
|
|
12
|
+
#define SE_ALIGNMENT 64u
|
|
13
|
+
|
|
14
|
+
#define SE_TOKENIZER_BERT_WORDPIECE_V1 1u
|
|
15
|
+
#define SE_DTYPE_F32 1u
|
|
16
|
+
#define SE_POOLING_MEAN 1u
|
|
17
|
+
#define SE_NORMALIZATION_NONE 0u
|
|
18
|
+
#define SE_NORMALIZATION_L2 1u
|
|
19
|
+
|
|
20
|
+
#define SE_TRUNCATE_IDS_BEFORE_POOLING 1u
|
|
21
|
+
|
|
22
|
+
#define SE_UNK_INCLUDE 0u
|
|
23
|
+
#define SE_UNK_DROP 1u
|
|
24
|
+
|
|
25
|
+
#define SE_EMPTY_ZERO_VECTOR 0u
|
|
26
|
+
#define SE_EMPTY_RAISE 1u
|
|
27
|
+
|
|
28
|
+
#define SE_SLOT_EMPTY 0xFFFFFFFFu
|
|
29
|
+
|
|
30
|
+
enum {
|
|
31
|
+
SE_OFF_MAGIC = 0,
|
|
32
|
+
SE_OFF_FORMAT_VERSION = 8,
|
|
33
|
+
SE_OFF_HEADER_SIZE = 12,
|
|
34
|
+
SE_OFF_FLAGS = 16,
|
|
35
|
+
SE_OFF_DIM = 20,
|
|
36
|
+
SE_OFF_VOCAB_SIZE = 24,
|
|
37
|
+
SE_OFF_TOKENIZER_TYPE = 28,
|
|
38
|
+
SE_OFF_EMBEDDING_DTYPE = 32,
|
|
39
|
+
SE_OFF_POOLING_TYPE = 36,
|
|
40
|
+
SE_OFF_NORMALIZATION_TYPE = 40,
|
|
41
|
+
SE_OFF_MAX_TOKENS_DEFAULT = 44,
|
|
42
|
+
SE_OFF_TRUNCATION_POLICY = 48,
|
|
43
|
+
SE_OFF_ADD_SPECIAL_TOKENS = 52,
|
|
44
|
+
SE_OFF_UNK_POLICY = 56,
|
|
45
|
+
SE_OFF_EMPTY_POLICY = 60,
|
|
46
|
+
SE_OFF_DO_LOWER_CASE = 64,
|
|
47
|
+
SE_OFF_STRIP_ACCENTS = 68,
|
|
48
|
+
SE_OFF_HANDLE_CHINESE_CHARS = 72,
|
|
49
|
+
SE_OFF_CLEAN_TEXT = 76,
|
|
50
|
+
SE_OFF_MAX_INPUT_CHARS_PER_WORD = 80,
|
|
51
|
+
SE_OFF_PAD_ID = 84,
|
|
52
|
+
SE_OFF_UNK_ID = 88,
|
|
53
|
+
SE_OFF_CLS_ID = 92,
|
|
54
|
+
SE_OFF_SEP_ID = 96,
|
|
55
|
+
SE_OFF_MASK_ID = 100,
|
|
56
|
+
SE_OFF_HASH_TABLE_SIZE = 104,
|
|
57
|
+
SE_OFF_HASH_SEED = 108,
|
|
58
|
+
SE_OFF_SUBWORD_PREFIX_LEN = 112,
|
|
59
|
+
SE_OFF_SUBWORD_PREFIX = 116,
|
|
60
|
+
SE_OFF_MAX_TOKEN_CHARS = 124,
|
|
61
|
+
SE_OFF_SEC_VOCAB_STRINGS = 128,
|
|
62
|
+
SE_OFF_SEC_VOCAB_HASH = 144,
|
|
63
|
+
SE_OFF_SEC_EMBEDDINGS = 160,
|
|
64
|
+
SE_OFF_SEC_NORM_TABLES = 176,
|
|
65
|
+
SE_OFF_SEC_PROVENANCE = 192,
|
|
66
|
+
SE_OFF_SEC_ROOT_TRIE = 208,
|
|
67
|
+
SE_OFF_SEC_CONT_TRIE = 224,
|
|
68
|
+
SE_OFF_CHECKSUM = 240,
|
|
69
|
+
SE_OFF_MAX_PROBE = 304
|
|
70
|
+
};
|
|
71
|
+
|
|
72
|
+
typedef struct {
|
|
73
|
+
uint32_t hash;
|
|
74
|
+
uint32_t str_off;
|
|
75
|
+
uint32_t str_len;
|
|
76
|
+
uint32_t token_id;
|
|
77
|
+
} se_vocab_slot_t;
|
|
78
|
+
|
|
79
|
+
typedef struct {
|
|
80
|
+
uint32_t edge_start;
|
|
81
|
+
uint32_t edge_count;
|
|
82
|
+
uint32_t token_id;
|
|
83
|
+
uint32_t reserved;
|
|
84
|
+
} se_trie_node_t;
|
|
85
|
+
|
|
86
|
+
typedef struct {
|
|
87
|
+
uint32_t byte;
|
|
88
|
+
uint32_t child;
|
|
89
|
+
} se_trie_edge_t;
|
|
90
|
+
|
|
91
|
+
typedef struct {
|
|
92
|
+
const se_trie_node_t *nodes;
|
|
93
|
+
const se_trie_edge_t *edges;
|
|
94
|
+
uint32_t node_count;
|
|
95
|
+
uint32_t edge_count;
|
|
96
|
+
} se_trie_t;
|
|
97
|
+
|
|
98
|
+
typedef struct {
|
|
99
|
+
uint32_t cp;
|
|
100
|
+
uint32_t len;
|
|
101
|
+
uint32_t out[4];
|
|
102
|
+
} se_map_entry_t;
|
|
103
|
+
|
|
104
|
+
typedef struct {
|
|
105
|
+
uint32_t lo;
|
|
106
|
+
uint32_t hi;
|
|
107
|
+
} se_range_t;
|
|
108
|
+
|
|
109
|
+
typedef struct {
|
|
110
|
+
const se_map_entry_t *lower;
|
|
111
|
+
uint32_t lower_count;
|
|
112
|
+
const se_map_entry_t *nfd;
|
|
113
|
+
uint32_t nfd_count;
|
|
114
|
+
const se_range_t *mn;
|
|
115
|
+
uint32_t mn_count;
|
|
116
|
+
const se_range_t *punct;
|
|
117
|
+
uint32_t punct_count;
|
|
118
|
+
const se_range_t *control;
|
|
119
|
+
uint32_t control_count;
|
|
120
|
+
const se_range_t *whitespace;
|
|
121
|
+
uint32_t whitespace_count;
|
|
122
|
+
} se_norm_tables_t;
|
|
123
|
+
|
|
124
|
+
typedef struct {
|
|
125
|
+
uint32_t dim;
|
|
126
|
+
uint32_t vocab_size;
|
|
127
|
+
uint32_t tokenizer_type;
|
|
128
|
+
uint32_t embedding_dtype;
|
|
129
|
+
uint32_t pooling_type;
|
|
130
|
+
uint32_t normalization_type;
|
|
131
|
+
uint32_t max_tokens_default;
|
|
132
|
+
uint32_t truncation_policy;
|
|
133
|
+
uint32_t add_special_tokens;
|
|
134
|
+
uint32_t unk_policy;
|
|
135
|
+
uint32_t empty_policy;
|
|
136
|
+
uint32_t do_lower_case;
|
|
137
|
+
uint32_t strip_accents;
|
|
138
|
+
uint32_t handle_chinese_chars;
|
|
139
|
+
uint32_t clean_text;
|
|
140
|
+
uint32_t max_input_chars_per_word;
|
|
141
|
+
uint32_t max_token_chars;
|
|
142
|
+
uint32_t max_probe;
|
|
143
|
+
uint32_t pad_id;
|
|
144
|
+
uint32_t unk_id;
|
|
145
|
+
uint32_t cls_id;
|
|
146
|
+
uint32_t sep_id;
|
|
147
|
+
uint32_t mask_id;
|
|
148
|
+
uint32_t hash_table_size;
|
|
149
|
+
uint32_t hash_seed;
|
|
150
|
+
uint32_t subword_prefix_len;
|
|
151
|
+
char subword_prefix[8];
|
|
152
|
+
} se_meta_t;
|
|
153
|
+
|
|
154
|
+
typedef struct {
|
|
155
|
+
void *map_base;
|
|
156
|
+
size_t map_size;
|
|
157
|
+
int mapped;
|
|
158
|
+
|
|
159
|
+
se_meta_t meta;
|
|
160
|
+
se_norm_tables_t norm;
|
|
161
|
+
|
|
162
|
+
const char *vocab_strings;
|
|
163
|
+
size_t vocab_strings_size;
|
|
164
|
+
const se_vocab_slot_t *vocab_hash;
|
|
165
|
+
se_trie_t root_trie;
|
|
166
|
+
se_trie_t cont_trie;
|
|
167
|
+
const float *embeddings;
|
|
168
|
+
const char *provenance;
|
|
169
|
+
size_t provenance_size;
|
|
170
|
+
} se_model_t;
|
|
171
|
+
|
|
172
|
+
typedef struct {
|
|
173
|
+
uint32_t *cps;
|
|
174
|
+
size_t cps_cap;
|
|
175
|
+
uint32_t *cps2;
|
|
176
|
+
size_t cps2_cap;
|
|
177
|
+
uint8_t *bytes;
|
|
178
|
+
size_t bytes_cap;
|
|
179
|
+
uint32_t *ids;
|
|
180
|
+
size_t ids_cap;
|
|
181
|
+
float *acc;
|
|
182
|
+
size_t acc_cap;
|
|
183
|
+
} se_scratch_t;
|
|
184
|
+
|
|
185
|
+
typedef enum {
|
|
186
|
+
SE_OK = 0,
|
|
187
|
+
SE_ERR_INVALID_FORMAT,
|
|
188
|
+
SE_ERR_UNSUPPORTED_VERSION,
|
|
189
|
+
SE_ERR_UNSUPPORTED_TOKENIZER,
|
|
190
|
+
SE_ERR_INVALID_UTF8,
|
|
191
|
+
SE_ERR_OOM,
|
|
192
|
+
SE_ERR_IO,
|
|
193
|
+
SE_ERR_EMPTY_INPUT,
|
|
194
|
+
SE_ERR_INTERNAL
|
|
195
|
+
} se_status_t;
|
|
196
|
+
|
|
197
|
+
typedef struct {
|
|
198
|
+
se_status_t status;
|
|
199
|
+
char message[256];
|
|
200
|
+
} se_error_t;
|
|
201
|
+
|
|
202
|
+
#if defined(__GNUC__) || defined(__clang__)
|
|
203
|
+
#define SE_PRINTF_FORMAT(fmt_index, first_arg) __attribute__((format(printf, fmt_index, first_arg)))
|
|
204
|
+
#else
|
|
205
|
+
#define SE_PRINTF_FORMAT(fmt_index, first_arg)
|
|
206
|
+
#endif
|
|
207
|
+
|
|
208
|
+
void se_error_clear(se_error_t *err);
|
|
209
|
+
void se_error_set(se_error_t *err, se_status_t status, const char *fmt, ...) SE_PRINTF_FORMAT(3, 4);
|
|
210
|
+
|
|
211
|
+
se_status_t se_model_open(se_model_t *model, const char *path, se_error_t *err);
|
|
212
|
+
void se_model_close(se_model_t *model);
|
|
213
|
+
size_t se_model_memsize(const se_model_t *model);
|
|
214
|
+
uint32_t se_hash_bytes(const uint8_t *data, size_t len, uint32_t seed);
|
|
215
|
+
int se_vocab_lookup(const se_model_t *model, const uint8_t *bytes, size_t len, uint32_t *id_out);
|
|
216
|
+
int se_vocab_lookup_piece(const se_model_t *model, const uint8_t *prefix, size_t prefix_len,
|
|
217
|
+
const uint8_t *bytes, size_t len, uint32_t *id_out);
|
|
218
|
+
int se_trie_longest_match(const se_trie_t *trie, const uint8_t *bytes, size_t len, uint32_t *id_out,
|
|
219
|
+
size_t *matched_len_out);
|
|
220
|
+
|
|
221
|
+
size_t se_utf8_decode(const uint8_t *src, size_t len, uint32_t *out, size_t out_cap, int *ok);
|
|
222
|
+
size_t se_utf8_encode(uint32_t cp, uint8_t *dst);
|
|
223
|
+
int se_range_contains(const se_range_t *ranges, uint32_t count, uint32_t cp);
|
|
224
|
+
const se_map_entry_t *se_map_lookup(const se_map_entry_t *entries, uint32_t count, uint32_t cp);
|
|
225
|
+
int se_is_cjk(uint32_t cp);
|
|
226
|
+
static inline int se_is_ascii_punct(uint32_t cp) {
|
|
227
|
+
return (cp >= 33 && cp <= 47) || (cp >= 58 && cp <= 64) || (cp >= 91 && cp <= 96) ||
|
|
228
|
+
(cp >= 123 && cp <= 126);
|
|
229
|
+
}
|
|
230
|
+
|
|
231
|
+
static inline int se_is_ascii_whitespace(uint32_t cp) {
|
|
232
|
+
return cp == ' ' || cp == '\t' || cp == '\n' || cp == '\r';
|
|
233
|
+
}
|
|
234
|
+
|
|
235
|
+
static inline int se_is_ascii_boundary(uint32_t cp) {
|
|
236
|
+
return cp < 0x80 && (se_is_ascii_whitespace(cp) || se_is_ascii_punct(cp));
|
|
237
|
+
}
|
|
238
|
+
|
|
239
|
+
void se_scratch_init(se_scratch_t *s);
|
|
240
|
+
void se_scratch_free(se_scratch_t *s);
|
|
241
|
+
int se_scratch_reserve(se_scratch_t *s, uint32_t dim);
|
|
242
|
+
|
|
243
|
+
size_t se_prefix_boundary_len(const se_model_t *model, const uint8_t *input, size_t input_len,
|
|
244
|
+
size_t target, size_t backscan);
|
|
245
|
+
|
|
246
|
+
typedef struct {
|
|
247
|
+
uint32_t token_count;
|
|
248
|
+
uint32_t unk_count;
|
|
249
|
+
uint32_t truncated;
|
|
250
|
+
} se_token_stats_t;
|
|
251
|
+
|
|
252
|
+
se_status_t se_tokenize(const se_model_t *model, se_scratch_t *scratch, const uint8_t *input,
|
|
253
|
+
size_t input_len, uint32_t max_tokens, se_token_stats_t *stats,
|
|
254
|
+
se_error_t *err, volatile sig_atomic_t *cancelled);
|
|
255
|
+
|
|
256
|
+
se_status_t se_embed_one(const se_model_t *model, se_scratch_t *scratch, const uint8_t *input,
|
|
257
|
+
size_t input_len, uint32_t max_tokens, float *out, se_token_stats_t *stats,
|
|
258
|
+
se_error_t *err, volatile sig_atomic_t *cancelled);
|
|
259
|
+
se_status_t se_embed_ids(const se_model_t *model, se_scratch_t *scratch, const uint32_t *ids,
|
|
260
|
+
size_t n_ids, float *out, se_token_stats_t *stats, se_error_t *err,
|
|
261
|
+
volatile sig_atomic_t *cancelled);
|
|
262
|
+
|
|
263
|
+
void se_l2_normalize(float *vec, uint32_t dim);
|
|
264
|
+
size_t se_model_warmup(const se_model_t *model);
|
|
265
|
+
|
|
266
|
+
#endif
|