55using integer_type = std::uint64_t;
65template <
typename Metric>
68 integer_type operator()(std::string_view s, std::string_view t)
const {
69 return (
static_cast<Metric
const *
>(
this))->compute_distance(s, t);
83 integer_type compute_distance(std::string_view s, std::string_view t)
const noexcept {
84 return s.length() > t.length() ? s.length() - t.length() : t.length() - s.length();
96 integer_type compute_distance(std::string_view, std::string_view)
const noexcept {
97 return integer_type{1};
109 integer_type m_alphabet_size;
112 explicit LeeDistance(integer_type alphabet_size = BK_LEE_ALPHABET_SIZE)
113 : m_alphabet_size(alphabet_size) {};
114 integer_type compute_distance(std::string_view s, std::string_view t)
const noexcept {
115 const integer_type M = s.length(), N = t.length();
117 return std::numeric_limits<integer_type>::max();
119 const integer_type m_comparison_size = M;
120 integer_type counter = 0, diff;
121 for (integer_type i = 0; i < m_comparison_size; ++i) {
122 diff = std::abs(s[i] - t[i]);
123 counter += std::min(diff, m_alphabet_size - diff);
145 mutable std::vector<integer_type> m_current, m_previous;
149 : m_current(initial_size), m_previous(initial_size) {};
150 integer_type compute_distance(std::string_view s, std::string_view t)
const noexcept {
151 const integer_type M = s.length(), N = t.length();
152 if (M == 0 || N == 0) {
155 if (m_current.size() <= N || m_previous.size() <= N) {
156 m_current.resize(N + 1);
157 m_previous.resize(N + 1);
159 std::fill(m_previous.begin(), m_previous.end(), 0);
160 for (integer_type i = 1; i <= M; ++i) {
161 for (integer_type j = 1; j <= N; ++j) {
162 if (s[i - 1] == t[j - 1]) {
163 m_current[j] = m_previous[j - 1] + 1;
165 m_current[j] = std::max(m_previous[j], m_current[j - 1]);
168 m_previous = m_current;
170 return M + N - 2 * m_previous[N];
184 integer_type compute_distance(std::string_view s, std::string_view t)
const noexcept {
185 const integer_type M = s.length(), N = t.length();
187 return std::numeric_limits<integer_type>::max();
189 const integer_type m_comparison_size = M;
190 integer_type counter = 0;
191 for (integer_type i = 0; i < m_comparison_size; ++i) {
192 counter += (s[i] != t[i]);
223 mutable std::vector<std::vector<integer_type>> m_matrix;
226 explicit EditDistance(
size_t initial_size = BK_ED_MATRIX_INITIAL_SIZE)
227 : m_matrix(initial_size, std::vector<integer_type>(initial_size)) {};
228 integer_type compute_distance(std::string_view s, std::string_view t)
const noexcept {
229 const integer_type M = s.length(), N = t.length();
230 if (M == 0 || N == 0) {
233 if (m_matrix.size() <= M || m_matrix[0].size() <= N) {
234 std::vector<std::vector<integer_type>> a_matrix(M + 1,
235 std::vector<integer_type>(N + 1));
236 m_matrix.swap(a_matrix);
238 for (integer_type i = 1; i <= M; ++i) {
241 for (integer_type j = 1; j <= N; ++j) {
244 for (integer_type j = 1; j <= N; ++j) {
245 for (integer_type i = 1; i <= M; ++i) {
246 m_matrix[i][j] = std::min(
247 {m_matrix[i][j - 1] + 1 , m_matrix[i - 1][j] + 1 ,
248 m_matrix[i - 1][j - 1] + (s[i - 1] != t[j - 1]) });
251 return m_matrix[M][N];
261 mutable std::vector<std::vector<integer_type>> m_matrix;
265 : m_matrix(initial_size, std::vector<integer_type>(initial_size)) {};
266 integer_type compute_distance(std::string_view s, std::string_view t)
const noexcept {
267 const integer_type M = s.length(), N = t.length();
268 if (M == 0 || N == 0) {
271 if (m_matrix.size() <= M || m_matrix[0].size() <= N) {
272 std::vector<std::vector<integer_type>> a_matrix(M + 1,
273 std::vector<integer_type>(N + 1));
274 m_matrix.swap(a_matrix);
276 for (integer_type i = 1; i <= M; ++i) {
279 for (integer_type j = 1; j <= N; ++j) {
282 for (integer_type j = 1; j <= N; ++j) {
283 for (integer_type i = 1; i <= M; ++i) {
284 m_matrix[i][j] = std::min(
285 {m_matrix[i][j - 1] + 1 , m_matrix[i - 1][j] + 1 ,
286 m_matrix[i - 1][j - 1] + (s[i - 1] == t[j - 1] ? 0 : 1) });
287 if (i > 1 && j > 1 && s[i - 1] == t[j - 2] && s[i - 2] == t[j - 1]) {
289 std::min(m_matrix[i][j], m_matrix[i - 2][j - 2] + 1 );
293 return m_matrix[M][N];
300std::false_type is_metric_impl(...);
302template <
typename Metric>
305template <
typename Metric>
306using is_metric =
decltype(is_metric_impl(std::declval<Metric &>()));
309template <
typename Metric>
311template <
typename Metric>
314using ResultEntry = std::pair<std::string, integer_type>;
315using ResultList = std::vector<ResultEntry>;
317template <
typename Metric>
319 friend class BKTree<Metric>;
320 using metric_type = Metric;
321 using node_type = BKTreeNode<metric_type>;
323 BKTreeNode(std::string_view value) : m_word(value) {}
324 bool _insert(std::string_view value,
const metric_type &distance);
325 bool _erase(std::string_view value,
const metric_type &distance);
326 void _find(ResultList &output, std::string_view value, integer_type limit,
327 const metric_type &metric)
const;
328 ResultList _find_wrapper(std::string_view value, integer_type limit,
329 const metric_type &metric)
const;
331 std::map<integer_type, std::unique_ptr<node_type>> m_children;
334 friend std::ostream &operator<<(std::ostream &oss,
const BKTreeNode &node) {
340 std::string_view word() const noexcept {
return m_word; }
346template <
typename Metric>
348 static_assert(helpers::is_metric<Metric>::value,
"Metric must be of type Distance");
350 using metric_type = Metric;
351 using node_type =
typename BKTreeNode<metric_type>::node_type;
359 using iterator_category = std::forward_iterator_tag;
360 using difference_type = std::ptrdiff_t;
361 using value_type = std::unique_ptr<node_type>;
362 using pointer = std::unique_ptr<node_type> *;
363 using reference = std::unique_ptr<node_type> &;
367 Iterator(pointer ptr) : m_pointer(ptr) {}
369 pointer operator->() {
return m_pointer; }
371 reference operator*()
const {
return *m_pointer; }
374 if (m_pointer ==
nullptr) {
375 throw std::out_of_range(
"No more tree node");
377 for (
auto &[_, child] : (*m_pointer)->m_children) {
378 m_queue.push(&child);
380 if (m_queue.empty()) {
383 m_pointer = m_queue.front();
396 return a.m_pointer == b.m_pointer;
400 return a.m_pointer != b.m_pointer;
405 std::queue<pointer> m_queue;
409 BKTree(
const metric_type &distance = Metric())
410 : m_root(nullptr), m_metric(distance), m_tree_size(BK_TREE_INITIAL_SIZE) {}
412 BKTree(std::initializer_list<std::string_view> list)
413 : m_root(nullptr), m_metric(Metric()), m_tree_size(BK_TREE_INITIAL_SIZE) {
414 for (
auto &str : list) {
419 BKTree(
const BKTree &other) : BKTree(other.m_metric) {
420 if (other.m_root ==
nullptr) {
423 std::queue<std::unique_ptr<node_type>
const *> bq;
424 bq.push(&(other.m_root));
425 while (!bq.empty()) {
426 auto *nptr = bq.front();
428 this->insert((*nptr)->m_word);
429 for (
auto &[_, child_node] : (*nptr)->m_children) {
430 bq.push(&child_node);
435 BKTree(BKTree &&other) noexcept
436 : m_root(std::exchange(other.m_root,
nullptr)), m_tree_size(other.m_tree_size) {}
438 BKTree &operator=(
const BKTree &other) {
439 if (
this == &other) {
443 std::swap(m_root, temp.m_root);
444 std::swap(m_tree_size, temp.m_tree_size);
448 BKTree &operator=(BKTree &&other)
noexcept {
449 std::swap(m_root, other.m_root);
450 std::swap(m_tree_size, other.m_tree_size);
456 bool insert(std::string_view value);
457 bool erase(std::string_view value);
458 size_t size() const noexcept {
return m_tree_size; }
459 bool empty() const noexcept {
return m_tree_size == 0; }
460 [[nodiscard]] ResultList find(std::string_view value, integer_type limit)
const;
461 [[nodiscard]] ResultList find(std::string_view value,
int limit)
const;
462 template <
typename Integer>
463 requires std::is_integral_v<Integer>
464 [[nodiscard]] ResultList find(std::string_view value, Integer limit)
const {
465 if constexpr (std::is_signed_v<Integer>) {
470 return find(value,
static_cast<integer_type
>(limit));
473 Iterator begin() {
return Iterator(&m_root); }
474 Iterator end() {
return Iterator(); }
477 std::unique_ptr<node_type> m_root;
478 const metric_type m_metric;
482template <
typename Metric>
483bool BKTreeNode<Metric>::_insert(std::string_view value,
484 const metric_type &distance_metric) {
485 const integer_type distance_between = distance_metric(value, m_word);
486 if (distance_between == std::numeric_limits<integer_type>::max()) {
489 auto it = m_children.find(distance_between);
490 if (it == m_children.end()) {
491 m_children.emplace(std::make_pair(
492 distance_between, std::unique_ptr<node_type>(
new node_type(value))));
495 return it->second->_insert(value, distance_metric);
498template <
typename Metric>
499bool BKTreeNode<Metric>::_erase(std::string_view value,
500 const metric_type &distance_metric) {
502 const integer_type distance_between = distance_metric(value, m_word);
503 auto it = m_children.find(distance_between);
504 if (it != m_children.end()) {
505 if (it->second->m_word == value) {
506 auto node = std::move(it->second);
507 m_children.erase(it);
508 std::queue<std::unique_ptr<node_type>
const *> bq;
509 for (
auto const &[_, child_node] : node->m_children) {
510 bq.push(&child_node);
512 while (!bq.empty()) {
513 auto *node = bq.front();
515 for (
auto const &[_, child_node] : (*node)->m_children) {
516 bq.push(&child_node);
518 _insert((*node)->m_word, distance_metric);
522 erased = it->second->_erase(value, distance_metric);
525 for (
auto const &[_, child] : m_children) {
526 if (child->_erase(value, distance_metric)) {
534template <
typename Metric>
535void BKTreeNode<Metric>::_find(ResultList &output, std::string_view value,
536 integer_type limit,
const metric_type &metric)
const {
537 const integer_type distance = metric(value, m_word);
538 if (distance != std::numeric_limits<integer_type>::max() && distance <= limit) {
539 output.push_back({m_word, distance});
541 for (
auto const &[dist, node] : m_children) {
542 const integer_type difference = dist > distance ? dist - distance : distance - dist;
543 if (difference <= limit) {
544 node->_find(output, value, limit, metric);
549template <
typename Metric>
550ResultList BKTreeNode<Metric>::_find_wrapper(std::string_view value, integer_type limit,
551 const metric_type &metric)
const {
553 _find(output, value, limit, metric);
557template <
typename Metric>
558bool BKTree<Metric>::insert(std::string_view value) {
559 bool inserted =
false;
560 if (m_root ==
nullptr) {
561 m_root = std::unique_ptr<node_type>(
new node_type(value));
564 }
else if (m_root->_insert(value, m_metric)) {
571template <
typename Metric>
572bool BKTree<Metric>::erase(std::string_view value) {
574 if (m_root ==
nullptr) {
576 }
else if (m_root->m_word == value) {
577 if (m_tree_size > 1) {
578 auto &replacement_node = m_root->m_children.begin()->second;
579 std::queue<std::unique_ptr<node_type>
const *> bq;
580 for (
bool first =
true;
auto const &[_, node] : m_root->m_children) {
587 while (!bq.empty()) {
588 auto node = bq.front();
590 for (
auto const &[_, child] : (*node)->m_children) {
593 replacement_node->_insert((*node)->m_word, m_metric);
595 m_root = std::move(replacement_node);
597 m_root.reset(
nullptr);
601 }
else if (m_root->_erase(value, m_metric)) {
608template <
typename Metric>
609ResultList BKTree<Metric>::find(std::string_view value, integer_type limit)
const {
610 if (m_root ==
nullptr) {
613 return m_root->_find_wrapper(value, limit, m_metric);
616template <
typename Metric>
617ResultList BKTree<Metric>::find(std::string_view value,
int limit)
const {
621 return find(value,
static_cast<integer_type
>(limit));