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.
Files changed (35) hide show
  1. checksums.yaml +4 -4
  2. data/CHANGELOG.md +50 -0
  3. data/README.md +17 -2
  4. data/docs/ARCHITECTURE.md +14 -5
  5. data/docs/MODEL_AUDIT.md +110 -0
  6. data/lib/models/demo.semb +0 -0
  7. data/lib/static_embeddings/bert_wordpiece.rb +191 -0
  8. data/lib/static_embeddings/canonical.rb +50 -0
  9. data/lib/static_embeddings/cli.rb +88 -63
  10. data/lib/static_embeddings/codec.rb +45 -0
  11. data/lib/static_embeddings/conversion.rb +58 -0
  12. data/lib/static_embeddings/errors.rb +2 -2
  13. data/lib/static_embeddings/format/constants.rb +109 -0
  14. data/lib/static_embeddings/format/hash_table.rb +69 -0
  15. data/lib/static_embeddings/format/trie.rb +78 -0
  16. data/lib/static_embeddings/format/verifier.rb +41 -0
  17. data/lib/static_embeddings/format/writer.rb +131 -0
  18. data/lib/static_embeddings/format.rb +3 -338
  19. data/lib/static_embeddings/importers/model2vec.rb +52 -0
  20. data/lib/static_embeddings/importers/sentence_transformers_static.rb +103 -0
  21. data/lib/static_embeddings/importers/support.rb +111 -0
  22. data/lib/static_embeddings/importers.rb +50 -0
  23. data/lib/static_embeddings/model.rb +35 -20
  24. data/lib/static_embeddings/paths.rb +8 -12
  25. data/lib/static_embeddings/provenance.rb +58 -0
  26. data/lib/static_embeddings/reference.rb +50 -40
  27. data/lib/static_embeddings/row_prefix_payload.rb +59 -0
  28. data/lib/static_embeddings/version.rb +1 -1
  29. data/lib/static_embeddings.rb +29 -55
  30. data/static_embeddings.gemspec +2 -2
  31. data/tools/check_model2vec_parity.rb +5 -1
  32. data/tools/check_st_parity.rb +125 -0
  33. data/tools/eval_retrieval.rb +58 -0
  34. metadata +24 -6
  35. 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 "digest"
2
-
3
- module StaticEmbeddings
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