static_embeddings 0.1.4 → 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 +118 -0
- data/README.md +67 -25
- data/Rakefile +1 -1
- data/docs/ARCHITECTURE.md +54 -34
- data/docs/LIMITATIONS.md +12 -6
- data/docs/MODEL_AUDIT.md +195 -52
- data/ext/static_embeddings/se_embed.c +2 -1
- data/ext/static_embeddings/se_format.c +41 -4
- data/ext/static_embeddings/se_internal.h +17 -5
- data/ext/static_embeddings/se_tokenizer.c +155 -18
- data/ext/static_embeddings/se_unicode.c +1 -1
- data/ext/static_embeddings/static_embeddings.c +32 -2
- 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 -62
- 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 -300
- 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 +17 -4
- data/lib/static_embeddings/provenance.rb +58 -0
- data/lib/static_embeddings/reference.rb +90 -33
- data/lib/static_embeddings/row_prefix_payload.rb +59 -0
- data/lib/static_embeddings/safetensors.rb +178 -34
- data/lib/static_embeddings/version.rb +1 -1
- data/lib/static_embeddings.rb +29 -57
- data/static_embeddings.gemspec +2 -2
- data/tools/check_model2vec_parity.rb +89 -54
- 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 -284
|
@@ -1,16 +1,16 @@
|
|
|
1
|
+
require "fileutils"
|
|
1
2
|
require "json"
|
|
2
3
|
require "optparse"
|
|
3
|
-
require "fileutils"
|
|
4
4
|
|
|
5
5
|
module StaticEmbeddings
|
|
6
|
-
|
|
6
|
+
module CLI
|
|
7
7
|
COMMANDS = {
|
|
8
|
-
"convert" => :
|
|
9
|
-
"verify" => :
|
|
10
|
-
"inspect" => :
|
|
11
|
-
"tokenize" => :
|
|
12
|
-
"embed" => :
|
|
13
|
-
"cache-path" => :
|
|
8
|
+
"convert" => :convert_command,
|
|
9
|
+
"verify" => :verify_command,
|
|
10
|
+
"inspect" => :inspect_command,
|
|
11
|
+
"tokenize" => :tokenize_command,
|
|
12
|
+
"embed" => :embed_command,
|
|
13
|
+
"cache-path" => :cache_path_command,
|
|
14
14
|
"help" => :usage,
|
|
15
15
|
"-h" => :usage,
|
|
16
16
|
"--help" => :usage,
|
|
@@ -20,10 +20,14 @@ module StaticEmbeddings
|
|
|
20
20
|
HELP = <<~TEXT
|
|
21
21
|
static_embeddings <command> [options]
|
|
22
22
|
|
|
23
|
-
convert SOURCE_DIR Convert
|
|
23
|
+
convert SOURCE_DIR Convert Model2Vec or Sentence Transformers StaticEmbedding to .semb
|
|
24
24
|
--out PATH Output file (default: <cache>/models/<id>.semb)
|
|
25
25
|
--id ID Model id recorded in provenance
|
|
26
|
-
--max-tokens N
|
|
26
|
+
--max-tokens N Model2Vec default: 512; Sentence Transformers default: unlimited
|
|
27
|
+
--max-tokens unlimited
|
|
28
|
+
--dimensions N Keep the first N dims (Matryoshka prefix slice)
|
|
29
|
+
--revision SHA Source revision recorded in provenance
|
|
30
|
+
--trained-mrl-dims 1024,512,256,... training dims recorded in provenance
|
|
27
31
|
|
|
28
32
|
verify PATH Recompute the SHA-256 embedded in the header
|
|
29
33
|
inspect PATH Print header fields and provenance
|
|
@@ -32,66 +36,78 @@ module StaticEmbeddings
|
|
|
32
36
|
cache-path Print the model cache directory
|
|
33
37
|
TEXT
|
|
34
38
|
|
|
35
|
-
|
|
36
|
-
new.run(argv)
|
|
37
|
-
end
|
|
39
|
+
module_function
|
|
38
40
|
|
|
39
41
|
def run(argv)
|
|
40
|
-
|
|
41
|
-
return
|
|
42
|
+
command = COMMANDS[argv.shift]
|
|
43
|
+
return unknown_command unless command
|
|
42
44
|
|
|
43
|
-
|
|
44
|
-
rescue StaticEmbeddings::Error => e
|
|
45
|
+
public_send(command, argv)
|
|
46
|
+
rescue StaticEmbeddings::Error, ArgumentError, OptionParser::ParseError => e
|
|
45
47
|
warn "#{e.class.name.split('::').last}: #{e.message}"
|
|
46
48
|
1
|
|
47
49
|
end
|
|
48
50
|
|
|
49
|
-
private
|
|
50
|
-
|
|
51
51
|
def usage(*)
|
|
52
52
|
puts HELP
|
|
53
53
|
0
|
|
54
54
|
end
|
|
55
55
|
|
|
56
|
-
def
|
|
56
|
+
def unknown_command
|
|
57
57
|
warn "unknown command"
|
|
58
58
|
usage
|
|
59
59
|
1
|
|
60
60
|
end
|
|
61
61
|
|
|
62
|
-
def
|
|
62
|
+
def cache_path_command(*)
|
|
63
63
|
puts StaticEmbeddings.cache_dir
|
|
64
64
|
0
|
|
65
65
|
end
|
|
66
66
|
|
|
67
|
-
def
|
|
68
|
-
options =
|
|
67
|
+
def convert_command(argv)
|
|
68
|
+
options = convert_options(argv)
|
|
69
69
|
source = required_arg(argv, "usage: static_embeddings convert SOURCE_DIR [--out PATH]")
|
|
70
70
|
model_id = options[:id] || File.basename(File.expand_path(source))
|
|
71
|
-
|
|
72
|
-
FileUtils.mkdir_p(File.dirname(
|
|
71
|
+
output = options[:out] || StaticEmbeddings.model_path(model_id)
|
|
72
|
+
FileUtils.mkdir_p(File.dirname(output))
|
|
73
73
|
|
|
74
|
-
report = StaticEmbeddings.convert(source, output_path:
|
|
75
|
-
|
|
76
|
-
puts conversion_report(out, report, options[:max_tokens])
|
|
74
|
+
report = StaticEmbeddings.convert(source, output_path: output, model_id: model_id, **conversion_options(options))
|
|
75
|
+
puts conversion_report(output, report)
|
|
77
76
|
0
|
|
78
77
|
end
|
|
79
78
|
|
|
80
|
-
def
|
|
81
|
-
|
|
82
|
-
|
|
83
|
-
|
|
84
|
-
|
|
85
|
-
|
|
86
|
-
|
|
87
|
-
|
|
79
|
+
def convert_options(argv)
|
|
80
|
+
{}.tap do |options|
|
|
81
|
+
OptionParser.new do |parser|
|
|
82
|
+
parser.on("--out PATH") { |value| options[:out] = value }
|
|
83
|
+
parser.on("--id ID") { |value| options[:id] = value }
|
|
84
|
+
parser.on("--max-tokens N") { |value| options[:max_tokens] = parse_max_tokens(value) }
|
|
85
|
+
parser.on("--dimensions N", Integer) { |value| options[:dimensions] = value }
|
|
86
|
+
parser.on("--revision SHA") { |value| options[:source_revision] = value }
|
|
87
|
+
parser.on("--trained-mrl-dims LIST") { |value| options[:trained_mrl_dims] = value }
|
|
88
|
+
end.parse!(argv)
|
|
89
|
+
end
|
|
88
90
|
end
|
|
89
91
|
|
|
90
|
-
def
|
|
92
|
+
def conversion_options(options)
|
|
93
|
+
options.select { |key, _| %i[max_tokens dimensions source_revision trained_mrl_dims].include?(key) }
|
|
94
|
+
end
|
|
95
|
+
|
|
96
|
+
def parse_max_tokens(value)
|
|
97
|
+
return :unlimited if value == "unlimited" || value == "0"
|
|
98
|
+
|
|
99
|
+
integer = Integer(value)
|
|
100
|
+
raise OptionParser::InvalidArgument, "--max-tokens must be positive or unlimited" unless integer.positive?
|
|
101
|
+
integer
|
|
102
|
+
end
|
|
103
|
+
|
|
104
|
+
def conversion_report(path, report)
|
|
105
|
+
max_tokens = report[:max_tokens].to_i.zero? ? "unlimited" : report[:max_tokens]
|
|
106
|
+
dim = report[:native_dim] == report[:dim] ? report[:dim].to_s : "#{report[:dim]} (from native #{report[:native_dim]})"
|
|
91
107
|
[
|
|
92
108
|
"wrote #{path}",
|
|
93
109
|
" vocab #{report[:vocab_size]}",
|
|
94
|
-
" dim #{
|
|
110
|
+
" dim #{dim}",
|
|
95
111
|
" bytes #{report[:bytes]}",
|
|
96
112
|
" sha256 #{report[:sha256]}",
|
|
97
113
|
" max_tokens #{max_tokens}",
|
|
@@ -100,20 +116,25 @@ module StaticEmbeddings
|
|
|
100
116
|
].join("\n")
|
|
101
117
|
end
|
|
102
118
|
|
|
103
|
-
def
|
|
119
|
+
def verify_command(argv)
|
|
104
120
|
result = StaticEmbeddings.verify(required_arg(argv, "usage: static_embeddings verify PATH"))
|
|
105
|
-
|
|
106
|
-
|
|
107
|
-
|
|
108
|
-
|
|
109
|
-
|
|
110
|
-
|
|
111
|
-
|
|
112
|
-
|
|
113
|
-
|
|
121
|
+
if result[:ok]
|
|
122
|
+
puts "ok #{result[:expected]}"
|
|
123
|
+
0
|
|
124
|
+
else
|
|
125
|
+
warn "CHECKSUM MISMATCH"
|
|
126
|
+
warn " stored #{result[:stored]}"
|
|
127
|
+
warn " computed #{result[:expected]}"
|
|
128
|
+
1
|
|
129
|
+
end
|
|
130
|
+
end
|
|
131
|
+
|
|
132
|
+
def inspect_command(argv)
|
|
114
133
|
model = StaticEmbeddings.load(required_arg(argv, "usage: static_embeddings inspect PATH"))
|
|
115
134
|
puts JSON.pretty_generate(model_summary(model))
|
|
116
135
|
0
|
|
136
|
+
ensure
|
|
137
|
+
model&.close
|
|
117
138
|
end
|
|
118
139
|
|
|
119
140
|
def model_summary(model)
|
|
@@ -130,18 +151,20 @@ module StaticEmbeddings
|
|
|
130
151
|
}
|
|
131
152
|
end
|
|
132
153
|
|
|
133
|
-
def
|
|
134
|
-
|
|
135
|
-
|
|
136
|
-
|
|
154
|
+
def tokenize_command(argv)
|
|
155
|
+
with_model_and_text(argv, "usage: static_embeddings tokenize PATH TEXT") do |model, text|
|
|
156
|
+
ids = model.tokenize(text)
|
|
157
|
+
puts JSON.generate("ids" => ids, "count" => ids.length, "unk" => ids.count(model.unk_id))
|
|
158
|
+
end
|
|
137
159
|
0
|
|
138
160
|
end
|
|
139
161
|
|
|
140
|
-
def
|
|
141
|
-
|
|
142
|
-
|
|
143
|
-
|
|
144
|
-
|
|
162
|
+
def embed_command(argv)
|
|
163
|
+
with_model_and_text(argv, "usage: static_embeddings embed PATH TEXT") do |model, text|
|
|
164
|
+
stats = model.embed_with_stats(text)
|
|
165
|
+
warn_high_unk(stats) if high_unk?(stats)
|
|
166
|
+
puts JSON.generate(stats_payload(model, stats))
|
|
167
|
+
end
|
|
145
168
|
0
|
|
146
169
|
end
|
|
147
170
|
|
|
@@ -150,7 +173,7 @@ module StaticEmbeddings
|
|
|
150
173
|
"token_count" => stats[:token_count],
|
|
151
174
|
"unk_count" => stats[:unk_count],
|
|
152
175
|
"truncated" => stats[:truncated],
|
|
153
|
-
"vector" => StaticEmbeddings.unpack(stats[:vector], model.dim).first.map { |
|
|
176
|
+
"vector" => StaticEmbeddings.unpack(stats[:vector], model.dim).first.map { |value| value.round(6) }
|
|
154
177
|
}
|
|
155
178
|
end
|
|
156
179
|
|
|
@@ -163,16 +186,19 @@ module StaticEmbeddings
|
|
|
163
186
|
warn "warning: #{(ratio * 100).round}% of tokens are [UNK] — wrong model for this language?"
|
|
164
187
|
end
|
|
165
188
|
|
|
166
|
-
def
|
|
189
|
+
def with_model_and_text(argv, usage)
|
|
167
190
|
path = argv.shift
|
|
168
191
|
text = argv.join(" ")
|
|
169
|
-
|
|
192
|
+
raise InvalidOptionError, usage if path.nil? || text.empty?
|
|
170
193
|
|
|
171
|
-
|
|
194
|
+
model = StaticEmbeddings.load(path)
|
|
195
|
+
yield model, text
|
|
196
|
+
ensure
|
|
197
|
+
model&.close
|
|
172
198
|
end
|
|
173
199
|
|
|
174
200
|
def required_arg(argv, usage)
|
|
175
|
-
argv.shift ||
|
|
201
|
+
argv.shift || raise(InvalidOptionError, usage)
|
|
176
202
|
end
|
|
177
203
|
end
|
|
178
204
|
end
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
module StaticEmbeddings
|
|
2
|
+
module Codec
|
|
3
|
+
module_function
|
|
4
|
+
|
|
5
|
+
def unpack(blob, dim, format: :f32)
|
|
6
|
+
dim = Integer(dim)
|
|
7
|
+
raise InvalidOptionError, "dim must be positive" unless dim.positive?
|
|
8
|
+
|
|
9
|
+
values = decode(blob, normalize_format(format))
|
|
10
|
+
raise InvalidOptionError, "blob is not a multiple of dim" unless (values.length % dim).zero?
|
|
11
|
+
|
|
12
|
+
values.each_slice(dim).to_a
|
|
13
|
+
end
|
|
14
|
+
|
|
15
|
+
def pack(rows, format: :f32)
|
|
16
|
+
values = rows.first.is_a?(Array) ? rows.flatten(1) : rows
|
|
17
|
+
floats = values.map(&:to_f)
|
|
18
|
+
|
|
19
|
+
case normalize_format(format)
|
|
20
|
+
when :f32 then floats.pack("e*")
|
|
21
|
+
when :f16 then StaticEmbeddings.encode_f16(floats)
|
|
22
|
+
end
|
|
23
|
+
end
|
|
24
|
+
|
|
25
|
+
def decode(blob, format)
|
|
26
|
+
case format
|
|
27
|
+
when :f32
|
|
28
|
+
raise InvalidOptionError, "f32 blob byte size must be a multiple of 4" unless (blob.bytesize % 4).zero?
|
|
29
|
+
blob.unpack("e*")
|
|
30
|
+
when :f16
|
|
31
|
+
raise InvalidOptionError, "f16 blob byte size must be a multiple of 2" unless (blob.bytesize % 2).zero?
|
|
32
|
+
StaticEmbeddings.decode_f16(blob)
|
|
33
|
+
end
|
|
34
|
+
end
|
|
35
|
+
|
|
36
|
+
def normalize_format(format)
|
|
37
|
+
case format&.to_sym
|
|
38
|
+
when nil, :f32, :float32 then :f32
|
|
39
|
+
when :f16, :float16 then :f16
|
|
40
|
+
else
|
|
41
|
+
raise InvalidOptionError, "unsupported embedding format #{format.inspect} (expected :f32 or :f16)"
|
|
42
|
+
end
|
|
43
|
+
end
|
|
44
|
+
end
|
|
45
|
+
end
|
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
require "json"
|
|
2
|
+
require "static_embeddings/format/writer"
|
|
3
|
+
require "static_embeddings/importers"
|
|
4
|
+
require "static_embeddings/provenance"
|
|
5
|
+
require "static_embeddings/unicode_tables"
|
|
6
|
+
|
|
7
|
+
module StaticEmbeddings
|
|
8
|
+
module Conversion
|
|
9
|
+
module_function
|
|
10
|
+
|
|
11
|
+
def call(source_dir, output_path:, **options)
|
|
12
|
+
model = Importers.import(source_dir, **options)
|
|
13
|
+
provenance = Provenance.json(model, unicode_source: UnicodeTables.source_stamp)
|
|
14
|
+
result = Format::Writer.call(
|
|
15
|
+
path: output_path,
|
|
16
|
+
meta: format_meta(model),
|
|
17
|
+
tokens: model.tokens,
|
|
18
|
+
matrix: model.matrix,
|
|
19
|
+
norm_tables: UnicodeTables.packed,
|
|
20
|
+
provenance: provenance
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
result.merge(
|
|
24
|
+
vocab_size: model.tokens.length,
|
|
25
|
+
dim: model.dimensions.output,
|
|
26
|
+
native_dim: model.dimensions.native,
|
|
27
|
+
max_tokens: model.runtime.max_tokens,
|
|
28
|
+
provenance: JSON.parse(provenance)
|
|
29
|
+
)
|
|
30
|
+
end
|
|
31
|
+
|
|
32
|
+
def format_meta(model)
|
|
33
|
+
tokenizer = model.tokenizer
|
|
34
|
+
runtime = model.runtime
|
|
35
|
+
{
|
|
36
|
+
dim: model.dimensions.output,
|
|
37
|
+
normalization_type: runtime.normalization,
|
|
38
|
+
max_tokens_default: runtime.max_tokens,
|
|
39
|
+
add_special_tokens: runtime.add_special_tokens,
|
|
40
|
+
unk_policy: runtime.unk_policy,
|
|
41
|
+
empty_policy: runtime.empty_policy,
|
|
42
|
+
do_lower_case: tokenizer.fetch(:do_lower_case),
|
|
43
|
+
strip_accents: tokenizer.fetch(:strip_accents),
|
|
44
|
+
handle_chinese_chars: tokenizer.fetch(:handle_chinese_chars),
|
|
45
|
+
clean_text: tokenizer.fetch(:clean_text),
|
|
46
|
+
added_token_mask: tokenizer.fetch(:added_token_mask),
|
|
47
|
+
max_input_chars_per_word: tokenizer.fetch(:max_input_chars_per_word),
|
|
48
|
+
max_token_chars: tokenizer.fetch(:max_token_chars),
|
|
49
|
+
subword_prefix: tokenizer.fetch(:subword_prefix),
|
|
50
|
+
pad_id: tokenizer.fetch(:pad_id),
|
|
51
|
+
unk_id: tokenizer.fetch(:unk_id),
|
|
52
|
+
cls_id: tokenizer.fetch(:cls_id),
|
|
53
|
+
sep_id: tokenizer.fetch(:sep_id),
|
|
54
|
+
mask_id: tokenizer.fetch(:mask_id)
|
|
55
|
+
}
|
|
56
|
+
end
|
|
57
|
+
end
|
|
58
|
+
end
|
|
@@ -0,0 +1,109 @@
|
|
|
1
|
+
module StaticEmbeddings
|
|
2
|
+
module Format
|
|
3
|
+
MAGIC = "SEMBv1\0\0"
|
|
4
|
+
VERSION = 3
|
|
5
|
+
VERSION_WORDPIECE = VERSION
|
|
6
|
+
HEADER_SIZE = 320
|
|
7
|
+
ALIGNMENT = 64
|
|
8
|
+
|
|
9
|
+
TOKENIZER_BERT_WORDPIECE_V1 = 1
|
|
10
|
+
DTYPE_F32 = 1
|
|
11
|
+
POOLING_MEAN = 1
|
|
12
|
+
NORMALIZATION_NONE = 0
|
|
13
|
+
NORMALIZATION_L2 = 1
|
|
14
|
+
TRUNCATE_USABLE_IDS_BEFORE_POOLING = 2
|
|
15
|
+
UNK_INCLUDE = 0
|
|
16
|
+
UNK_DROP = 1
|
|
17
|
+
EMPTY_ZERO_VECTOR = 0
|
|
18
|
+
EMPTY_RAISE = 1
|
|
19
|
+
|
|
20
|
+
SLOT_EMPTY = 0xFFFFFFFF
|
|
21
|
+
HASH_SEED = 2_166_136_261
|
|
22
|
+
LOAD_FACTOR = 0.70
|
|
23
|
+
UINT32_MAX = 0xFFFFFFFF
|
|
24
|
+
UINT64_MAX = 0xFFFFFFFFFFFFFFFF
|
|
25
|
+
|
|
26
|
+
MAX_TOKEN_CHARS_OFFSET = 124
|
|
27
|
+
CHECKSUM_OFFSET = 240
|
|
28
|
+
CHECKSUM_SIZE = 32
|
|
29
|
+
MAX_PROBE_OFFSET = 304
|
|
30
|
+
ADDED_TOKEN_MASK_OFFSET = 308
|
|
31
|
+
|
|
32
|
+
ADDED_PAD = 1 << 0
|
|
33
|
+
ADDED_UNK = 1 << 1
|
|
34
|
+
ADDED_CLS = 1 << 2
|
|
35
|
+
ADDED_SEP = 1 << 3
|
|
36
|
+
ADDED_MASK = 1 << 4
|
|
37
|
+
ADDED_TOKEN_MASK_ALL = ADDED_PAD | ADDED_UNK | ADDED_CLS | ADDED_SEP | ADDED_MASK
|
|
38
|
+
|
|
39
|
+
SECTION_FIELDS = {
|
|
40
|
+
vocab_strings: 128,
|
|
41
|
+
vocab_hash: 144,
|
|
42
|
+
embeddings: 160,
|
|
43
|
+
norm_tables: 176,
|
|
44
|
+
provenance: 192,
|
|
45
|
+
root_trie: 208,
|
|
46
|
+
continuation_trie: 224
|
|
47
|
+
}.freeze
|
|
48
|
+
|
|
49
|
+
HEADER_U32 = {
|
|
50
|
+
8 => VERSION,
|
|
51
|
+
12 => HEADER_SIZE,
|
|
52
|
+
16 => 1,
|
|
53
|
+
28 => TOKENIZER_BERT_WORDPIECE_V1,
|
|
54
|
+
32 => DTYPE_F32,
|
|
55
|
+
36 => POOLING_MEAN,
|
|
56
|
+
48 => TRUNCATE_USABLE_IDS_BEFORE_POOLING,
|
|
57
|
+
108 => HASH_SEED
|
|
58
|
+
}.freeze
|
|
59
|
+
|
|
60
|
+
META_U32 = {
|
|
61
|
+
20 => :dim,
|
|
62
|
+
40 => :normalization_type,
|
|
63
|
+
44 => :max_tokens_default,
|
|
64
|
+
56 => :unk_policy,
|
|
65
|
+
60 => :empty_policy,
|
|
66
|
+
80 => :max_input_chars_per_word,
|
|
67
|
+
84 => :pad_id,
|
|
68
|
+
88 => :unk_id,
|
|
69
|
+
92 => :cls_id,
|
|
70
|
+
96 => :sep_id,
|
|
71
|
+
100 => :mask_id
|
|
72
|
+
}.freeze
|
|
73
|
+
|
|
74
|
+
META_BOOL = {
|
|
75
|
+
52 => :add_special_tokens,
|
|
76
|
+
64 => :do_lower_case,
|
|
77
|
+
68 => :strip_accents,
|
|
78
|
+
72 => :handle_chinese_chars,
|
|
79
|
+
76 => :clean_text
|
|
80
|
+
}.freeze
|
|
81
|
+
|
|
82
|
+
module_function
|
|
83
|
+
|
|
84
|
+
def binary_string(capacity = nil)
|
|
85
|
+
string = capacity ? String.new(capacity: capacity) : +""
|
|
86
|
+
string.force_encoding(Encoding::BINARY)
|
|
87
|
+
end
|
|
88
|
+
|
|
89
|
+
def verify(path)
|
|
90
|
+
require "static_embeddings/format/verifier"
|
|
91
|
+
Verifier.call(path)
|
|
92
|
+
end
|
|
93
|
+
|
|
94
|
+
def write(**kwargs)
|
|
95
|
+
require "static_embeddings/format/writer"
|
|
96
|
+
Writer.call(**kwargs)
|
|
97
|
+
end
|
|
98
|
+
|
|
99
|
+
def hash_bytes(string, seed = HASH_SEED)
|
|
100
|
+
require "static_embeddings/format/hash_table"
|
|
101
|
+
HashTable.hash_bytes(string, seed)
|
|
102
|
+
end
|
|
103
|
+
|
|
104
|
+
def next_power_of_two(value)
|
|
105
|
+
require "static_embeddings/format/hash_table"
|
|
106
|
+
HashTable.next_power_of_two(value)
|
|
107
|
+
end
|
|
108
|
+
end
|
|
109
|
+
end
|
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
require "static_embeddings/format/constants"
|
|
2
|
+
|
|
3
|
+
module StaticEmbeddings
|
|
4
|
+
module Format
|
|
5
|
+
module HashTable
|
|
6
|
+
module_function
|
|
7
|
+
|
|
8
|
+
def build(tokens)
|
|
9
|
+
size = next_power_of_two((tokens.length / LOAD_FACTOR).ceil + 1)
|
|
10
|
+
slots = Array.new(size)
|
|
11
|
+
strings = Format.binary_string
|
|
12
|
+
max_probe = 0
|
|
13
|
+
|
|
14
|
+
tokens.each_with_index do |token, id|
|
|
15
|
+
bytes = token.b
|
|
16
|
+
offset = strings.bytesize
|
|
17
|
+
strings << bytes
|
|
18
|
+
probe = insert(slots, tokens, size, bytes, hash_bytes(bytes), offset, id)
|
|
19
|
+
max_probe = probe if probe > max_probe
|
|
20
|
+
end
|
|
21
|
+
|
|
22
|
+
[size, strings, pack(slots), max_probe]
|
|
23
|
+
end
|
|
24
|
+
|
|
25
|
+
def hash_bytes(string, seed = HASH_SEED)
|
|
26
|
+
string.each_byte.reduce(seed) { |hash, byte| ((hash ^ byte) * 16_777_619) & UINT32_MAX }
|
|
27
|
+
end
|
|
28
|
+
|
|
29
|
+
def next_power_of_two(value)
|
|
30
|
+
1 << (value - 1).bit_length
|
|
31
|
+
end
|
|
32
|
+
|
|
33
|
+
def insert(slots, tokens, size, bytes, hash, offset, id)
|
|
34
|
+
position = hash & (size - 1)
|
|
35
|
+
probe = 1
|
|
36
|
+
|
|
37
|
+
loop do
|
|
38
|
+
slot = slots[position]
|
|
39
|
+
unless slot
|
|
40
|
+
slots[position] = [hash, offset, bytes.bytesize, id]
|
|
41
|
+
return probe
|
|
42
|
+
end
|
|
43
|
+
|
|
44
|
+
if slot[0] == hash && slot[2] == bytes.bytesize && tokens.fetch(slot[3]).b == bytes
|
|
45
|
+
raise ArgumentError, "duplicate token in vocabulary: #{tokens.fetch(slot[3]).inspect}"
|
|
46
|
+
end
|
|
47
|
+
|
|
48
|
+
position = (position + 1) & (size - 1)
|
|
49
|
+
probe += 1
|
|
50
|
+
end
|
|
51
|
+
end
|
|
52
|
+
|
|
53
|
+
def pack(slots)
|
|
54
|
+
empty = [0, 0, 0, SLOT_EMPTY].pack("V4")
|
|
55
|
+
slots.each_with_object(Format.binary_string(slots.length * 16)) do |slot, packed|
|
|
56
|
+
packed << (slot ? slot.pack("V4") : empty)
|
|
57
|
+
end
|
|
58
|
+
end
|
|
59
|
+
end
|
|
60
|
+
|
|
61
|
+
def self.hash_bytes(string, seed = HASH_SEED)
|
|
62
|
+
HashTable.hash_bytes(string, seed)
|
|
63
|
+
end
|
|
64
|
+
|
|
65
|
+
def self.next_power_of_two(value)
|
|
66
|
+
HashTable.next_power_of_two(value)
|
|
67
|
+
end
|
|
68
|
+
end
|
|
69
|
+
end
|
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
require "static_embeddings/format/constants"
|
|
2
|
+
|
|
3
|
+
module StaticEmbeddings
|
|
4
|
+
module Format
|
|
5
|
+
class Trie
|
|
6
|
+
Node = Struct.new(:terminal, :children, keyword_init: true)
|
|
7
|
+
|
|
8
|
+
def initialize
|
|
9
|
+
@nodes = [Node.new(terminal: SLOT_EMPTY, children: {})]
|
|
10
|
+
end
|
|
11
|
+
|
|
12
|
+
def insert(bytes, id)
|
|
13
|
+
return if bytes.empty?
|
|
14
|
+
|
|
15
|
+
node_index = 0
|
|
16
|
+
bytes.each_byte do |byte|
|
|
17
|
+
node = @nodes[node_index]
|
|
18
|
+
node_index = node.children[byte] ||= append_node
|
|
19
|
+
end
|
|
20
|
+
|
|
21
|
+
node = @nodes[node_index]
|
|
22
|
+
raise ArgumentError, "duplicate trie key for token id #{id}" unless node.terminal == SLOT_EMPTY
|
|
23
|
+
|
|
24
|
+
node.terminal = id
|
|
25
|
+
end
|
|
26
|
+
|
|
27
|
+
def pack
|
|
28
|
+
edges = []
|
|
29
|
+
nodes = @nodes.map do |node|
|
|
30
|
+
start = edges.length
|
|
31
|
+
node.children.sort.each { |byte, child| edges << [byte, child] }
|
|
32
|
+
[start, node.children.length, node.terminal, 0]
|
|
33
|
+
end
|
|
34
|
+
|
|
35
|
+
nodes.each_with_object(header(nodes.length, edges.length)) { |record, out| out << record.pack("V4") }
|
|
36
|
+
.tap { |out| edges.each { |edge| out << edge.pack("V2") } }
|
|
37
|
+
end
|
|
38
|
+
|
|
39
|
+
private
|
|
40
|
+
|
|
41
|
+
def append_node
|
|
42
|
+
@nodes << Node.new(terminal: SLOT_EMPTY, children: {})
|
|
43
|
+
@nodes.length - 1
|
|
44
|
+
end
|
|
45
|
+
|
|
46
|
+
def header(node_count, edge_count)
|
|
47
|
+
capacity = 16 + node_count * 16 + edge_count * 8
|
|
48
|
+
Format.binary_string(capacity).tap { |out| out << [node_count, edge_count, 0, 0].pack("V4") }
|
|
49
|
+
end
|
|
50
|
+
end
|
|
51
|
+
|
|
52
|
+
module WordPieceTrie
|
|
53
|
+
module_function
|
|
54
|
+
|
|
55
|
+
def build(tokens, prefix)
|
|
56
|
+
root = Trie.new
|
|
57
|
+
continuation = Trie.new
|
|
58
|
+
prefix_bytes = prefix.b
|
|
59
|
+
|
|
60
|
+
tokens.each_with_index do |token, id|
|
|
61
|
+
bytes = token.b
|
|
62
|
+
if continuation?(bytes, prefix_bytes)
|
|
63
|
+
body = bytes.byteslice(prefix_bytes.bytesize, bytes.bytesize - prefix_bytes.bytesize)
|
|
64
|
+
continuation.insert(body, id)
|
|
65
|
+
else
|
|
66
|
+
root.insert(bytes, id)
|
|
67
|
+
end
|
|
68
|
+
end
|
|
69
|
+
|
|
70
|
+
[root.pack, continuation.pack]
|
|
71
|
+
end
|
|
72
|
+
|
|
73
|
+
def continuation?(bytes, prefix)
|
|
74
|
+
!prefix.empty? && bytes.start_with?(prefix) && bytes.bytesize > prefix.bytesize
|
|
75
|
+
end
|
|
76
|
+
end
|
|
77
|
+
end
|
|
78
|
+
end
|
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
require "digest"
|
|
2
|
+
require "static_embeddings/errors"
|
|
3
|
+
require "static_embeddings/format/constants"
|
|
4
|
+
|
|
5
|
+
module StaticEmbeddings
|
|
6
|
+
module Format
|
|
7
|
+
module Verifier
|
|
8
|
+
CHUNK_BYTES = 1024 * 1024
|
|
9
|
+
|
|
10
|
+
module_function
|
|
11
|
+
|
|
12
|
+
def call(path)
|
|
13
|
+
raise InvalidModelError, "file too small" if File.size(path) < HEADER_SIZE
|
|
14
|
+
|
|
15
|
+
File.open(path, "rb") do |io|
|
|
16
|
+
header = io.read(HEADER_SIZE)
|
|
17
|
+
raise InvalidModelError, "bad magic" unless header.byteslice(0, 8) == MAGIC.b
|
|
18
|
+
|
|
19
|
+
stored = header.byteslice(CHECKSUM_OFFSET, CHECKSUM_SIZE)
|
|
20
|
+
actual = checksum(io, header)
|
|
21
|
+
{ ok: stored == actual, expected: actual.unpack1("H*"), stored: stored.unpack1("H*") }
|
|
22
|
+
end
|
|
23
|
+
end
|
|
24
|
+
|
|
25
|
+
def checksum(io, header)
|
|
26
|
+
digest = Digest::SHA256.new
|
|
27
|
+
zeroed = header.dup
|
|
28
|
+
zeroed[CHECKSUM_OFFSET, CHECKSUM_SIZE] = "\0".b * CHECKSUM_SIZE
|
|
29
|
+
digest << zeroed
|
|
30
|
+
|
|
31
|
+
buffer = String.new(capacity: CHUNK_BYTES)
|
|
32
|
+
digest << buffer while io.read(CHUNK_BYTES, buffer)
|
|
33
|
+
digest.digest
|
|
34
|
+
end
|
|
35
|
+
end
|
|
36
|
+
|
|
37
|
+
def self.verify(path)
|
|
38
|
+
Verifier.call(path)
|
|
39
|
+
end
|
|
40
|
+
end
|
|
41
|
+
end
|