static_embeddings 0.1.5 → 1.5.6
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 +4 -4
- data/CHANGELOG.md +50 -0
- data/README.md +17 -2
- data/docs/ARCHITECTURE.md +14 -5
- data/docs/MODEL_AUDIT.md +110 -0
- data/lib/models/demo.semb +0 -0
- data/lib/static_embeddings/bert_wordpiece.rb +191 -0
- data/lib/static_embeddings/canonical.rb +50 -0
- data/lib/static_embeddings/cli.rb +88 -63
- data/lib/static_embeddings/codec.rb +45 -0
- data/lib/static_embeddings/conversion.rb +58 -0
- data/lib/static_embeddings/errors.rb +2 -2
- data/lib/static_embeddings/format/constants.rb +109 -0
- data/lib/static_embeddings/format/hash_table.rb +69 -0
- data/lib/static_embeddings/format/trie.rb +78 -0
- data/lib/static_embeddings/format/verifier.rb +41 -0
- data/lib/static_embeddings/format/writer.rb +131 -0
- data/lib/static_embeddings/format.rb +3 -338
- data/lib/static_embeddings/importers/model2vec.rb +52 -0
- data/lib/static_embeddings/importers/sentence_transformers_static.rb +103 -0
- data/lib/static_embeddings/importers/support.rb +111 -0
- data/lib/static_embeddings/importers.rb +50 -0
- data/lib/static_embeddings/model.rb +35 -20
- data/lib/static_embeddings/paths.rb +8 -12
- data/lib/static_embeddings/provenance.rb +58 -0
- data/lib/static_embeddings/reference.rb +50 -40
- data/lib/static_embeddings/row_prefix_payload.rb +59 -0
- data/lib/static_embeddings/version.rb +1 -1
- data/lib/static_embeddings.rb +29 -55
- data/static_embeddings.gemspec +2 -2
- data/tools/check_model2vec_parity.rb +5 -1
- data/tools/check_st_parity.rb +125 -0
- data/tools/eval_retrieval.rb +58 -0
- metadata +24 -6
- data/lib/static_embeddings/converter.rb +0 -328
|
@@ -0,0 +1,131 @@
|
|
|
1
|
+
require "digest"
|
|
2
|
+
require "static_embeddings/format/constants"
|
|
3
|
+
require "static_embeddings/format/hash_table"
|
|
4
|
+
require "static_embeddings/format/trie"
|
|
5
|
+
|
|
6
|
+
module StaticEmbeddings
|
|
7
|
+
module Format
|
|
8
|
+
module Writer
|
|
9
|
+
module_function
|
|
10
|
+
|
|
11
|
+
def call(path:, meta:, tokens:, matrix:, norm_tables:, provenance:)
|
|
12
|
+
hash_size, strings, hash_blob, max_probe = HashTable.build(tokens)
|
|
13
|
+
root_trie, continuation_trie = WordPieceTrie.build(tokens, meta.fetch(:subword_prefix))
|
|
14
|
+
payloads = payloads(strings, hash_blob, matrix, norm_tables, provenance, root_trie, continuation_trie)
|
|
15
|
+
sections, file_size = layout(payloads)
|
|
16
|
+
header = build_header(meta, tokens.length, hash_size, max_probe, sections)
|
|
17
|
+
checksum = write_file(path, header, payloads, sections)
|
|
18
|
+
|
|
19
|
+
{ bytes: file_size, sha256: checksum.unpack1("H*"), hash_table_size: hash_size }
|
|
20
|
+
end
|
|
21
|
+
|
|
22
|
+
def payloads(strings, hash_blob, matrix, norm_tables, provenance, root_trie, continuation_trie)
|
|
23
|
+
{
|
|
24
|
+
vocab_strings: strings,
|
|
25
|
+
vocab_hash: hash_blob,
|
|
26
|
+
embeddings: matrix,
|
|
27
|
+
norm_tables: norm_tables,
|
|
28
|
+
provenance: provenance,
|
|
29
|
+
root_trie: root_trie,
|
|
30
|
+
continuation_trie: continuation_trie
|
|
31
|
+
}
|
|
32
|
+
end
|
|
33
|
+
|
|
34
|
+
def layout(payloads)
|
|
35
|
+
offset = HEADER_SIZE
|
|
36
|
+
sections = payloads.each_with_object({}) do |(name, payload), out|
|
|
37
|
+
offset += (ALIGNMENT - (offset % ALIGNMENT)) % ALIGNMENT
|
|
38
|
+
out[name] = [offset, payload.bytesize]
|
|
39
|
+
offset += payload.bytesize
|
|
40
|
+
end
|
|
41
|
+
[sections, offset]
|
|
42
|
+
end
|
|
43
|
+
|
|
44
|
+
def write_file(path, header, payloads, sections)
|
|
45
|
+
digest = Digest::SHA256.new
|
|
46
|
+
File.open(path, "wb") do |io|
|
|
47
|
+
write_chunk(io, digest, header)
|
|
48
|
+
offset = HEADER_SIZE
|
|
49
|
+
|
|
50
|
+
payloads.each do |name, payload|
|
|
51
|
+
target = sections.fetch(name).first
|
|
52
|
+
padding = target - offset
|
|
53
|
+
write_chunk(io, digest, "\0".b * padding) if padding.positive?
|
|
54
|
+
write_payload(io, digest, payload)
|
|
55
|
+
offset = target + payload.bytesize
|
|
56
|
+
end
|
|
57
|
+
end
|
|
58
|
+
|
|
59
|
+
digest.digest.tap do |checksum|
|
|
60
|
+
File.open(path, "r+b") do |io|
|
|
61
|
+
io.seek(CHECKSUM_OFFSET, IO::SEEK_SET)
|
|
62
|
+
io.write(checksum)
|
|
63
|
+
end
|
|
64
|
+
end
|
|
65
|
+
end
|
|
66
|
+
|
|
67
|
+
def write_payload(io, digest, payload)
|
|
68
|
+
if payload.respond_to?(:each_chunk)
|
|
69
|
+
payload.each_chunk { |chunk| write_chunk(io, digest, chunk) }
|
|
70
|
+
else
|
|
71
|
+
write_chunk(io, digest, payload)
|
|
72
|
+
end
|
|
73
|
+
end
|
|
74
|
+
|
|
75
|
+
def write_chunk(io, digest, chunk)
|
|
76
|
+
io.write(chunk)
|
|
77
|
+
digest << chunk
|
|
78
|
+
end
|
|
79
|
+
|
|
80
|
+
def build_header(meta, vocab_size, hash_size, max_probe, sections)
|
|
81
|
+
header = "\0".b * HEADER_SIZE
|
|
82
|
+
header[0, 8] = MAGIC.b
|
|
83
|
+
|
|
84
|
+
HEADER_U32.each { |offset, value| put_u32(header, offset, value) }
|
|
85
|
+
META_U32.each { |offset, key| put_u32(header, offset, meta.fetch(key)) }
|
|
86
|
+
META_BOOL.each { |offset, key| put_u32(header, offset, meta.fetch(key) ? 1 : 0) }
|
|
87
|
+
put_u32(header, 24, vocab_size)
|
|
88
|
+
put_u32(header, 104, hash_size)
|
|
89
|
+
put_u32(header, MAX_TOKEN_CHARS_OFFSET, meta.fetch(:max_token_chars))
|
|
90
|
+
put_u32(header, MAX_PROBE_OFFSET, max_probe)
|
|
91
|
+
put_u32(header, ADDED_TOKEN_MASK_OFFSET, meta.fetch(:added_token_mask, 0))
|
|
92
|
+
put_prefix(header, meta.fetch(:subword_prefix))
|
|
93
|
+
SECTION_FIELDS.each { |name, field| put_section(header, field, sections.fetch(name)) }
|
|
94
|
+
header
|
|
95
|
+
end
|
|
96
|
+
|
|
97
|
+
def put_prefix(header, prefix)
|
|
98
|
+
bytes = prefix.b
|
|
99
|
+
raise ArgumentError, "subword prefix too long" if bytes.bytesize > 8
|
|
100
|
+
|
|
101
|
+
put_u32(header, 112, bytes.bytesize)
|
|
102
|
+
header[116, 8] = bytes.ljust(8, "\0")
|
|
103
|
+
end
|
|
104
|
+
|
|
105
|
+
def put_section(header, offset, section)
|
|
106
|
+
put_u64(header, offset, section.fetch(0))
|
|
107
|
+
put_u64(header, offset + 8, section.fetch(1))
|
|
108
|
+
end
|
|
109
|
+
|
|
110
|
+
def put_u32(buffer, offset, value)
|
|
111
|
+
integer = Integer(value)
|
|
112
|
+
unless integer.between?(0, UINT32_MAX)
|
|
113
|
+
raise ArgumentError, "u32 value out of range: #{integer}"
|
|
114
|
+
end
|
|
115
|
+
buffer[offset, 4] = [integer].pack("V")
|
|
116
|
+
end
|
|
117
|
+
|
|
118
|
+
def put_u64(buffer, offset, value)
|
|
119
|
+
integer = Integer(value)
|
|
120
|
+
unless integer.between?(0, UINT64_MAX)
|
|
121
|
+
raise ArgumentError, "u64 value out of range: #{integer}"
|
|
122
|
+
end
|
|
123
|
+
buffer[offset, 8] = [integer & UINT32_MAX, integer >> 32].pack("V2")
|
|
124
|
+
end
|
|
125
|
+
end
|
|
126
|
+
|
|
127
|
+
def self.write(**kwargs)
|
|
128
|
+
Writer.call(**kwargs)
|
|
129
|
+
end
|
|
130
|
+
end
|
|
131
|
+
end
|
|
@@ -1,338 +1,3 @@
|
|
|
1
|
-
require "
|
|
2
|
-
|
|
3
|
-
|
|
4
|
-
module Format
|
|
5
|
-
MAGIC = "SEMBv1\0\0"
|
|
6
|
-
VERSION = 3
|
|
7
|
-
HEADER_SIZE = 320
|
|
8
|
-
ALIGNMENT = 64
|
|
9
|
-
|
|
10
|
-
TOKENIZER_BERT_WORDPIECE_V1 = 1
|
|
11
|
-
DTYPE_F32 = 1
|
|
12
|
-
POOLING_MEAN = 1
|
|
13
|
-
NORMALIZATION_NONE = 0
|
|
14
|
-
NORMALIZATION_L2 = 1
|
|
15
|
-
TRUNCATE_USABLE_IDS_BEFORE_POOLING = 2
|
|
16
|
-
UNK_INCLUDE = 0
|
|
17
|
-
UNK_DROP = 1
|
|
18
|
-
EMPTY_ZERO_VECTOR = 0
|
|
19
|
-
EMPTY_RAISE = 1
|
|
20
|
-
|
|
21
|
-
SLOT_EMPTY = 0xFFFFFFFF
|
|
22
|
-
HASH_SEED = 2_166_136_261
|
|
23
|
-
LOAD_FACTOR = 0.70
|
|
24
|
-
|
|
25
|
-
MAX_TOKEN_CHARS_OFFSET = 124
|
|
26
|
-
CHECKSUM_OFFSET = 240
|
|
27
|
-
CHECKSUM_SIZE = 32
|
|
28
|
-
MAX_PROBE_OFFSET = 304
|
|
29
|
-
ADDED_TOKEN_MASK_OFFSET = 308
|
|
30
|
-
|
|
31
|
-
ADDED_PAD = 1 << 0
|
|
32
|
-
ADDED_UNK = 1 << 1
|
|
33
|
-
ADDED_CLS = 1 << 2
|
|
34
|
-
ADDED_SEP = 1 << 3
|
|
35
|
-
ADDED_MASK = 1 << 4
|
|
36
|
-
ADDED_TOKEN_MASK_ALL = ADDED_PAD | ADDED_UNK | ADDED_CLS | ADDED_SEP | ADDED_MASK
|
|
37
|
-
|
|
38
|
-
SECTION_FIELDS = {
|
|
39
|
-
vocab_strings: 128,
|
|
40
|
-
vocab_hash: 144,
|
|
41
|
-
embeddings: 160,
|
|
42
|
-
norm_tables: 176,
|
|
43
|
-
provenance: 192,
|
|
44
|
-
root_trie: 208,
|
|
45
|
-
continuation_trie: 224
|
|
46
|
-
}.freeze
|
|
47
|
-
|
|
48
|
-
HEADER_U32 = {
|
|
49
|
-
8 => VERSION,
|
|
50
|
-
12 => HEADER_SIZE,
|
|
51
|
-
16 => 1,
|
|
52
|
-
28 => TOKENIZER_BERT_WORDPIECE_V1,
|
|
53
|
-
32 => DTYPE_F32,
|
|
54
|
-
36 => POOLING_MEAN,
|
|
55
|
-
48 => TRUNCATE_USABLE_IDS_BEFORE_POOLING,
|
|
56
|
-
108 => HASH_SEED
|
|
57
|
-
}.freeze
|
|
58
|
-
|
|
59
|
-
META_U32 = {
|
|
60
|
-
20 => :dim,
|
|
61
|
-
40 => :normalization_type,
|
|
62
|
-
44 => :max_tokens_default,
|
|
63
|
-
56 => :unk_policy,
|
|
64
|
-
60 => :empty_policy,
|
|
65
|
-
80 => :max_input_chars_per_word,
|
|
66
|
-
84 => :pad_id,
|
|
67
|
-
88 => :unk_id,
|
|
68
|
-
92 => :cls_id,
|
|
69
|
-
96 => :sep_id,
|
|
70
|
-
100 => :mask_id
|
|
71
|
-
}.freeze
|
|
72
|
-
|
|
73
|
-
META_BOOL = {
|
|
74
|
-
52 => :add_special_tokens,
|
|
75
|
-
64 => :do_lower_case,
|
|
76
|
-
68 => :strip_accents,
|
|
77
|
-
72 => :handle_chinese_chars,
|
|
78
|
-
76 => :clean_text
|
|
79
|
-
}.freeze
|
|
80
|
-
|
|
81
|
-
VERIFY_CHUNK_BYTES = 1024 * 1024
|
|
82
|
-
|
|
83
|
-
module_function
|
|
84
|
-
|
|
85
|
-
def hash_bytes(str, seed = HASH_SEED)
|
|
86
|
-
str.each_byte.reduce(seed) { |h, b| ((h ^ b) * 16_777_619) & 0xFFFFFFFF }
|
|
87
|
-
end
|
|
88
|
-
|
|
89
|
-
def next_power_of_two(n)
|
|
90
|
-
1 << (n - 1).bit_length
|
|
91
|
-
end
|
|
92
|
-
|
|
93
|
-
def build_hash_table(tokens)
|
|
94
|
-
size = next_power_of_two((tokens.length / LOAD_FACTOR).ceil + 1)
|
|
95
|
-
slots = Array.new(size)
|
|
96
|
-
strings = binary_string
|
|
97
|
-
max_probe = 0
|
|
98
|
-
|
|
99
|
-
tokens.each_with_index do |token, id|
|
|
100
|
-
bytes = token.b
|
|
101
|
-
offset = strings.bytesize
|
|
102
|
-
strings << bytes
|
|
103
|
-
probe = insert_slot!(slots, tokens, size, bytes, hash_bytes(bytes), offset, id)
|
|
104
|
-
max_probe = probe if probe > max_probe
|
|
105
|
-
end
|
|
106
|
-
|
|
107
|
-
[size, strings, pack_slots(slots), max_probe]
|
|
108
|
-
end
|
|
109
|
-
|
|
110
|
-
def write(path:, meta:, tokens:, matrix:, norm_tables:, provenance:)
|
|
111
|
-
hash_size, strings, hash_blob, max_probe = build_hash_table(tokens)
|
|
112
|
-
root_trie, continuation_trie = build_wordpiece_tries(tokens, meta.fetch(:subword_prefix))
|
|
113
|
-
payloads = {
|
|
114
|
-
vocab_strings: strings,
|
|
115
|
-
vocab_hash: hash_blob,
|
|
116
|
-
embeddings: matrix,
|
|
117
|
-
norm_tables: norm_tables,
|
|
118
|
-
provenance: provenance,
|
|
119
|
-
root_trie: root_trie,
|
|
120
|
-
continuation_trie: continuation_trie
|
|
121
|
-
}
|
|
122
|
-
sections, file_size = layout_sections(payloads)
|
|
123
|
-
header = build_header(meta, tokens.length, hash_size, max_probe, sections)
|
|
124
|
-
|
|
125
|
-
digest = Digest::SHA256.new
|
|
126
|
-
File.open(path, "wb") do |io|
|
|
127
|
-
io.write(header)
|
|
128
|
-
digest << header
|
|
129
|
-
offset = HEADER_SIZE
|
|
130
|
-
|
|
131
|
-
payloads.each do |name, payload|
|
|
132
|
-
target = sections.fetch(name).first
|
|
133
|
-
padding = target - offset
|
|
134
|
-
if padding.positive?
|
|
135
|
-
zeros = "\0".b * padding
|
|
136
|
-
io.write(zeros)
|
|
137
|
-
digest << zeros
|
|
138
|
-
end
|
|
139
|
-
write_payload(io, digest, payload)
|
|
140
|
-
offset = target + payload.bytesize
|
|
141
|
-
end
|
|
142
|
-
end
|
|
143
|
-
|
|
144
|
-
checksum = digest.digest
|
|
145
|
-
File.open(path, "r+b") do |io|
|
|
146
|
-
io.seek(CHECKSUM_OFFSET, IO::SEEK_SET)
|
|
147
|
-
io.write(checksum)
|
|
148
|
-
end
|
|
149
|
-
|
|
150
|
-
{ bytes: file_size, sha256: checksum.unpack1("H*"), hash_table_size: hash_size }
|
|
151
|
-
end
|
|
152
|
-
|
|
153
|
-
def write_payload(io, digest, payload)
|
|
154
|
-
if payload.respond_to?(:each_chunk)
|
|
155
|
-
payload.each_chunk do |chunk|
|
|
156
|
-
io.write(chunk)
|
|
157
|
-
digest << chunk
|
|
158
|
-
end
|
|
159
|
-
else
|
|
160
|
-
io.write(payload)
|
|
161
|
-
digest << payload
|
|
162
|
-
end
|
|
163
|
-
end
|
|
164
|
-
|
|
165
|
-
def verify(path)
|
|
166
|
-
raise InvalidModelError, "file too small" if File.size(path) < HEADER_SIZE
|
|
167
|
-
|
|
168
|
-
File.open(path, "rb") do |io|
|
|
169
|
-
header = io.read(HEADER_SIZE)
|
|
170
|
-
raise InvalidModelError, "bad magic" unless header.byteslice(0, 8) == MAGIC.b
|
|
171
|
-
|
|
172
|
-
stored = header.byteslice(CHECKSUM_OFFSET, CHECKSUM_SIZE)
|
|
173
|
-
actual = streaming_checksum(io, header)
|
|
174
|
-
|
|
175
|
-
{ ok: stored == actual, expected: actual.unpack1("H*"), stored: stored.unpack1("H*") }
|
|
176
|
-
end
|
|
177
|
-
end
|
|
178
|
-
|
|
179
|
-
def streaming_checksum(io, header)
|
|
180
|
-
digest = Digest::SHA256.new
|
|
181
|
-
zeroed = header.dup
|
|
182
|
-
zeroed[CHECKSUM_OFFSET, CHECKSUM_SIZE] = "\0".b * CHECKSUM_SIZE
|
|
183
|
-
digest << zeroed
|
|
184
|
-
|
|
185
|
-
buffer = String.new(capacity: VERIFY_CHUNK_BYTES)
|
|
186
|
-
digest << buffer while io.read(VERIFY_CHUNK_BYTES, buffer)
|
|
187
|
-
digest.digest
|
|
188
|
-
end
|
|
189
|
-
|
|
190
|
-
def binary_string(capacity = nil)
|
|
191
|
-
str = capacity ? String.new(capacity: capacity) : +""
|
|
192
|
-
str.force_encoding(Encoding::BINARY)
|
|
193
|
-
end
|
|
194
|
-
|
|
195
|
-
def insert_slot!(slots, tokens, size, bytes, hash, offset, id)
|
|
196
|
-
pos = hash & (size - 1)
|
|
197
|
-
probe = 1
|
|
198
|
-
loop do
|
|
199
|
-
slot = slots[pos]
|
|
200
|
-
unless slot
|
|
201
|
-
slots[pos] = [hash, offset, bytes.bytesize, id]
|
|
202
|
-
return probe
|
|
203
|
-
end
|
|
204
|
-
|
|
205
|
-
if slot[0] == hash && slot[2] == bytes.bytesize && tokens.fetch(slot[3]).b == bytes
|
|
206
|
-
raise ArgumentError, "duplicate token in vocabulary: #{tokens.fetch(slot[3]).inspect}"
|
|
207
|
-
end
|
|
208
|
-
|
|
209
|
-
pos = (pos + 1) & (size - 1)
|
|
210
|
-
probe += 1
|
|
211
|
-
end
|
|
212
|
-
end
|
|
213
|
-
|
|
214
|
-
def pack_slots(slots)
|
|
215
|
-
empty = [0, 0, 0, SLOT_EMPTY].pack("V4")
|
|
216
|
-
packed = binary_string(slots.length * 16)
|
|
217
|
-
slots.each { |slot| packed << (slot ? slot.pack("V4") : empty) }
|
|
218
|
-
packed
|
|
219
|
-
end
|
|
220
|
-
|
|
221
|
-
def build_wordpiece_tries(tokens, prefix)
|
|
222
|
-
root = TrieBuilder.new
|
|
223
|
-
continuation = TrieBuilder.new
|
|
224
|
-
prefix_bytes = prefix.b
|
|
225
|
-
|
|
226
|
-
tokens.each_with_index do |token, id|
|
|
227
|
-
bytes = token.b
|
|
228
|
-
if !prefix_bytes.empty? && bytes.start_with?(prefix_bytes) && bytes.bytesize > prefix_bytes.bytesize
|
|
229
|
-
continuation.insert(bytes.byteslice(prefix_bytes.bytesize, bytes.bytesize - prefix_bytes.bytesize), id)
|
|
230
|
-
else
|
|
231
|
-
root.insert(bytes, id)
|
|
232
|
-
end
|
|
233
|
-
end
|
|
234
|
-
|
|
235
|
-
[root.pack, continuation.pack]
|
|
236
|
-
end
|
|
237
|
-
|
|
238
|
-
class TrieBuilder
|
|
239
|
-
Node = Struct.new(:terminal, :children, keyword_init: true)
|
|
240
|
-
|
|
241
|
-
def initialize
|
|
242
|
-
@nodes = [Node.new(terminal: SLOT_EMPTY, children: {})]
|
|
243
|
-
end
|
|
244
|
-
|
|
245
|
-
def insert(bytes, id)
|
|
246
|
-
return if bytes.empty?
|
|
247
|
-
|
|
248
|
-
node_index = 0
|
|
249
|
-
bytes.each_byte do |byte|
|
|
250
|
-
node = @nodes[node_index]
|
|
251
|
-
child = node.children[byte]
|
|
252
|
-
unless child
|
|
253
|
-
child = @nodes.length
|
|
254
|
-
node.children[byte] = child
|
|
255
|
-
@nodes << Node.new(terminal: SLOT_EMPTY, children: {})
|
|
256
|
-
end
|
|
257
|
-
node_index = child
|
|
258
|
-
end
|
|
259
|
-
|
|
260
|
-
node = @nodes[node_index]
|
|
261
|
-
raise ArgumentError, "duplicate trie key for token id #{id}" unless node.terminal == SLOT_EMPTY
|
|
262
|
-
|
|
263
|
-
node.terminal = id
|
|
264
|
-
end
|
|
265
|
-
|
|
266
|
-
def pack
|
|
267
|
-
edges = []
|
|
268
|
-
node_records = @nodes.map do |node|
|
|
269
|
-
start = edges.length
|
|
270
|
-
node.children.sort_by { |byte, _| byte }.each do |byte, child|
|
|
271
|
-
edges << [byte, child]
|
|
272
|
-
end
|
|
273
|
-
[start, node.children.length, node.terminal, 0]
|
|
274
|
-
end
|
|
275
|
-
|
|
276
|
-
packed = Format.binary_string(16 + node_records.length * 16 + edges.length * 8)
|
|
277
|
-
packed << [node_records.length, edges.length, 0, 0].pack("V4")
|
|
278
|
-
node_records.each { |record| packed << record.pack("V4") }
|
|
279
|
-
edges.each { |edge| packed << edge.pack("V2") }
|
|
280
|
-
packed
|
|
281
|
-
end
|
|
282
|
-
end
|
|
283
|
-
|
|
284
|
-
def layout_sections(payloads)
|
|
285
|
-
offset = HEADER_SIZE
|
|
286
|
-
sections = {}
|
|
287
|
-
payloads.each do |name, payload|
|
|
288
|
-
padding = (ALIGNMENT - (offset % ALIGNMENT)) % ALIGNMENT
|
|
289
|
-
offset += padding
|
|
290
|
-
sections[name] = [offset, payload.bytesize]
|
|
291
|
-
offset += payload.bytesize
|
|
292
|
-
end
|
|
293
|
-
[sections, offset]
|
|
294
|
-
end
|
|
295
|
-
|
|
296
|
-
def build_header(meta, vocab_size, hash_size, max_probe, sections)
|
|
297
|
-
header = "\0".b * HEADER_SIZE
|
|
298
|
-
header[0, 8] = MAGIC.b
|
|
299
|
-
|
|
300
|
-
HEADER_U32.each { |offset, value| put_u32(header, offset, value) }
|
|
301
|
-
META_U32.each { |offset, key| put_u32(header, offset, meta.fetch(key)) }
|
|
302
|
-
META_BOOL.each { |offset, key| put_u32(header, offset, meta.fetch(key) ? 1 : 0) }
|
|
303
|
-
|
|
304
|
-
put_u32(header, 24, vocab_size)
|
|
305
|
-
put_u32(header, 104, hash_size)
|
|
306
|
-
put_u32(header, MAX_TOKEN_CHARS_OFFSET, meta.fetch(:max_token_chars))
|
|
307
|
-
put_u32(header, MAX_PROBE_OFFSET, max_probe)
|
|
308
|
-
put_u32(header, ADDED_TOKEN_MASK_OFFSET, meta.fetch(:added_token_mask, 0))
|
|
309
|
-
put_prefix(header, meta.fetch(:subword_prefix))
|
|
310
|
-
SECTION_FIELDS.each { |name, field| put_section(header, field, sections.fetch(name)) }
|
|
311
|
-
|
|
312
|
-
header
|
|
313
|
-
end
|
|
314
|
-
|
|
315
|
-
def put_prefix(header, prefix)
|
|
316
|
-
bytes = prefix.b
|
|
317
|
-
raise ArgumentError, "subword prefix too long" if bytes.bytesize > 8
|
|
318
|
-
|
|
319
|
-
put_u32(header, 112, bytes.bytesize)
|
|
320
|
-
header[116, 8] = bytes.ljust(8, "\0")
|
|
321
|
-
end
|
|
322
|
-
|
|
323
|
-
def put_section(header, offset, section)
|
|
324
|
-
off, size = section
|
|
325
|
-
put_u64(header, offset, off)
|
|
326
|
-
put_u64(header, offset + 8, size)
|
|
327
|
-
end
|
|
328
|
-
|
|
329
|
-
def put_u32(buffer, offset, value)
|
|
330
|
-
buffer[offset, 4] = [value].pack("V")
|
|
331
|
-
end
|
|
332
|
-
|
|
333
|
-
def put_u64(buffer, offset, value)
|
|
334
|
-
buffer[offset, 8] = [value & 0xFFFFFFFF, value >> 32].pack("V2")
|
|
335
|
-
end
|
|
336
|
-
|
|
337
|
-
end
|
|
338
|
-
end
|
|
1
|
+
require "static_embeddings/format/constants"
|
|
2
|
+
require "static_embeddings/format/verifier"
|
|
3
|
+
require "static_embeddings/format/writer"
|
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
module StaticEmbeddings
|
|
2
|
+
module Importers
|
|
3
|
+
module Model2Vec
|
|
4
|
+
SOURCE_FILES = %w[tokenizer.json config.json tokenizer_config.json model.safetensors].freeze
|
|
5
|
+
DEFAULT_MAX_TOKENS = 512
|
|
6
|
+
FAMILY = "model2vec"
|
|
7
|
+
ORACLE = "model2vec.StaticModel"
|
|
8
|
+
|
|
9
|
+
module_function
|
|
10
|
+
|
|
11
|
+
def call(root, model_id: nil, max_tokens: nil, dimensions: nil,
|
|
12
|
+
source_revision: nil, trained_mrl_dims: nil)
|
|
13
|
+
tokenizer = Support.load_object(File.join(root, "tokenizer.json"))
|
|
14
|
+
config = Support.load_object(File.join(root, "config.json"), optional: true) || {}
|
|
15
|
+
tokenizer_config = Support.load_object(File.join(root, "tokenizer_config.json"), optional: true) || {}
|
|
16
|
+
profile, tokens = BertWordPiece.compile(tokenizer, tokenizer_config)
|
|
17
|
+
payload, native_dim = Support.extract_matrix(File.join(root, "model.safetensors"), tokens.length)
|
|
18
|
+
output_dim = Support.resolve_dimensions(dimensions, native_dim)
|
|
19
|
+
|
|
20
|
+
Canonical.model(
|
|
21
|
+
tokens: tokens,
|
|
22
|
+
matrix: Support.slice_matrix(payload, native_dim, output_dim),
|
|
23
|
+
dimensions: Canonical.dimensions(
|
|
24
|
+
native: native_dim,
|
|
25
|
+
output: output_dim,
|
|
26
|
+
trained: Support.normalize_mrl_dims(trained_mrl_dims, native_dim: native_dim)
|
|
27
|
+
),
|
|
28
|
+
runtime: Canonical.runtime(
|
|
29
|
+
normalization: normalization(config),
|
|
30
|
+
unk_policy: Format::UNK_DROP,
|
|
31
|
+
empty_policy: Format::EMPTY_ZERO_VECTOR,
|
|
32
|
+
max_tokens: Support.resolve_max_tokens(max_tokens, DEFAULT_MAX_TOKENS)
|
|
33
|
+
),
|
|
34
|
+
tokenizer: BertWordPiece.runtime_meta(profile, tokens),
|
|
35
|
+
source: Canonical.source(
|
|
36
|
+
family: FAMILY,
|
|
37
|
+
model: model_id || File.basename(root),
|
|
38
|
+
revision: source_revision,
|
|
39
|
+
oracle: ORACLE,
|
|
40
|
+
files_sha256: Support.digest_files(SOURCE_FILES.map { |name| File.join(root, name) }),
|
|
41
|
+
tokenizer_class: profile[:tokenizer_class],
|
|
42
|
+
config_seq_length: config["seq_length"]
|
|
43
|
+
)
|
|
44
|
+
)
|
|
45
|
+
end
|
|
46
|
+
|
|
47
|
+
def normalization(config)
|
|
48
|
+
config.fetch("normalize", false) ? Format::NORMALIZATION_L2 : Format::NORMALIZATION_NONE
|
|
49
|
+
end
|
|
50
|
+
end
|
|
51
|
+
end
|
|
52
|
+
end
|
|
@@ -0,0 +1,103 @@
|
|
|
1
|
+
require "pathname"
|
|
2
|
+
|
|
3
|
+
module StaticEmbeddings
|
|
4
|
+
module Importers
|
|
5
|
+
module SentenceTransformersStatic
|
|
6
|
+
TYPE = "sentence_transformers.models.StaticEmbedding"
|
|
7
|
+
FAMILY = "sentence_transformers_static"
|
|
8
|
+
ORACLE = "sentence_transformers.SentenceTransformer.encode"
|
|
9
|
+
DEFAULT_MAX_TOKENS = 0
|
|
10
|
+
|
|
11
|
+
module_function
|
|
12
|
+
|
|
13
|
+
def call(root, model_id: nil, max_tokens: nil, dimensions: nil,
|
|
14
|
+
source_revision: nil, trained_mrl_dims: nil)
|
|
15
|
+
module_dir = module_directory(root)
|
|
16
|
+
tokenizer = Support.load_object(File.join(module_dir, "tokenizer.json"))
|
|
17
|
+
tokenizer_config = Support.load_object(File.join(module_dir, "tokenizer_config.json"), optional: true) || {}
|
|
18
|
+
config = Support.load_object(File.join(module_dir, "config.json"), optional: true) || {}
|
|
19
|
+
profile, tokens = BertWordPiece.compile(tokenizer, tokenizer_config)
|
|
20
|
+
payload, native_dim = Support.extract_matrix(File.join(module_dir, "model.safetensors"), tokens.length)
|
|
21
|
+
output_dim = Support.resolve_dimensions(dimensions, native_dim)
|
|
22
|
+
|
|
23
|
+
Canonical.model(
|
|
24
|
+
tokens: tokens,
|
|
25
|
+
matrix: Support.slice_matrix(payload, native_dim, output_dim),
|
|
26
|
+
dimensions: Canonical.dimensions(
|
|
27
|
+
native: native_dim,
|
|
28
|
+
output: output_dim,
|
|
29
|
+
trained: Support.normalize_mrl_dims(trained_mrl_dims, native_dim: native_dim)
|
|
30
|
+
),
|
|
31
|
+
runtime: Canonical.runtime(
|
|
32
|
+
normalization: Format::NORMALIZATION_NONE,
|
|
33
|
+
unk_policy: Format::UNK_INCLUDE,
|
|
34
|
+
empty_policy: Format::EMPTY_ZERO_VECTOR,
|
|
35
|
+
max_tokens: Support.resolve_max_tokens(max_tokens, DEFAULT_MAX_TOKENS)
|
|
36
|
+
),
|
|
37
|
+
tokenizer: BertWordPiece.runtime_meta(profile, tokens),
|
|
38
|
+
source: Canonical.source(
|
|
39
|
+
family: FAMILY,
|
|
40
|
+
model: model_id || File.basename(root),
|
|
41
|
+
revision: source_revision,
|
|
42
|
+
oracle: ORACLE,
|
|
43
|
+
files_sha256: source_digests(root, module_dir),
|
|
44
|
+
tokenizer_class: profile[:tokenizer_class],
|
|
45
|
+
config_seq_length: config["seq_length"]
|
|
46
|
+
)
|
|
47
|
+
)
|
|
48
|
+
end
|
|
49
|
+
|
|
50
|
+
def module_directory(root)
|
|
51
|
+
spec = module_spec(root)
|
|
52
|
+
contained_directory(root, spec.fetch("path"))
|
|
53
|
+
end
|
|
54
|
+
|
|
55
|
+
def module_spec(root)
|
|
56
|
+
modules = Support.load_json(File.join(root, "modules.json"))
|
|
57
|
+
reject_source("modules.json is not an array") unless modules.is_a?(Array)
|
|
58
|
+
reject_source("modules.json declares #{modules.length} modules, expected exactly one StaticEmbedding") unless modules.length == 1
|
|
59
|
+
|
|
60
|
+
spec = modules.first
|
|
61
|
+
reject_source("modules.json[0] is not an object") unless spec.is_a?(Hash)
|
|
62
|
+
reject_source("modules.json[0].type is #{spec['type'].inspect}, expected #{TYPE}") unless spec["type"] == TYPE
|
|
63
|
+
reject_source("modules.json[0] has no path") unless spec.key?("path")
|
|
64
|
+
reject_source("modules.json[0].path is not a string") unless spec["path"].is_a?(String)
|
|
65
|
+
spec
|
|
66
|
+
end
|
|
67
|
+
|
|
68
|
+
def contained_directory(root, relative)
|
|
69
|
+
relative = relative.to_s
|
|
70
|
+
reject_source("module path contains a NUL byte") if relative.include?("\0")
|
|
71
|
+
reject_source("module path #{relative.inspect} is absolute") if Pathname.new(relative).absolute?
|
|
72
|
+
|
|
73
|
+
root_real = File.realpath(root)
|
|
74
|
+
candidate = relative.empty? ? root : File.expand_path(relative, root)
|
|
75
|
+
reject_source("module path #{relative.inspect} is not a directory") unless File.directory?(candidate)
|
|
76
|
+
|
|
77
|
+
real = File.realpath(candidate)
|
|
78
|
+
prefix = root_real.end_with?(File::SEPARATOR) ? root_real : "#{root_real}#{File::SEPARATOR}"
|
|
79
|
+
reject_source("module path #{relative.inspect} escapes the source directory") unless real == root_real || real.start_with?(prefix)
|
|
80
|
+
real
|
|
81
|
+
rescue Errno::ENOENT
|
|
82
|
+
reject_source("module path #{relative.inspect} does not exist")
|
|
83
|
+
end
|
|
84
|
+
|
|
85
|
+
def source_digests(root, module_dir)
|
|
86
|
+
paths = [
|
|
87
|
+
File.join(root, "modules.json"),
|
|
88
|
+
File.join(module_dir, "tokenizer.json"),
|
|
89
|
+
File.join(module_dir, "tokenizer_config.json"),
|
|
90
|
+
File.join(module_dir, "config.json"),
|
|
91
|
+
File.join(module_dir, "model.safetensors")
|
|
92
|
+
]
|
|
93
|
+
Support.digest_files(paths, root: root)
|
|
94
|
+
end
|
|
95
|
+
|
|
96
|
+
def reject_source(message)
|
|
97
|
+
raise UnsupportedModelError,
|
|
98
|
+
"#{message}. Sentence Transformers conversion accepts exactly one #{TYPE} module " \
|
|
99
|
+
"using #{BertWordPiece::TOKENIZER_PROFILE}."
|
|
100
|
+
end
|
|
101
|
+
end
|
|
102
|
+
end
|
|
103
|
+
end
|