// Modified from xgrammar/nanobind/nanobind.cc from xgrammar project. /*! * Copyright (c) 2024 by Contributors * \file xgrammar/nanobind/nanobind.cc */ #include #include #include #include #include #include #include #include #include #include "src/turbomind/core/check.h" namespace py = pybind11; using namespace xgrammar; using namespace pybind11::literals; namespace { static const std::vector CommonEncodedVocabType(const py::typing::List>& lst) { std::vector out; out.reserve(lst.size()); for (const auto& h : lst) { if (py::isinstance(h)) { out.emplace_back(h.cast()); } else if (py::isinstance(h)) { out.emplace_back(h.cast()); } else { throw std::invalid_argument("encoded_vocab items must be str or bytes"); } } return out; } TokenizerInfo TokenizerInfo_Init(const std::vector& encoded_vocab, int vocab_type, std::optional vocab_size, std::optional> stop_token_ids, bool add_prefix_space) { TM_CHECK(vocab_type == 0 || vocab_type == 1 || vocab_type == 2) << "Invalid vocab type: " << vocab_type; return TokenizerInfo( encoded_vocab, static_cast(vocab_type), vocab_size, stop_token_ids, add_prefix_space); } int TokenizerInfo_GetVocabType(const TokenizerInfo& tokenizer) { return static_cast(tokenizer.GetVocabType()); } std::vector TokenizerInfo_GetDecodedVocab(const TokenizerInfo& tokenizer) { const auto& decoded_vocab = tokenizer.GetDecodedVocab(); std::vector py_result; py_result.reserve(decoded_vocab.size()); for (const auto& item : decoded_vocab) { py_result.emplace_back(py::bytes(item.c_str())); } return py_result; } } // namespace PYBIND11_MODULE(_xgrammar, m) { py::class_>(m, "TokenizerInfo") .def(py::init([](const py::typing::List>& encoded_vocab, int vocab_type, std::optional vocab_size, std::optional> stop_token_ids, bool add_prefix_space) { return TokenizerInfo{TokenizerInfo_Init(CommonEncodedVocabType(encoded_vocab), vocab_type, vocab_size, std::move(stop_token_ids), add_prefix_space)}; }), py::arg("encoded_vocab"), py::arg("vocab_type"), py::arg("vocab_size") = py::none(), py::arg("stop_token_ids") = py::none(), py::arg("add_prefix_space")) .def_property_readonly("vocab_type", &TokenizerInfo_GetVocabType) .def_property_readonly("vocab_size", &TokenizerInfo::GetVocabSize) .def_property_readonly("add_prefix_space", &TokenizerInfo::GetAddPrefixSpace) .def_property_readonly("decoded_vocab", &TokenizerInfo_GetDecodedVocab) .def_property_readonly("stop_token_ids", &TokenizerInfo::GetStopTokenIds) .def_property_readonly("special_token_ids", &TokenizerInfo::GetSpecialTokenIds) .def("dump_metadata", &TokenizerInfo::DumpMetadata) .def_static("from_vocab_and_metadata", [](const py::typing::List>& encoded_vocab, const std::string& metadata) { return TokenizerInfo::FromVocabAndMetadata(CommonEncodedVocabType(encoded_vocab), metadata); }) .def_static("_detect_metadata_from_hf", &TokenizerInfo::DetectMetadataFromHF); py::class_(m, "CompiledGrammar"); py::class_ pyGrammarCompiler(m, "GrammarCompiler"); pyGrammarCompiler .def(py::init(), py::arg("tokenizer_info"), py::arg("max_threads") = 8, py::arg("cache_enabled") = true, py::arg("max_memory_bytes") = -1) .def("compile_json_schema", &GrammarCompiler::CompileJSONSchema, py::call_guard(), py::arg("schema"), py::arg("any_whitespace") = false, py::arg("indent") = py::none(), py::arg("separators") = py::none(), py::arg("strict_mode") = true, py::arg("max_whitespace_cnt") = py::none()) .def("compile_regex", &GrammarCompiler::CompileRegex, py::call_guard(), py::arg("schema")); }