//! What a sentence-transformers BERT snapshot says about itself. //! //! Everything that differs between models of this family is read from the //! snapshot directory rather than assumed: shape (config.json), pooling //! (1_Pooling/config.json), output normalization (modules.json), length cap //! (sentence_bert_config.json) and tensor naming. Anything embedz cannot run //! faithfully is rejected here with a specific error, so a model swap fails at //! load instead of producing plausible wrong vectors. const std = @import("std"); const Allocator = std.mem.Allocator; const Io = std.Io; const json = std.json; const ops = @import("ops.zig"); pub const Pooling = enum { cls, mean }; pub const Spec = struct { vocab: usize, /// rows of the token-type embedding table; embedz always uses type 0 type_vocab: usize, hidden: usize, heads: usize, layers: usize, intermediate: usize, max_positions: usize, eps: f32, /// tokens per text including [CLS] and [SEP] max_seq: usize, pooling: Pooling, normalize: bool, /// "" or "bert.", whichever the checkpoint's tensor names use tensor_prefix: []const u8 = "", pub fn headDim(self: Spec) usize { return self.hidden / self.heads; } }; pub const Error = error{ UnsupportedArchitecture, UnsupportedActivation, UnsupportedPositionEmbedding, UnsupportedHeadDim, UnsupportedPooling, UnsupportedTokenizer, MissingPoolingConfig, BadConfig, }; /// Read and validate the spec from a snapshot directory. pub fn read(gpa: Allocator, io: Io, dir: Io.Dir) !Spec { var arena_state = std.heap.ArenaAllocator.init(gpa); defer arena_state.deinit(); const a = arena_state.allocator(); const config = try readJson(a, io, dir, "config.json") orelse return error.BadConfig; const pooling_cfg = try readJson(a, io, dir, "1_Pooling/config.json"); const modules = try readJson(a, io, dir, "modules.json"); const st_cfg = try readJson(a, io, dir, "sentence_bert_config.json"); const tokenizer = try readJson(a, io, dir, "tokenizer.json") orelse return error.UnsupportedTokenizer; return fromJson(config, pooling_cfg, modules, st_cfg, tokenizer); } pub fn fromJson(config: json.Value, pooling_cfg: ?json.Value, modules: ?json.Value, st_cfg: ?json.Value, tokenizer: json.Value) Error!Spec { const archs = field(config, "architectures") orelse return error.UnsupportedArchitecture; if (archs != .array) return error.UnsupportedArchitecture; const is_bert = for (archs.array.items) |arch| { if (arch == .string and std.mem.eql(u8, arch.string, "BertModel")) break true; } else false; if (!is_bert) return error.UnsupportedArchitecture; if (!strIs(config, "hidden_act", "gelu")) return error.UnsupportedActivation; if (field(config, "position_embedding_type")) |p| { if (p != .string or !std.mem.eql(u8, p.string, "absolute")) return error.UnsupportedPositionEmbedding; } const hidden = try int(config, "hidden_size"); const heads = try int(config, "num_attention_heads"); if (heads == 0 or hidden % heads != 0) return error.BadConfig; if (std.mem.indexOfScalar(usize, &ops.head_dims, hidden / heads) == null) return error.UnsupportedHeadDim; const max_positions = try int(config, "max_position_embeddings"); const max_seq = if (st_cfg) |c| (if (field(c, "max_seq_length")) |v| try asInt(v) else max_positions) else max_positions; try checkTokenizer(tokenizer); return .{ .vocab = try int(config, "vocab_size"), .type_vocab = if (field(config, "type_vocab_size")) |v| try asInt(v) else 2, .hidden = hidden, .heads = heads, .layers = try int(config, "num_hidden_layers"), .intermediate = try int(config, "intermediate_size"), .max_positions = max_positions, .eps = if (field(config, "layer_norm_eps")) |v| try asFloat(v) else 1e-12, .max_seq = @min(max_seq, max_positions), .pooling = try pooling(pooling_cfg orelse return error.MissingPoolingConfig), .normalize = if (modules) |m| hasNormalize(m) else false, }; } fn pooling(cfg: json.Value) Error!Pooling { const modes = [_]struct { key: []const u8, mode: ?Pooling }{ .{ .key = "pooling_mode_cls_token", .mode = .cls }, .{ .key = "pooling_mode_mean_tokens", .mode = .mean }, .{ .key = "pooling_mode_max_tokens", .mode = null }, .{ .key = "pooling_mode_mean_sqrt_len_tokens", .mode = null }, .{ .key = "pooling_mode_weightedmean_tokens", .mode = null }, .{ .key = "pooling_mode_lasttoken", .mode = null }, }; var chosen: ?Pooling = null; for (modes) |m| { const on = field(cfg, m.key) orelse continue; if (on != .bool or !on.bool) continue; const mode = m.mode orelse return error.UnsupportedPooling; if (chosen != null) return error.UnsupportedPooling; chosen = mode; } return chosen orelse error.UnsupportedPooling; } fn hasNormalize(modules: json.Value) bool { if (modules != .array) return false; for (modules.array.items) |m| { const t = field(m, "type") orelse continue; if (t == .string and std.mem.endsWith(u8, t.string, ".Normalize")) return true; } return false; } /// src/unicode_table.zig was generated from an uncased BertNormalizer and /// BertPreTokenizer, and the tokenizer implements "##" WordPiece; anything else /// would tokenize differently without failing. fn checkTokenizer(tok: json.Value) Error!void { const norm = field(tok, "normalizer") orelse return error.UnsupportedTokenizer; if (!strIs(norm, "type", "BertNormalizer")) return error.UnsupportedTokenizer; if (!boolIs(norm, "lowercase", true) or !boolIs(norm, "clean_text", true) or !boolIs(norm, "handle_chinese_chars", true)) return error.UnsupportedTokenizer; if (field(norm, "strip_accents")) |s| switch (s) { .null => {}, .bool => |b| if (!b) return error.UnsupportedTokenizer, else => return error.UnsupportedTokenizer, }; const pre = field(tok, "pre_tokenizer") orelse return error.UnsupportedTokenizer; if (!strIs(pre, "type", "BertPreTokenizer")) return error.UnsupportedTokenizer; const model = field(tok, "model") orelse return error.UnsupportedTokenizer; if (!strIs(model, "type", "WordPiece") or !strIs(model, "continuing_subword_prefix", "##") or !strIs(model, "unk_token", "[UNK]")) return error.UnsupportedTokenizer; } fn readJson(a: Allocator, io: Io, dir: Io.Dir, path: []const u8) !?json.Value { const bytes = dir.readFileAlloc(io, path, a, .limited(64 << 20)) catch |err| switch (err) { error.FileNotFound => return null, else => return err, }; return try json.parseFromSliceLeaky(json.Value, a, bytes, .{}); } fn field(v: json.Value, key: []const u8) ?json.Value { if (v != .object) return null; return v.object.get(key); } fn strIs(v: json.Value, key: []const u8, want: []const u8) bool { const f = field(v, key) orelse return false; return f == .string and std.mem.eql(u8, f.string, want); } fn boolIs(v: json.Value, key: []const u8, want: bool) bool { const f = field(v, key) orelse return false; return f == .bool and f.bool == want; } fn int(v: json.Value, key: []const u8) Error!usize { return asInt(field(v, key) orelse return error.BadConfig); } fn asInt(v: json.Value) Error!usize { if (v != .integer or v.integer <= 0) return error.BadConfig; return @intCast(v.integer); } fn asFloat(v: json.Value) Error!f32 { return switch (v) { .float => |f| @floatCast(f), .integer => |i| @floatFromInt(i), else => error.BadConfig, }; } pub const test_tokenizer = \\{"normalizer":{"type":"BertNormalizer","clean_text":true,"handle_chinese_chars":true,"strip_accents":null,"lowercase":true}, \\ "pre_tokenizer":{"type":"BertPreTokenizer"}, \\ "model":{"type":"WordPiece","unk_token":"[UNK]","continuing_subword_prefix":"##"}} ; fn parse(a: Allocator, s: []const u8) !json.Value { return json.parseFromSliceLeaky(json.Value, a, s, .{}); } test "reads a mean-pooled 6-layer model and its length cap" { var arena = std.heap.ArenaAllocator.init(std.testing.allocator); defer arena.deinit(); const a = arena.allocator(); const spec = try fromJson( try parse(a, "{\"architectures\":[\"BertModel\"],\"hidden_act\":\"gelu\",\"hidden_size\":384,\"num_attention_heads\":12,\"num_hidden_layers\":6,\"intermediate_size\":1536,\"max_position_embeddings\":512,\"vocab_size\":30522,\"layer_norm_eps\":1e-12}"), try parse(a, "{\"pooling_mode_cls_token\":false,\"pooling_mode_mean_tokens\":true}"), try parse(a, "[{\"type\":\"sentence_transformers.models.Transformer\"},{\"type\":\"sentence_transformers.models.Normalize\"}]"), try parse(a, "{\"max_seq_length\":256}"), try parse(a, test_tokenizer), ); try std.testing.expectEqual(@as(usize, 6), spec.layers); try std.testing.expectEqual(Pooling.mean, spec.pooling); try std.testing.expect(spec.normalize); try std.testing.expectEqual(@as(usize, 256), spec.max_seq); try std.testing.expectEqual(@as(usize, 32), spec.headDim()); } test "rejects what embedz cannot run faithfully" { var arena = std.heap.ArenaAllocator.init(std.testing.allocator); defer arena.deinit(); const a = arena.allocator(); const bert = "{\"architectures\":[\"BertModel\"],\"hidden_act\":\"gelu\",\"hidden_size\":384,\"num_attention_heads\":12,\"num_hidden_layers\":6,\"intermediate_size\":1536,\"max_position_embeddings\":512,\"vocab_size\":30522}"; const cls = try parse(a, "{\"pooling_mode_cls_token\":true}"); const tok = try parse(a, test_tokenizer); try std.testing.expectError(error.UnsupportedArchitecture, fromJson(try parse(a, "{\"architectures\":[\"XLMRobertaModel\"]}"), cls, null, null, tok)); try std.testing.expectError(error.UnsupportedPooling, fromJson(try parse(a, bert), try parse(a, "{\"pooling_mode_max_tokens\":true}"), null, null, tok)); try std.testing.expectError(error.MissingPoolingConfig, fromJson(try parse(a, bert), null, null, null, tok)); try std.testing.expectError(error.UnsupportedHeadDim, fromJson(try parse(a, "{\"architectures\":[\"BertModel\"],\"hidden_act\":\"gelu\",\"hidden_size\":480,\"num_attention_heads\":12,\"num_hidden_layers\":6,\"intermediate_size\":1536,\"max_position_embeddings\":512,\"vocab_size\":30522}"), cls, null, null, tok)); const cased = try parse(a, "{\"normalizer\":{\"type\":\"BertNormalizer\",\"clean_text\":true,\"handle_chinese_chars\":true,\"strip_accents\":null,\"lowercase\":false},\"pre_tokenizer\":{\"type\":\"BertPreTokenizer\"},\"model\":{\"type\":\"WordPiece\",\"unk_token\":\"[UNK]\",\"continuing_subword_prefix\":\"##\"}}"); try std.testing.expectError(error.UnsupportedTokenizer, fromJson(try parse(a, bert), cls, null, null, cased)); }