diff --git a/modules/dnn/src/tokenizer/tokenizer.cpp b/modules/dnn/src/tokenizer/tokenizer.cpp index 2f332604d1..5d933754f0 100644 --- a/modules/dnn/src/tokenizer/tokenizer.cpp +++ b/modules/dnn/src/tokenizer/tokenizer.cpp @@ -73,6 +73,35 @@ struct GemmaBpeTokenizerImpl : public Tokenizer::Impl { } }; +struct SentencePieceTokenizerImpl : public Tokenizer::Impl { + CoreGemmaBPE model; + std::unordered_set allowedSpecial; + int bosTokenId; + + explicit SentencePieceTokenizerImpl(CoreGemmaBPE m, + std::unordered_set special = {}, + int bos = -1) + : model(std::move(m)), allowedSpecial(std::move(special)), bosTokenId(bos) {} + + std::vector encode(const std::string& text) override { + std::vector ids = model.encode(text, allowedSpecial); + // HuggingFace SentencePiece tokenizers (Gemma2) prepend automatically + if (bosTokenId >= 0) { + ids.insert(ids.begin(), bosTokenId); + } + return ids; + } + + std::string decode(const std::vector& tokens) override { + // Skip the bos token if present at the beginning + if (bosTokenId >= 0 && !tokens.empty() && tokens.front() == bosTokenId) { + std::vector stripped(tokens.begin() + 1, tokens.end()); + return model.decode(stripped); + } + return model.decode(tokens); + } +}; + static Ptr buildGemmaFromJson( const std::string& json_path, std::unordered_set* outSpecial = nullptr) { @@ -147,6 +176,93 @@ static Ptr buildGemmaFromJson( return makePtr(std::move(gemma), std::move(special)); } +static Ptr buildSentencePieceFromJson( + const std::string& json_path, + std::unordered_set* outSpecial = nullptr) { + + cv::FileStorage fs(json_path, cv::FileStorage::READ | cv::FileStorage::FORMAT_JSON); + if (!fs.isOpened()) + CV_Error(cv::Error::StsError, "Failed to open tokenizer.json: " + json_path); + + cv::FileNode model_node = fs["model"]; + CV_CheckFalse(model_node.empty(), "tokenizer.json missing 'model'"); + + std::string model_type; + model_node["type"] >> model_type; + if (model_type != "BPE") + CV_Error(cv::Error::StsError, + "Expected BPE model in tokenizer.json for SentencePiece, got: " + model_type); + + CoreGemmaBPE gemma; + + cv::FileNode vocab_node = model_node["vocab"]; + CV_CheckFalse(vocab_node.empty(), "tokenizer.json model missing 'vocab'"); + + int maxId = -1; + for (auto it = vocab_node.begin(); it != vocab_node.end(); ++it) { + cv::FileNode entry = *it; + std::string piece = entry.name(); + int id = (int)entry; + if (id > maxId) maxId = id; + gemma.pieceToId[piece] = id; + } + + gemma.idToPiece.resize(maxId + 1); + for (const auto& kv : gemma.pieceToId) + gemma.idToPiece[kv.second] = kv.first; + + cv::FileNode merges_node = model_node["merges"]; + if (!merges_node.empty()) { + uint32_t rank = 0; + for (auto it = merges_node.begin(); it != merges_node.end(); ++it) { + cv::FileNode entry = *it; + std::string a, b; + if (entry.isString()) { + // SentencePiece format: "a b" (single string with space separator) + std::string merge_str; + entry >> merge_str; + size_t sp = merge_str.find(' '); + if (sp == std::string::npos) { + ++rank; + continue; + } + a = merge_str.substr(0, sp); + b = merge_str.substr(sp + 1); + } else if (entry.size() == 2) { + // Array format: ["a", "b"] + entry[0] >> a; + entry[1] >> b; + } else { + ++rank; + continue; + } + gemma.addMerge(a, b, rank); + ++rank; + } + } + + std::unordered_set special; + int bosId = -1; + cv::FileNode added = fs["added_tokens"]; + if (!added.empty()) { + for (auto it = added.begin(); it != added.end(); ++it) { + cv::FileNode t = *it; + int id = -1; t["id"] >> id; + std::string content; t["content"] >> content; + if (id >= 0 && !content.empty()) { + gemma.specialToId[content] = id; + gemma.idToSpecial[id] = content; + special.insert(content); + if (outSpecial) outSpecial->insert(content); + // Capture token id for SentencePiece auto-prepend + if (content == "") bosId = id; + } + } + } + + return makePtr(std::move(gemma), std::move(special), bosId); +} + static void registerDefaultTokenizers() { auto& reg = tokenizerRegistry(); if (reg.find("BPE") == reg.end()) { @@ -175,6 +291,14 @@ static void registerDefaultTokenizers() { return buildGemmaFromJson(tok_json, &special); }; } + + if (reg.find("SentencePiece") == reg.end()) { + reg["SentencePiece"] = [](const FileStorage& /*cfg*/, const std::string& dir) -> Ptr { + std::string tok_json = dir + "tokenizer.json"; + std::unordered_set special; + return buildSentencePieceFromJson(tok_json, &special); + }; + } } Tokenizer::Tokenizer() : impl_(nullptr) {} @@ -275,7 +399,7 @@ Tokenizer Tokenizer::load(const std::string& model_config) { auto it = reg.find(methodType); if (it == reg.end()) CV_Error(cv::Error::StsError, - "Unsupported tokenizer method: '" + methodType + "'. Supported: BPE, Gemma"); + "Unsupported tokenizer method: '" + methodType + "'. Supported: BPE, Gemma, SentencePiece"); Tokenizer tok; tok.impl_ = it->second(cfg, dir); diff --git a/modules/dnn/test/test_tokenizer.cpp b/modules/dnn/test/test_tokenizer.cpp index ad0965f94f..38e3005df7 100644 --- a/modules/dnn/test/test_tokenizer.cpp +++ b/modules/dnn/test/test_tokenizer.cpp @@ -197,4 +197,44 @@ TEST(Tokenizer_Gemma, Tokenizer_Gemma3_Roundtrip) { } } +// Gemma2 tests (SentencePiece tokenizer) +TEST(Tokenizer_SentencePiece, Tokenizer_Gemma2_English) { + std::string model = _tf("gemma2/config.json"); + Tokenizer tok = Tokenizer::load(model); + EXPECT_EQ(tok.encode("Hello world"), (std::vector{2, 4521, 2134})); +} + +TEST(Tokenizer_SentencePiece, Tokenizer_Gemma2_Phrase) { + std::string model = _tf("gemma2/config.json"); + Tokenizer tok = Tokenizer::load(model); + EXPECT_EQ(tok.encode("the quick brown fox"), + (std::vector{2, 1175, 4320, 8426, 25341})); +} + +TEST(Tokenizer_SentencePiece, Tokenizer_Gemma2_Mixed) { + std::string model = _tf("gemma2/config.json"); + Tokenizer tok = Tokenizer::load(model); + EXPECT_EQ(tok.encode("OpenCV"), (std::vector{2, 6047, 17813})); +} + +TEST(Tokenizer_SentencePiece, Tokenizer_Gemma2_Numbers) { + std::string model = _tf("gemma2/config.json"); + Tokenizer tok = Tokenizer::load(model); + EXPECT_EQ(tok.encode("2024"), (std::vector{2, 235284, 235276, 235284, 235310})); +} + +TEST(Tokenizer_SentencePiece, Tokenizer_Gemma2_Roundtrip) { + std::string model = _tf("gemma2/config.json"); + Tokenizer tok = Tokenizer::load(model); + std::vector cases = { + "Hello world", + "the quick brown fox", + "OpenCV", + "hello world", + }; + for (const auto& text : cases) { + EXPECT_EQ(tok.decode(tok.encode(text)), text); + } +} + }}