From 71f23f5162a243714191829760efd8d9068901f8 Mon Sep 17 00:00:00 2001 From: Zijian Yi Date: Wed, 22 Jul 2026 17:54:58 -0700 Subject: [PATCH] [gemma_cpp] Add support for Qwen3 models (0.6B, 1.7B and 4B). Qwen3 shares a very similar architecture as Gemma3 ([comparison](https://sebastianraschka.com/llm-architecture-gallery/?compare=qwen3-0-6b%2Cgemma-3-270m#architecture-diff-tool)). Differences include: - *Tokenizer* : Qwen3 uses bytelevel BPE while Gemma3 uses `Sentencepiece`; - *Embedding* : Qwen3 0.6B and 1.7B do not share the `input_embedding` and `lm_head` weights, and there is no Embedding Scaling for Qwen3; - *Activation* : Qwen3 uses SiLU while gemma3 uses GeLU; - *Norm* : Qwen3 only has two pre-norms, and they are not Zero-centered (we substract 1.0 while converting the weights to keep the kernel untouched); - *Attention* : Qwen3 uses full-attention for all the layers. - *Prompt Wrapping* : Qwen3 uses different special tokens and does not require . PiperOrigin-RevId: 952434363 --- deepseek/deepseek_tensors.cc | 10 -- gemma/configs.cc | 100 +++++++++++ gemma/configs.h | 12 ++ gemma/gemma.cc | 8 +- gemma/tensor_info.cc | 10 ++ gemma/tokenizer.cc | 10 +- gemma/tokenizer.h | 1 + gemma/weights.h | 5 +- python/configs.cc | 3 + python/convert_from_safetensors.py | 246 ++++++++++++++++++++++++- tokenizer/bpe_tokenizer.cc | 280 ++++++++++++++++++++++++++++- 11 files changed, 657 insertions(+), 28 deletions(-) diff --git a/deepseek/deepseek_tensors.cc b/deepseek/deepseek_tensors.cc index 424a2994..dc182d3b 100644 --- a/deepseek/deepseek_tensors.cc +++ b/deepseek/deepseek_tensors.cc @@ -58,16 +58,6 @@ void TensorInfoRegistry::AddDeepSeekModelTensors(const ModelConfig& config) { }); }; - if (config.HasMLA()) { - // DeepSeek models have an untied output head. - Add(no_suffix, { - .base_name = "lm_head", - .source_names = {"lm_head.weight", "lm_head"}, - .axes = {0, 1}, - .shape = {config.vocab_size, config.model_dim}, - .min_size = Type::kBF16, - }); - } if (config.hc_mult > 1) { add_hc_collapse("hc_head", ""); } diff --git a/gemma/configs.cc b/gemma/configs.cc index 6e96650f..89297b0e 100644 --- a/gemma/configs.cc +++ b/gemma/configs.cc @@ -666,6 +666,85 @@ LayerConfig ModelConfig::MTPLayerConfig() const { return lc; } +static ModelConfig ConfigBaseQwen3() { + ModelConfig config = ConfigNoSSM(); + config.vocab_size = 151936; + config.max_seq_len = 32768; + config.eos_id = 151645; + config.secondary_eos_id = 151643; + return config; +} + +static LayerConfig LayerConfigQwen3(size_t model_dim, size_t ff_hidden_dim, + size_t heads, size_t kv_heads, + size_t qkv_dim) { + LayerConfig config; + config.model_dim = model_dim; + config.ff_hidden_dim = ff_hidden_dim; + config.heads = heads; + config.kv_heads = kv_heads; + config.qkv_dim = qkv_dim; + config.optimized_gating = false; + config.post_norm = PostNormType::None; + config.activation = ActivationType::Silu; + config.use_qk_norm = true; + return config; +} + +static ModelConfig ConfigQWEN3_600M() { + ModelConfig config = ConfigBaseQwen3(); + config.display_name = "Qwen3_0.6B"; + config.model = Model::QWEN3_600M; + config.wrapping = PromptWrapping::GEMMA_IT; + config.model_dim = 1024; + + LayerConfig layer_config = + LayerConfigQwen3(config.model_dim, 3072, 16, 8, 128); + config.num_layers = 28; + config.layer_configs = {config.num_layers, layer_config}; + config.query_scale = QueryScaleType::SqrtKeySize; + config.use_global_timescale = true; + config.attention_window_sizes = + FixedAttentionWindowSizes<28>(config.max_seq_len); + return config; +} + +static ModelConfig ConfigQWEN3_2B() { + ModelConfig config = ConfigBaseQwen3(); + config.display_name = "Qwen3_1.7B"; + config.model = Model::QWEN3_2B; + config.wrapping = PromptWrapping::GEMMA_IT; + config.model_dim = 2048; + + LayerConfig layer_config = + LayerConfigQwen3(config.model_dim, 6144, 16, 8, 128); + config.num_layers = 28; + config.layer_configs = {config.num_layers, layer_config}; + config.query_scale = QueryScaleType::SqrtKeySize; + config.use_global_timescale = true; + config.attention_window_sizes = + FixedAttentionWindowSizes<28>(config.max_seq_len); + return config; +} + +static ModelConfig ConfigQwen3_4B() { + ModelConfig config = ConfigBaseQwen3(); + config.display_name = "Qwen3_4B"; + config.model = Model::QWEN3_4B; + config.wrapping = PromptWrapping::GEMMA_IT; + config.model_dim = 2560; + + LayerConfig layer_config = + LayerConfigQwen3(config.model_dim, 9728, 32, 8, 128); + config.num_layers = 36; + config.layer_configs = {config.num_layers, layer_config}; + config.query_scale = QueryScaleType::SqrtKeySize; + config.use_global_timescale = true; + config.attention_window_sizes = + FixedAttentionWindowSizes<36>(config.max_seq_len); + return config; +} + static ModelConfig ConfigFromModel(Model model) { switch (model) { case Model::GEMMA2_2B: @@ -704,6 +783,12 @@ static ModelConfig ConfigFromModel(Model model) { return ConfigGemma4_2B(); case Model::DEEPSEEK4_FLASH: return ConfigDeepSeek4_Flash(); + case Model::QWEN3_600M: + return ConfigQWEN3_600M(); + case Model::QWEN3_2B: + return ConfigQWEN3_2B(); + case Model::QWEN3_4B: + return ConfigQwen3_4B(); default: HWY_ABORT("Model type %d unknown.", static_cast(model)); } @@ -749,6 +834,12 @@ const char* ModelPrefix(Model model) { return "gemma4-2b"; case Model::DEEPSEEK4_FLASH: return "deepseek4-flash"; + case Model::QWEN3_600M: + return "qwen3-0_6b"; + case Model::QWEN3_2B: + return "qwen3-2b"; + case Model::QWEN3_4B: + return "qwen3-4b"; default: HWY_ABORT("Model type %d unknown.", static_cast(model)); } @@ -943,6 +1034,13 @@ Model DeduceModel(const Path& blob_path, size_t layers, int layer_types) { case 27: return (layer_types & kDeduced448) ? Model::PALIGEMMA2_3B_448 : Model::PALIGEMMA2_3B_224; + case 28: + if (blob_path.path.find("qwen3-2b") != std::string::npos || + blob_path.path.find("qwen3-1_7b") != std::string::npos) { + return Model::QWEN3_2B; + } + return Model::QWEN3_600M; + case 30: return Model::GEMMA4_26B_MOE; @@ -951,6 +1049,8 @@ Model DeduceModel(const Path& blob_path, size_t layers, int layer_types) { : Model::GEMMA3_4B_LM; case 35: return Model::GEMMA4_2B; + case 36: + return Model::QWEN3_4B; case 42: if (layer_types & kDeducedViT) { return (layer_types & kDeduced448) ? Model::PALIGEMMA2_10B_448 diff --git a/gemma/configs.h b/gemma/configs.h index 52b856f6..58a9345a 100644 --- a/gemma/configs.h +++ b/gemma/configs.h @@ -273,6 +273,9 @@ enum class Model { GEMMA4_26B_MOE, GEMMA4_2B, DEEPSEEK4_FLASH, + QWEN3_600M, + QWEN3_2B, // 1.7B rounded up for readability. + QWEN3_4B, kSentinel, }; @@ -632,6 +635,15 @@ struct ModelConfig : public IFields { return false; } + bool IsQwen3() const { + return model == Model::QWEN3_600M || model == Model::QWEN3_2B || + model == Model::QWEN3_4B; + } + + bool HasLmHead() const { + return model == Model::QWEN3_600M || model == Model::QWEN3_2B || HasMLA(); + } + // Synthesized config for the multi-token-prediction block (DeepSeek V4): // a full extra layer (dense MLA + MoE) used for speculative decoding, not // part of the main stack. Only valid when `num_mtp_layers > 0`. diff --git a/gemma/gemma.cc b/gemma/gemma.cc index 5d596d52..f482c689 100644 --- a/gemma/gemma.cc +++ b/gemma/gemma.cc @@ -220,6 +220,10 @@ static float EmbeddingScaling(size_t model_dim) { hwy::ConvertScalarTo(sqrtf(static_cast(model_dim)))); } +static bool HasEmbeddingScaling(const ModelConfig& model_config) { + return !(model_config.IsQwen3() || model_config.HasMLA()); +} + // `x_row` indicates which row of `x` to write to. // `pos` is the *token*'s position for `AddAbsolutePositionalEmbeddings`, not // the start of the batch, because this is called for batches of tokens in @@ -256,9 +260,9 @@ HWY_NOINLINE size_t EmbedMMToken(int token, size_t x_row, size_t pos, } const size_t model_dim = model_config.model_dim; - // DeepSeek does not scale embeddings by sqrt(model_dim). + // Qwen3/DeepSeekV4 does not scale embeddings by sqrt(model_dim). const float emb_scaling = - model_config.HasMLA() ? 1.0f : EmbeddingScaling(model_dim); + HasEmbeddingScaling(model_config) ? EmbeddingScaling(model_dim) : 1.0f; HWY_DASSERT(token >= 0); HWY_DASSERT(token < static_cast(model_config.vocab_size)); diff --git a/gemma/tensor_info.cc b/gemma/tensor_info.cc index 37d5822b..190bb119 100644 --- a/gemma/tensor_info.cc +++ b/gemma/tensor_info.cc @@ -59,6 +59,16 @@ void TensorInfoRegistry::AddModelTensors(const ModelConfig& config) { .min_size = Type::kBF16, }); AddDeepSeekModelTensors(config); + if (config.HasLmHead()) { + Add(no_suffix, + { + .base_name = "lm_head", + .source_names = {"lm_head/weight", "lm_head.weight", "lm_head"}, + .axes = {0, 1}, + .shape = {config.vocab_size, config.model_dim}, + .min_size = Type::kBF16, + }); + } Add(no_suffix, { .base_name = "enc_norm_bias", .source_names = {"img/Transformer/encoder_norm/bias"}, diff --git a/gemma/tokenizer.cc b/gemma/tokenizer.cc index a6b7b273..8396b1b3 100644 --- a/gemma/tokenizer.cc +++ b/gemma/tokenizer.cc @@ -75,6 +75,12 @@ GemmaChatTemplate::GemmaChatTemplate(const GemmaTokenizer& tokenizer, sot_user_ = {105, 2364, 107}; sot_model_ = {105, 4368, 107}; eot_ = {106, 107}; + } else if (model == Model::QWEN3_600M || model == Model::QWEN3_2B || + model == Model::QWEN3_4B) { + prepend_bos_ = false; + HWY_ASSERT(tokenizer.Encode("<|im_start|>user\n", &sot_user_)); + HWY_ASSERT(tokenizer.Encode("<|im_start|>assistant\n", &sot_model_)); + HWY_ASSERT(tokenizer.Encode("<|im_end|>\n", &eot_)); } else { sot_user_.reserve(3); if (!tokenizer.Encode("user\n", &sot_user_)) return; @@ -101,7 +107,9 @@ std::vector GemmaChatTemplate::Apply(size_t pos, // Start with BOS, or prepend end_of_turn if this is a continuation. if (pos == 0) { - out.push_back(BOS_ID); + if (prepend_bos_) { + out.push_back(BOS_ID); + } } else { out.insert(out.cend(), eot_.cbegin(), eot_.cend()); } diff --git a/gemma/tokenizer.h b/gemma/tokenizer.h index 89eaf4b2..fa1e3894 100644 --- a/gemma/tokenizer.h +++ b/gemma/tokenizer.h @@ -77,6 +77,7 @@ class GemmaChatTemplate { std::vector pali_sep_; std::vector vlm_soi_; std::vector vlm_eoi_; + bool prepend_bos_ = true; }; std::vector WrapAndTokenize(const GemmaTokenizer& tokenizer, diff --git a/gemma/weights.h b/gemma/weights.h index 0c4ce439..39660ee4 100644 --- a/gemma/weights.h +++ b/gemma/weights.h @@ -16,7 +16,6 @@ #ifndef THIRD_PARTY_GEMMA_CPP_GEMMA_WEIGHTS_H_ #define THIRD_PARTY_GEMMA_CPP_GEMMA_WEIGHTS_H_ -#include // isnan #include #include @@ -252,7 +251,7 @@ struct LayerWeightsPtrs { MatPtr key_norm_scale; // at least BF16. MatPtr query_norm_scale; // at least BF16. - + MatPtr router_scale; MatPtr p_expert_sc; MatPtr post_ffw1_ns; @@ -594,7 +593,7 @@ struct WeightsPtrs { LayerWeightsPtrs* other_layer2 = nullptr; func(TENSOR_ARGS(embedder_input_embedding, kMustRead)); func(TENSOR_ARGS(final_norm_scale, kMustRead)); - if (config_.HasMLA()) { + if (config_.HasLmHead()) { func(TENSOR_ARGS(lm_head, kMustRead)); } if (config_.hc_mult > 1) { diff --git a/python/configs.cc b/python/configs.cc index aa3ad27d..60a7255a 100644 --- a/python/configs.cc +++ b/python/configs.cc @@ -104,6 +104,9 @@ PYBIND11_MODULE(configs, py_module) { .value("GEMMA3_12B_LM", gcpp::Model::GEMMA3_12B_LM) .value("GEMMA3_27B_LM", gcpp::Model::GEMMA3_27B_LM) .value("DEEPSEEK4_FLASH", gcpp::Model::DEEPSEEK4_FLASH) + .value("QWEN3_600M", gcpp::Model::QWEN3_600M) + .value("QWEN3_2B", gcpp::Model::QWEN3_2B) + .value("QWEN3_4B", gcpp::Model::QWEN3_4B) // Insert new models above this line. .value("PALIGEMMA_448", gcpp::Model::PALIGEMMA_448); diff --git a/python/convert_from_safetensors.py b/python/convert_from_safetensors.py index 518c030d..2e2e8e70 100644 --- a/python/convert_from_safetensors.py +++ b/python/convert_from_safetensors.py @@ -114,9 +114,23 @@ def pack_bpe_tokenizer(tokenizer_json_path: str) -> bytes: if left in vocab_map and right in vocab_map and (left + right) in vocab_map: merge_ranks.append((rank, vocab_map[left], vocab_map[right])) - bpe_flags = 1 << 1 # kFlagSpaceReplace - if any(f"<0x{b:02X}>" in vocab_map for b in range(256)): - bpe_flags |= 1 << 0 # kFlagByteFallback + def _json_has_byte_level(obj) -> bool: + if isinstance(obj, dict): + if obj.get("type") == "ByteLevel": + return True + return any(_json_has_byte_level(v) for v in obj.values()) + elif isinstance(obj, list): + return any(_json_has_byte_level(v) for v in obj) + return False + + if _json_has_byte_level(j.get("pre_tokenizer")) or _json_has_byte_level( + j.get("decoder") + ): + bpe_flags = 1 << 2 # kFlagByteLevel + else: + bpe_flags = 1 << 1 # kFlagSpaceReplace + if any(f"<0x{b:02X}>" in vocab_map for b in range(256)): + bpe_flags |= 1 << 0 # kFlagByteFallback unk = model.get("unk_token", "") unk_id = vocab_map.get(unk, 0) @@ -639,7 +653,9 @@ def export_gemma3_lm_sbs( for k in f.keys(): # TranslateGemma checkpoints sometimes still ship the vision tower / # projector tensors. Silently drop them — this is the LM-only path. - if k.startswith("vision_tower.") or k.startswith("multi_modal_projector."): + if k.startswith("vision_tower.") or k.startswith( + "multi_modal_projector." + ): continue params[k] = f.get_tensor(k) @@ -884,6 +900,222 @@ def add_gating_einsum(i): ) +def export_qwen3_lm_sbs( + model_specifier: str, + load_path: str, + tokenizer_file: str, + csv_file: str, + sbs_file: str, +) -> None: + """Exports sbs file from a text-only Qwen 3 safetensors checkpoint.""" + if load_path.endswith(".json"): + with open(load_path, "r") as f: + j_obj = json.load(f) + files = list(set(j_obj["weight_map"].values())) + files = [os.path.join(os.path.dirname(load_path), f) for f in files] + else: + files = [load_path] + + params: Dict[str, Any] = {} + for file in files: + with safetensors.safe_open(file, framework="pt") as f: + for k in f.keys(): + params[k] = f.get_tensor(k) + + if "model.embed_tokens.weight" not in params: + raise ValueError( + "Could not locate 'model.embed_tokens.weight' in checkpoint." + ) + llm_prefix = "model." + + embed_tokens = params[f"{llm_prefix}embed_tokens.weight"] + vocab_size, model_dim = embed_tokens.shape + hidden_dim = params[f"{llm_prefix}layers.0.mlp.gate_proj.weight"].shape[0] + + has_qk_norm = f"{llm_prefix}layers.0.self_attn.q_norm.weight" in params + head_dim = params[f"{llm_prefix}layers.0.self_attn.q_norm.weight"].shape[0] + + num_q_heads = ( + params[f"{llm_prefix}layers.0.self_attn.q_proj.weight"].shape[0] + // head_dim + ) + num_kv_heads = ( + params[f"{llm_prefix}layers.0.self_attn.k_proj.weight"].shape[0] + // head_dim + ) + num_layers = len( + set([k for k in params.keys() if k.endswith("input_layernorm.weight")]) + ) + + print( + f"Qwen3 LM: vocab={vocab_size} dim={model_dim} hidden={hidden_dim} " + f"q_heads={num_q_heads} kv_heads={num_kv_heads} " + f"head_dim={head_dim} layers={num_layers} qk_norm={has_qk_norm}" + ) + + writer = compression.SbsWriter(sbs_file) + metadata = [] + scales = {} + + def add_data(param_name, data, expected_shape, sbs_name, layer_index=None): + if not isinstance(expected_shape, tuple): + expected_shape = (expected_shape,) + print(f"Writing {param_name} with shape {data.shape} e:{expected_shape}") + assert data.shape == expected_shape, param_name + + assert isinstance(data, torch.Tensor) + data = data.to(torch.float32).numpy() + data = np.array(data) + + if layer_index is not None: + param_name = param_name % layer_index + sbs_name = sbs_name + f"_{layer_index}" + + value = flatten_f32(data) + scale = compute_scale(value) + both_names = param_name + "::" + sbs_name + metadata.append((both_names, data.dtype, data.shape, scale)) + + if _is_float_param(sbs_name): + packed = configs.Type.kF32 + elif _is_bf16_param(sbs_name): + packed = configs.Type.kBF16 + else: + packed = configs.Type.kSFP + assert scale == 1.0, f"Scale for {both_names} is not 1.0" + scales[sbs_name] = scale + sys.stdout.flush() + + info = configs.TensorInfo() + info.name = sbs_name + info.shape = data.shape + writer.insert(sbs_name, value, packed, info) + + def add_qkv_einsum(i): + q = params.pop(f"{llm_prefix}layers.{i}.self_attn.q_proj.weight") + k = params.pop(f"{llm_prefix}layers.{i}.self_attn.k_proj.weight") + v = params.pop(f"{llm_prefix}layers.{i}.self_attn.v_proj.weight") + n_kv = k.shape[0] // head_dim + q = q.reshape(num_q_heads, head_dim, model_dim) + k = k.reshape(n_kv, head_dim, model_dim) + v = v.reshape(n_kv, head_dim, model_dim) + stacked = torch.stack((k, v), dim=0) # (2, K, H, D) + transposed = stacked.transpose(0, 1) # (K, 2, H, D) + reshaped = transposed.reshape(2 * n_kv, head_dim, model_dim) + qkv = torch.cat([q, reshaped], dim=0) + add_data( + f"{llm_prefix}layers.%d.self_attn.qkv_proj.weight", + qkv, + (num_q_heads + 2 * n_kv, head_dim, model_dim), + "qkv_ein", + i, + ) + + def add_att_einsum(i): + o = params.pop(f"{llm_prefix}layers.{i}.self_attn.o_proj.weight") + o = o.reshape(model_dim, num_q_heads, head_dim).permute(1, 0, 2) + add_data( + f"{llm_prefix}layers.%d.self_attn.o_proj.weight", + o, + (num_q_heads, model_dim, head_dim), + "att_ein", + i, + ) + + def add_gating_einsum(i): + gate = params.pop(f"{llm_prefix}layers.{i}.mlp.gate_proj.weight") + up = params.pop(f"{llm_prefix}layers.{i}.mlp.up_proj.weight") + assert gate.shape == up.shape == (hidden_dim, model_dim) + gating = torch.stack([gate, up], dim=0) + add_data( + f"{llm_prefix}layers.%d.mlp.gating_einsum.weight", + gating, + (2, hidden_dim, model_dim), + "gating_ein", + i, + ) + + # Non-layer tensors. + add_data( + f"{llm_prefix}embed_tokens.weight", + params.pop(f"{llm_prefix}embed_tokens.weight"), + (vocab_size, model_dim), + "c_embedding", + ) + add_data( + f"{llm_prefix}norm.weight", + params.pop(f"{llm_prefix}norm.weight") - 1.0, + (model_dim,), + "c_final_norm", + ) + # 4B model has no lm_head.weight. + if "lm_head.weight" in params: + add_data( + "lm_head.weight", + params.pop("lm_head.weight"), + (vocab_size, model_dim), + "lm_head", + ) + + for i in range(num_layers): + add_att_einsum(i) + add_gating_einsum(i) + add_qkv_einsum(i) + add_data( + f"{llm_prefix}layers.%d.mlp.down_proj.weight", + params.pop(f"{llm_prefix}layers.{i}.mlp.down_proj.weight"), + (model_dim, hidden_dim), + "linear_w", + i, + ) + # Qwen3 has only two pre norms. + add_data( + f"{llm_prefix}layers.%d.input_layernorm.weight", + params.pop(f"{llm_prefix}layers.{i}.input_layernorm.weight") - 1.0, + (model_dim,), + "pre_att_ns", + i, + ) + add_data( + f"{llm_prefix}layers.%d.post_attention_layernorm.weight", + params.pop(f"{llm_prefix}layers.{i}.post_attention_layernorm.weight") + - 1.0, + (model_dim,), + "pre_ff_ns", + i, + ) + + if has_qk_norm: + add_data( + f"{llm_prefix}layers.%d.self_attn.q_norm.weight", + params.pop(f"{llm_prefix}layers.{i}.self_attn.q_norm.weight") - 1.0, + (head_dim,), + "query_norm", + i, + ) + add_data( + f"{llm_prefix}layers.%d.self_attn.k_norm.weight", + params.pop(f"{llm_prefix}layers.{i}.self_attn.k_norm.weight") - 1.0, + (head_dim,), + "key_norm", + i, + ) + + if params: + print(f"WARNING: leftover params not consumed: {list(params.keys())[:10]}") + + sbs_config = configs.ModelConfig(model_specifier) + if tokenizer_file.endswith(".json"): + sbs_config.tokenizer_kind = configs.TokenizerKind.kHfBpe + tokenizer_blob = pack_bpe_tokenizer(tokenizer_file) + else: + raise ValueError("Qwen3 LM requires a HF BPE tokenizer.") + writer.write(sbs_config, tokenizer_blob) + + with open(csv_file, "w") as csv_handle: + csv.writer(csv_handle).writerows(metadata) + + def main(argv: Sequence[str]) -> None: if len(argv) > 1: raise app.UsageError("Too many command-line arguments.") @@ -926,10 +1158,14 @@ def main(argv: Sequence[str]) -> None: export_gemma3_lm_sbs( model_specifier, load_path, tokenizer_file, metadata_file, sbs_file ) + elif model_specifier.startswith("qwen3-"): + export_qwen3_lm_sbs( + model_specifier, load_path, tokenizer_file, metadata_file, sbs_file + ) else: raise app.UsageError( f"Unsupported model_specifier {model_specifier!r}. Expected a " - "'paligemma*' or 'gemma3-*-lm-*' specifier." + "'paligemma*', 'gemma3-*-lm-*' or 'qwen3-*' specifier." ) diff --git a/tokenizer/bpe_tokenizer.cc b/tokenizer/bpe_tokenizer.cc index d901e3ca..092db644 100644 --- a/tokenizer/bpe_tokenizer.cc +++ b/tokenizer/bpe_tokenizer.cc @@ -46,7 +46,7 @@ constexpr const char* kSpaceRepl = "\xe2\x96\x81"; // (the corresponding behavior is fixed); reserved so readers can branch later. constexpr uint32_t kFlagByteFallback = 1u << 0; constexpr uint32_t kFlagSpaceReplace = 1u << 1; - +constexpr uint32_t kFlagByteLevel = 1u << 2; class BpeTokenizer : public Tokenizer { public: @@ -69,6 +69,7 @@ class BpeTokenizer : public Tokenizer { }; void LoadByteTokens(); + void BuildByteLevelMaps(); const std::string& IdToToken(int id) const; int IdByteValue(int id) const; int MatchAddedToken(std::string_view input, size_t i, size_t* len) const; @@ -76,6 +77,9 @@ class BpeTokenizer : public Tokenizer { std::vector* ids) const; std::string Normalize(std::string_view text) const; void EncodeSpan(std::string_view text, std::vector* ids) const; + void EncodeSpanByteLevel(std::string_view text, std::vector* ids) const; + bool DecodeByteLevel(hwy::Span ids, + std::string& detokenized) const; void MergeSymbols(std::vector* sym_id) const; std::unordered_map vocab_; @@ -86,7 +90,11 @@ class BpeTokenizer : public Tokenizer { std::unordered_map added_tokens_; std::unordered_set added_first_bytes_; std::vector added_lengths_; // distinct, descending + std::vector byte_to_token_; + std::unordered_map unicode_to_byte_; + int unk_id_ = 0; + bool byte_level_ = false; }; @@ -123,6 +131,196 @@ uint64_t PairKey(int left, int right) { static_cast(right); } +struct CodePoint { + uint32_t cp; + uint32_t off; + uint32_t len; +}; + +std::vector DecodeUtf8(std::string_view s) { + std::vector out; + out.reserve(s.size()); + size_t i = 0; + while (i < s.size()) { + const unsigned char c0 = static_cast(s[i]); + size_t len = Utf8Len(c0); + if (i + len > s.size()) len = 1; // truncated: treat lead as a single byte + uint32_t cp = c0; + if (len == 2) { + cp = ((c0 & 0x1Fu) << 6) | (static_cast(s[i + 1]) & 0x3Fu); + } else if (len == 3) { + cp = ((c0 & 0x0Fu) << 12) | + ((static_cast(s[i + 1]) & 0x3Fu) << 6) | + (static_cast(s[i + 2]) & 0x3Fu); + } else if (len == 4) { + cp = ((c0 & 0x07u) << 18) | + ((static_cast(s[i + 1]) & 0x3Fu) << 12) | + ((static_cast(s[i + 2]) & 0x3Fu) << 6) | + (static_cast(s[i + 3]) & 0x3Fu); + } + out.push_back({cp, static_cast(i), static_cast(len)}); + i += len; + } + return out; +} + +std::string Utf8Encode(uint32_t cp) { + std::string o; + if (cp < 0x80) { + o.push_back(static_cast(cp)); + } else if (cp < 0x800) { + o.push_back(static_cast(0xC0 | (cp >> 6))); + o.push_back(static_cast(0x80 | (cp & 0x3F))); + } else if (cp < 0x10000) { + o.push_back(static_cast(0xE0 | (cp >> 12))); + o.push_back(static_cast(0x80 | ((cp >> 6) & 0x3F))); + o.push_back(static_cast(0x80 | (cp & 0x3F))); + } else { + o.push_back(static_cast(0xF0 | (cp >> 18))); + o.push_back(static_cast(0x80 | ((cp >> 12) & 0x3F))); + o.push_back(static_cast(0x80 | ((cp >> 6) & 0x3F))); + o.push_back(static_cast(0x80 | (cp & 0x3F))); + } + return o; +} + +bool IsAsciiWs(uint32_t cp) { + return cp == ' ' || cp == '\t' || cp == '\n' || cp == '\r' || cp == '\f' || + cp == '\v'; +} + +bool IsWhitespace(uint32_t cp) { + if (IsAsciiWs(cp)) return true; + switch (cp) { + case 0x85: // NEL + case 0xA0: // NBSP + case 0x1680: // OGHAM SPACE MARK + case 0x2028: // LINE SEPARATOR + case 0x2029: // PARAGRAPH SEPARATOR + case 0x202F: // NARROW NBSP + case 0x205F: // MEDIUM MATHEMATICAL SPACE + case 0x3000: // IDEOGRAPHIC SPACE + return true; + default: + break; + } + return cp >= 0x2000 && cp <= 0x200A; +} + +bool IsNumber(uint32_t cp) { return cp >= '0' && cp <= '9'; } + +bool IsLetter(uint32_t cp) { + if ((cp >= 'a' && cp <= 'z') || (cp >= 'A' && cp <= 'Z')) return true; + if (cp < 0x80) return false; + return !IsWhitespace(cp) && !IsNumber(cp); +} + +uint32_t AsciiLower(uint32_t c) { return (c >= 'A' && c <= 'Z') ? c + 32 : c; } + +std::vector Gpt2Split(std::string_view text) { + std::vector pieces; + const std::vector cps = DecodeUtf8(text); + const size_t n = cps.size(); + const auto range = [&](size_t a, size_t b) -> std::string { + if (a >= b) return std::string(); + const size_t off = cps[a].off; + const size_t end = cps[b - 1].off + cps[b - 1].len; + return std::string(text.substr(off, end - off)); + }; + const auto is_d = [](uint32_t c) { + return !IsWhitespace(c) && !IsLetter(c) && !IsNumber(c); + }; + + size_t k = 0; + while (k < n) { + // A: contractions. + if (cps[k].cp == '\'' && k + 1 < n) { + const uint32_t c1 = AsciiLower(cps[k + 1].cp); + if (k + 2 < n) { + const uint32_t c2 = AsciiLower(cps[k + 2].cp); + if ((c1 == 'r' && c2 == 'e') || (c1 == 'v' && c2 == 'e') || + (c1 == 'l' && c2 == 'l')) { + pieces.push_back(range(k, k + 3)); + k += 3; + continue; + } + } + if (c1 == 's' || c1 == 't' || c1 == 'm' || c1 == 'd') { + pieces.push_back(range(k, k + 2)); + k += 2; + continue; + } + } + // B: optional non-letter/number lead, then one or more letters. + { + size_t p = k; + if (p < n && cps[p].cp != '\r' && cps[p].cp != '\n' && + !IsLetter(cps[p].cp) && !IsNumber(cps[p].cp) && p + 1 < n && + IsLetter(cps[p + 1].cp)) { + ++p; + } + if (p < n && IsLetter(cps[p].cp)) { + while (p < n && IsLetter(cps[p].cp)) ++p; + pieces.push_back(range(k, p)); + k = p; + continue; + } + } + // C: a single number. + if (IsNumber(cps[k].cp)) { + pieces.push_back(range(k, k + 1)); + ++k; + continue; + } + // D: optional space, then symbols, then trailing newlines. + { + const bool space = cps[k].cp == ' '; + const size_t q = space ? k + 1 : k; + if (q < n && is_d(cps[q].cp)) { + size_t r = q; + while (r < n && is_d(cps[r].cp)) ++r; + while (r < n && (cps[r].cp == '\r' || cps[r].cp == '\n')) ++r; + pieces.push_back(range(k, r)); + k = r; + continue; + } + } + // E/F/G: whitespace runs. + { + size_t p = k; + while (p < n && IsWhitespace(cps[p].cp)) ++p; + if (p > k) { + bool found_nl = false; + size_t last_nl = k; + for (size_t j = k; j < p; ++j) { + if (cps[j].cp == '\n' || cps[j].cp == '\r') { + found_nl = true; + last_nl = j; + } + } + if (found_nl) { // E: `\s*[\r\n]+` + pieces.push_back(range(k, last_nl + 1)); + k = last_nl + 1; + } else if (p == n) { // F: trailing whitespace at end of text + pieces.push_back(range(k, p)); + k = p; + } else if (p - k >= 2) { // F: leave the last space for the next word + pieces.push_back(range(k, p - 1)); + k = p - 1; + } else { // G: a lone space before a non-space + pieces.push_back(range(k, p)); + k = p; + } + continue; + } + } + // Safety net: never stall. + pieces.push_back(range(k, k + 1)); + ++k; + } + return pieces; +} + // A candidate merge popped from the priority queue: lower `rank` merges first, // ties broken by leftmost position, matching HuggingFace's BPE merge order. struct QueuedMerge { @@ -200,11 +398,15 @@ std::string BpeTokenizer::Serialize() const { BpeTokenizerBlob blob; blob.unk_id = static_cast(unk_id_); blob.vocab = id_to_token_; - blob.flags = kFlagSpaceReplace; - for (int b : byte_id_) { - if (b >= 0) { - blob.flags |= kFlagByteFallback; - break; + if (byte_level_) { + blob.flags |= kFlagByteLevel; + } else { + blob.flags = kFlagSpaceReplace; + for (int b : byte_id_) { + if (b >= 0) { + blob.flags |= kFlagByteFallback; + break; + } } } @@ -235,6 +437,9 @@ bool BpeTokenizer::Encode(std::string_view input, bool BpeTokenizer::Decode(hwy::Span ids, std::string& detokenized) const { + if (byte_level_) { + return DecodeByteLevel(ids, detokenized); + } detokenized.clear(); std::string pending_bytes; // accumulates consecutive byte-fallback tokens for (int id : ids) { @@ -329,7 +534,64 @@ bool BpeTokenizer::Deserialize(std::string_view data) { added_lengths_.assign(lengths.begin(), lengths.end()); std::sort(added_lengths_.begin(), added_lengths_.end(), std::greater<>()); - LoadByteTokens(); + byte_level_ = (blob.flags & kFlagByteLevel) != 0; + if (byte_level_) { + BuildByteLevelMaps(); + } else { + LoadByteTokens(); + } + return true; +} + +// Builds the GPT-2 `bytes_to_unicode` alphabet. +void BpeTokenizer::BuildByteLevelMaps() { + byte_to_token_.assign(256, std::string()); + unicode_to_byte_.clear(); + unicode_to_byte_.reserve(512); + std::vector direct(256, false); + const auto mark = [&](int lo, int hi) { + for (int b = lo; b <= hi; ++b) direct[b] = true; + }; + mark('!', '~'); // 0x21..0x7E + mark(0xA1, 0xAC); // ¡..¬ + mark(0xAE, 0xFF); // ®..ÿ + int n = 0; + for (int b = 0; b < 256; ++b) { + const uint32_t cp = direct[b] ? static_cast(b) : (256u + n++); + byte_to_token_[b] = Utf8Encode(cp); + unicode_to_byte_[cp] = b; + } +} + +void BpeTokenizer::EncodeSpanByteLevel(std::string_view text, + std::vector* ids) const { + for (const std::string& piece : Gpt2Split(text)) { + std::vector sym_id; + sym_id.reserve(piece.size()); + for (char c : piece) { + const std::string& tok = byte_to_token_[static_cast(c)]; + const auto it = vocab_.find(tok); + sym_id.push_back(it != vocab_.end() ? it->second : unk_id_); + } + MergeSymbols(&sym_id); + ids->insert(ids->end(), sym_id.begin(), sym_id.end()); + } +} + +bool BpeTokenizer::DecodeByteLevel(hwy::Span ids, + std::string& detokenized) const { + detokenized.clear(); + for (int id : ids) { + const std::string& tok = IdToToken(id); + for (const CodePoint& c : DecodeUtf8(tok)) { + const auto it = unicode_to_byte_.find(c.cp); + if (it != unicode_to_byte_.end()) { + detokenized.push_back(static_cast(it->second)); + } else { + detokenized.append(tok, c.off, c.len); + } + } + } return true; } @@ -403,6 +665,10 @@ std::string BpeTokenizer::Normalize(std::string_view text) const { void BpeTokenizer::EncodeSpan(std::string_view text, std::vector* ids) const { if (text.empty()) return; + if (byte_level_) { + EncodeSpanByteLevel(text, ids); + return; + } const std::string norm = Normalize(text); // Initial symbols: one per whole-character vocab entry, else byte-fallback.