Module: StaticEmbeddings::Format
- Defined in:
- lib/static_embeddings/format.rb
Defined Under Namespace
Classes: TrieBuilder
Constant Summary collapse
- MAGIC =
"SEMBv1\0\0"- VERSION =
2- HEADER_SIZE =
320- ALIGNMENT =
64- TOKENIZER_BERT_WORDPIECE_V1 =
1- DTYPE_F32 =
1- POOLING_MEAN =
1- NORMALIZATION_NONE =
0- NORMALIZATION_L2 =
1- TRUNCATE_IDS_BEFORE_POOLING =
1- UNK_INCLUDE =
0- UNK_DROP =
1- EMPTY_ZERO_VECTOR =
0- EMPTY_RAISE =
1- SLOT_EMPTY =
0xFFFFFFFF- HASH_SEED =
2_166_136_261- LOAD_FACTOR =
0.70- MAX_TOKEN_CHARS_OFFSET =
124- CHECKSUM_OFFSET =
240- CHECKSUM_SIZE =
32- MAX_PROBE_OFFSET =
304- SECTION_FIELDS =
{ vocab_strings: 128, vocab_hash: 144, embeddings: 160, norm_tables: 176, provenance: 192, root_trie: 208, continuation_trie: 224 }.freeze
- HEADER_U32 =
{ 8 => VERSION, 12 => HEADER_SIZE, 16 => 1, 28 => TOKENIZER_BERT_WORDPIECE_V1, 32 => DTYPE_F32, 36 => POOLING_MEAN, 48 => TRUNCATE_IDS_BEFORE_POOLING, 108 => HASH_SEED }.freeze
- META_U32 =
{ 20 => :dim, 40 => :normalization_type, 44 => :max_tokens_default, 56 => :unk_policy, 60 => :empty_policy, 80 => :max_input_chars_per_word, 84 => :pad_id, 88 => :unk_id, 92 => :cls_id, 96 => :sep_id, 100 => :mask_id }.freeze
- META_BOOL =
{ 52 => :add_special_tokens, 64 => :do_lower_case, 68 => :strip_accents, 72 => :handle_chinese_chars, 76 => :clean_text }.freeze
- VERIFY_CHUNK_BYTES =
1024 * 1024
Class Method Summary collapse
- .align_body!(body) ⇒ Object
- .binary_string(capacity = nil) ⇒ Object
- .build_body(payloads) ⇒ Object
- .build_hash_table(tokens) ⇒ Object
- .build_header(meta, vocab_size, hash_size, max_probe, sections) ⇒ Object
- .build_wordpiece_tries(tokens, prefix) ⇒ Object
- .hash_bytes(str, seed = HASH_SEED) ⇒ Object
- .insert_slot!(slots, tokens, size, bytes, hash, offset, id) ⇒ Object
- .next_power_of_two(n) ⇒ Object
- .pack_slots(slots) ⇒ Object
- .put_prefix(header, prefix) ⇒ Object
- .put_section(header, offset, section) ⇒ Object
- .put_u32(buffer, offset, value) ⇒ Object
- .put_u64(buffer, offset, value) ⇒ Object
- .streaming_checksum(io, header) ⇒ Object
- .verify(path) ⇒ Object
- .write(path:, meta:, tokens:, matrix:, norm_tables:, provenance:) ⇒ Object
Class Method Details
.align_body!(body) ⇒ Object
254 255 256 257 |
# File 'lib/static_embeddings/format.rb', line 254 def align_body!(body) padding = (ALIGNMENT - ((HEADER_SIZE + body.bytesize) % ALIGNMENT)) % ALIGNMENT body << "\0".b * padding if padding.positive? end |
.binary_string(capacity = nil) ⇒ Object
149 150 151 152 |
# File 'lib/static_embeddings/format.rb', line 149 def binary_string(capacity = nil) str = capacity ? String.new(capacity: capacity) : +"" str.force_encoding(Encoding::BINARY) end |
.build_body(payloads) ⇒ Object
243 244 245 246 247 248 249 250 251 252 |
# File 'lib/static_embeddings/format.rb', line 243 def build_body(payloads) body = binary_string sections = {} payloads.each do |name, payload| align_body!(body) sections[name] = [HEADER_SIZE + body.bytesize, payload.bytesize] body << payload end [sections, body] end |
.build_hash_table(tokens) ⇒ Object
85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 |
# File 'lib/static_embeddings/format.rb', line 85 def build_hash_table(tokens) size = next_power_of_two((tokens.length / LOAD_FACTOR).ceil + 1) slots = Array.new(size) strings = binary_string max_probe = 0 tokens.each_with_index do |token, id| bytes = token.b offset = strings.bytesize strings << bytes probe = insert_slot!(slots, tokens, size, bytes, hash_bytes(bytes), offset, id) max_probe = probe if probe > max_probe end [size, strings, pack_slots(slots), max_probe] end |
.build_header(meta, vocab_size, hash_size, max_probe, sections) ⇒ Object
259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 |
# File 'lib/static_embeddings/format.rb', line 259 def build_header(, vocab_size, hash_size, max_probe, sections) header = "\0".b * HEADER_SIZE header[0, 8] = MAGIC.b HEADER_U32.each { |offset, value| put_u32(header, offset, value) } META_U32.each { |offset, key| put_u32(header, offset, .fetch(key)) } META_BOOL.each { |offset, key| put_u32(header, offset, .fetch(key) ? 1 : 0) } put_u32(header, 24, vocab_size) put_u32(header, 104, hash_size) put_u32(header, MAX_TOKEN_CHARS_OFFSET, .fetch(:max_token_chars)) put_u32(header, MAX_PROBE_OFFSET, max_probe) put_prefix(header, .fetch(:subword_prefix)) SECTION_FIELDS.each { |name, field| put_section(header, field, sections.fetch(name)) } header end |
.build_wordpiece_tries(tokens, prefix) ⇒ Object
180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 |
# File 'lib/static_embeddings/format.rb', line 180 def build_wordpiece_tries(tokens, prefix) root = TrieBuilder.new continuation = TrieBuilder.new prefix_bytes = prefix.b tokens.each_with_index do |token, id| bytes = token.b if !prefix_bytes.empty? && bytes.start_with?(prefix_bytes) && bytes.bytesize > prefix_bytes.bytesize continuation.insert(bytes.byteslice(prefix_bytes.bytesize, bytes.bytesize - prefix_bytes.bytesize), id) else root.insert(bytes, id) end end [root.pack, continuation.pack] end |
.hash_bytes(str, seed = HASH_SEED) ⇒ Object
77 78 79 |
# File 'lib/static_embeddings/format.rb', line 77 def hash_bytes(str, seed = HASH_SEED) str.each_byte.reduce(seed) { |h, b| ((h ^ b) * 16_777_619) & 0xFFFFFFFF } end |
.insert_slot!(slots, tokens, size, bytes, hash, offset, id) ⇒ Object
154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 |
# File 'lib/static_embeddings/format.rb', line 154 def insert_slot!(slots, tokens, size, bytes, hash, offset, id) pos = hash & (size - 1) probe = 1 loop do slot = slots[pos] unless slot slots[pos] = [hash, offset, bytes.bytesize, id] return probe end if slot[0] == hash && slot[2] == bytes.bytesize && tokens.fetch(slot[3]).b == bytes raise ArgumentError, "duplicate token in vocabulary: #{tokens.fetch(slot[3]).inspect}" end pos = (pos + 1) & (size - 1) probe += 1 end end |
.next_power_of_two(n) ⇒ Object
81 82 83 |
# File 'lib/static_embeddings/format.rb', line 81 def next_power_of_two(n) 1 << (n - 1).bit_length end |
.pack_slots(slots) ⇒ Object
173 174 175 176 177 178 |
# File 'lib/static_embeddings/format.rb', line 173 def pack_slots(slots) empty = [0, 0, 0, SLOT_EMPTY].pack("V4") packed = binary_string(slots.length * 16) slots.each { |slot| packed << (slot ? slot.pack("V4") : empty) } packed end |
.put_prefix(header, prefix) ⇒ Object
277 278 279 280 281 282 283 |
# File 'lib/static_embeddings/format.rb', line 277 def put_prefix(header, prefix) bytes = prefix.b raise ArgumentError, "subword prefix too long" if bytes.bytesize > 8 put_u32(header, 112, bytes.bytesize) header[116, 8] = bytes.ljust(8, "\0") end |
.put_section(header, offset, section) ⇒ Object
285 286 287 288 289 |
# File 'lib/static_embeddings/format.rb', line 285 def put_section(header, offset, section) off, size = section put_u64(header, offset, off) put_u64(header, offset + 8, size) end |
.put_u32(buffer, offset, value) ⇒ Object
291 292 293 |
# File 'lib/static_embeddings/format.rb', line 291 def put_u32(buffer, offset, value) buffer[offset, 4] = [value].pack("V") end |
.put_u64(buffer, offset, value) ⇒ Object
295 296 297 |
# File 'lib/static_embeddings/format.rb', line 295 def put_u64(buffer, offset, value) buffer[offset, 8] = [value & 0xFFFFFFFF, value >> 32].pack("V2") end |
.streaming_checksum(io, header) ⇒ Object
138 139 140 141 142 143 144 145 146 147 |
# File 'lib/static_embeddings/format.rb', line 138 def streaming_checksum(io, header) digest = Digest::SHA256.new zeroed = header.dup zeroed[CHECKSUM_OFFSET, CHECKSUM_SIZE] = "\0".b * CHECKSUM_SIZE digest << zeroed buffer = String.new(capacity: VERIFY_CHUNK_BYTES) digest << buffer while io.read(VERIFY_CHUNK_BYTES, buffer) digest.digest end |
.verify(path) ⇒ Object
124 125 126 127 128 129 130 131 132 133 134 135 136 |
# File 'lib/static_embeddings/format.rb', line 124 def verify(path) raise InvalidModelError, "file too small" if File.size(path) < HEADER_SIZE File.open(path, "rb") do |io| header = io.read(HEADER_SIZE) raise InvalidModelError, "bad magic" unless header.byteslice(0, 8) == MAGIC.b stored = header.byteslice(CHECKSUM_OFFSET, CHECKSUM_SIZE) actual = streaming_checksum(io, header) { ok: stored == actual, expected: actual.unpack1("H*"), stored: stored.unpack1("H*") } end end |
.write(path:, meta:, tokens:, matrix:, norm_tables:, provenance:) ⇒ Object
102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 |
# File 'lib/static_embeddings/format.rb', line 102 def write(path:, meta:, tokens:, matrix:, norm_tables:, provenance:) hash_size, strings, hash_blob, max_probe = build_hash_table(tokens) root_trie, continuation_trie = build_wordpiece_tries(tokens, .fetch(:subword_prefix)) sections, body = build_body( vocab_strings: strings, vocab_hash: hash_blob, embeddings: matrix, norm_tables: norm_tables, provenance: provenance, root_trie: root_trie, continuation_trie: continuation_trie ) header = build_header(, tokens.length, hash_size, max_probe, sections) file = header << body digest = Digest::SHA256.digest(file) file[CHECKSUM_OFFSET, CHECKSUM_SIZE] = digest File.binwrite(path, file) { bytes: file.bytesize, sha256: digest.unpack1("H*"), hash_table_size: hash_size } end |