#include #include "py_backend.h" #include #include #include PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { std::string name = std::string("TableManager"); py::class_(m, name.c_str()) .def(py::init([](const py::array_t& seq_lens, const py::array_t& group_ids, const py::array_t& merge_orders, size_t window_size, size_t cache_id_offset, size_t detach_cache_id_offset, vector>& span_ids) { return new TableManager(seq_lens, group_ids, merge_orders, window_size, cache_id_offset, detach_cache_id_offset, span_ids); })) .def("step", &TableManager::step) .def("root_ids", &TableManager::root_ids) .def("is_finished", &TableManager::is_finished) .def("prepare_bilm", &TableManager::prepare_bilm) .def("prepare_generation", &TableManager::prepare_generation) .def("batch_size", &TableManager::batch_size); name = std::string("SpanTokenizer"); py::class_(m, name.c_str()) .def(py::init([](vector>& dictionary, int max_entry_id) { return new SpanTokenizer(dictionary, max_entry_id); })) .def("tokenize", &SpanTokenizer::tokenize); }