Something went wrong. Try again.
sentence embeddings in pure zig: bge-small with an HF-exact tokenizer and an SME matmul
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266//! BERT WordPiece tokenizer, matching HF `tokenizers` BertNormalizer (lowercase,//! accent strip, clean_text, CJK padding) + BertPreTokenizer + WordPiece.
const std = @import("std");const Allocator = std.mem.Allocator;const table = @import("unicode_table.zig");
pub const Tokenizer = struct { vocab: std.StringHashMapUnmanaged(u32), /// backing storage for vocab keys vocab_text: []u8, cls: u32, sep: u32, unk: u32,
pub const max_input_chars_per_word = 100;
/// `vocab_txt` is BERT's vocab.txt: one token per line, line index = id. pub fn init(gpa: Allocator, vocab_txt: []const u8) !Tokenizer { const text = try gpa.dupe(u8, vocab_txt); errdefer gpa.free(text); var vocab: std.StringHashMapUnmanaged(u32) = .empty; errdefer vocab.deinit(gpa);
var id: u32 = 0; var it = std.mem.splitScalar(u8, text, '\n'); while (it.next()) |line_raw| { const line = std.mem.trimEnd(u8, line_raw, "\r"); if (line.len == 0 and it.peek() == null) break; try vocab.put(gpa, line, id); id += 1; } return .{ .vocab = vocab, .vocab_text = text, .cls = vocab.get("[CLS]") orelse return error.MissingSpecialToken, .sep = vocab.get("[SEP]") orelse return error.MissingSpecialToken, .unk = vocab.get("[UNK]") orelse return error.MissingSpecialToken, }; }
pub fn deinit(self: *Tokenizer, gpa: Allocator) void { self.vocab.deinit(gpa); gpa.free(self.vocab_text); }
/// Encode `text` as [CLS] pieces... [SEP], truncated to `out.len` ids. /// `scratch` receives the normalized text and must hold 4 * text.len + 4 bytes. pub fn encode(self: *const Tokenizer, text: []const u8, scratch: []u8, out: []u32) usize { std.debug.assert(out.len >= 2); const norm = normalize(text, scratch); const limit = out.len - 1; var n: usize = 0; out[n] = self.cls; n += 1;
var words = WordIterator{ .text = norm }; while (words.next()) |word| { if (n >= limit) break; n = self.wordPiece(word, out, n, limit); } out[n] = self.sep; return n + 1; }
fn wordPiece(self: *const Tokenizer, word: []const u8, out: []u32, start: usize, limit: usize) usize { if ((std.unicode.utf8CountCodepoints(word) catch word.len) > max_input_chars_per_word) return push(out, start, limit, self.unk);
var buf: [4 * max_input_chars_per_word + 2]u8 = undefined; var pieces: [max_input_chars_per_word]u32 = undefined; var count: usize = 0; var pos: usize = 0; while (pos < word.len) { var end = word.len; var found: ?u32 = null; while (end > pos) { const piece = if (pos == 0) word[pos..end] else blk: { buf[0] = '#'; buf[1] = '#'; @memcpy(buf[2 .. 2 + end - pos], word[pos..end]); break :blk buf[0 .. 2 + end - pos]; }; if (self.vocab.get(piece)) |id| { found = id; break; } end = prevBoundary(word, end, pos); } const id = found orelse return push(out, start, limit, self.unk); pieces[count] = id; count += 1; pos = end; } var n = start; for (pieces[0..count]) |id| n = push(out, n, limit, id); return n; }};
fn push(out: []u32, n: usize, limit: usize, id: u32) usize { if (n >= limit) return n; out[n] = id; return n + 1;}
/// the start of the codepoint that ends at `end` (exclusive), not before `floor`fn prevBoundary(s: []const u8, end: usize, floor: usize) usize { var i = end - 1; while (i > floor and s[i] & 0xC0 == 0x80) i -= 1; return i;}
fn inRanges(comptime ranges: []const table.Range, cp: u21) bool { var lo: usize = 0; var hi: usize = ranges.len; while (lo < hi) { const mid = (lo + hi) / 2; if (cp < ranges[mid].lo) { hi = mid; } else if (cp > ranges[mid].hi) { lo = mid + 1; } else return true; } return false;}
fn mapping(cp: u21) ?[]const u8 { const m = &table.mappings; var lo: usize = 0; var hi: usize = m.len; while (lo < hi) { const mid = (lo + hi) / 2; if (cp < m[mid].cp) { hi = mid; } else if (cp > m[mid].cp) { lo = mid + 1; } else return table.blob[m[mid].off..][0..m[mid].len]; } return null;}
/// Apply BertNormalizer codepoint by codepoint. Invalid UTF-8 bytes are dropped.pub fn normalize(text: []const u8, out: []u8) []const u8 { var n: usize = 0; var i: usize = 0; while (i < text.len) { const len = std.unicode.utf8ByteSequenceLength(text[i]) catch { i += 1; continue; }; if (i + len > text.len) break; const cp = std.unicode.utf8Decode(text[i..][0..len]) catch { i += 1; continue; }; const raw = text[i..][0..len]; i += len;
if (cp < 0x80) { const c: u8 = @intCast(cp); if (c >= 'A' and c <= 'Z') { out[n] = c + 32; n += 1; continue; } if (c >= 0x20 and c < 0x7F) { out[n] = c; n += 1; continue; } } if (inRanges(&table.removed, cp)) continue; if (inRanges(&table.to_space, cp)) { out[n] = ' '; n += 1; continue; } if (inRanges(&table.cjk, cp)) { out[n] = ' '; @memcpy(out[n + 1 ..][0..len], raw); out[n + 1 + len] = ' '; n += len + 2; continue; } const bytes = mapping(cp) orelse raw; @memcpy(out[n..][0..bytes.len], bytes); n += bytes.len; } return out[0..n];}
/// Splits normalized text on whitespace, and isolates each punctuation codepoint.const WordIterator = struct { text: []const u8, pos: usize = 0,
const Class = enum { space, punct, word };
fn classAt(self: *const WordIterator, i: usize, len: *usize) Class { const b = self.text[i]; if (b < 0x80) { len.* = 1; return switch (b) { ' ', '\t', '\n', '\r', 0x0B, 0x0C => .space, '!'...'/', ':'...'@', '['...'`', '{'...'~' => .punct, else => .word, }; } const l = std.unicode.utf8ByteSequenceLength(b) catch 1; len.* = @min(l, self.text.len - i); const cp = std.unicode.utf8Decode(self.text[i..][0..len.*]) catch return .word; if (inRanges(&table.space, cp)) return .space; if (inRanges(&table.punct, cp)) return .punct; return .word; }
fn next(self: *WordIterator) ?[]const u8 { var len: usize = 0; while (self.pos < self.text.len and self.classAt(self.pos, &len) == .space) self.pos += len; if (self.pos >= self.text.len) return null; const start = self.pos; if (self.classAt(self.pos, &len) == .punct) { self.pos += len; return self.text[start..self.pos]; } while (self.pos < self.text.len and self.classAt(self.pos, &len) == .word) self.pos += len; return self.text[start..self.pos]; }};
test "normalize lowercases, strips accents, pads cjk, drops controls" { var buf: [128]u8 = undefined; try std.testing.expectEqualStrings("cafe naive", normalize("Café NAÏVE", &buf)); try std.testing.expectEqualStrings("a 中 b", normalize("a中b", &buf)); try std.testing.expectEqualStrings("ab", normalize("a\x00b", &buf)); try std.testing.expectEqualStrings("a b", normalize("a\u{3000}b", &buf));}
test "word iterator isolates punctuation and splits whitespace" { var it = WordIterator{ .text = "hello, world! x" }; const want = [_][]const u8{ "hello", ",", "world", "!", "x" }; for (want) |w| try std.testing.expectEqualStrings(w, it.next().?); try std.testing.expect(it.next() == null);}
test "wordpiece greedy longest match with ## continuation" { const gpa = std.testing.allocator; var tok = try Tokenizer.init(gpa, "[PAD]\n[UNK]\n[CLS]\n[SEP]\nun\n##aff\n##able\nhello\n"); defer tok.deinit(gpa); var scratch: [256]u8 = undefined; var ids: [16]u32 = undefined; const n = tok.encode("unaffable hello xyz", &scratch, &ids); try std.testing.expectEqualSlices(u32, &.{ 2, 4, 5, 6, 7, 1, 3 }, ids[0..n]);}
test "truncation keeps [SEP] as the last id" { const gpa = std.testing.allocator; var tok = try Tokenizer.init(gpa, "[PAD]\n[UNK]\n[CLS]\n[SEP]\na\n"); defer tok.deinit(gpa); var scratch: [256]u8 = undefined; var ids: [4]u32 = undefined; const n = tok.encode("a a a a a", &scratch, &ids); try std.testing.expectEqualSlices(u32, &.{ 2, 4, 4, 3 }, ids[0..n]);}