|
1 | 1 | import copy |
| 2 | +import importlib |
2 | 3 | import re |
3 | 4 | import traceback |
4 | 5 |
|
|
13 | 14 | from memos.memories.textual.tree_text_memory.retrieve.bm25_util import EnhancedBM25 |
14 | 15 | from memos.memories.textual.tree_text_memory.retrieve.retrieve_utils import ( |
15 | 16 | FastTokenizer, |
| 17 | + StopwordManager, |
16 | 18 | cosine_similarity_matrix, |
17 | 19 | detect_lang, |
18 | 20 | find_best_unrelated_subgroup, |
|
33 | 35 |
|
34 | 36 |
|
35 | 37 | logger = get_logger(__name__) |
36 | | -KEYWORD_EXTRACT_TOP_K = 12 |
| 38 | +KEYWORD_EXTRACT_TOP_K = 3 |
| 39 | +KEYWORD_ALLOW_POS = ("n", "nr", "nrt", "ns", "nt", "nz", "vn", "v", "t", "eng", "m") |
37 | 40 | COT_DICT = { |
38 | 41 | "fine": {"en": COT_PROMPT, "zh": COT_PROMPT_ZH}, |
39 | 42 | "fast": {"en": SIMPLE_COT_PROMPT, "zh": SIMPLE_COT_PROMPT_ZH}, |
@@ -516,27 +519,85 @@ def _require_keyword_user_name(user_name: str | None) -> str: |
516 | 519 | ) |
517 | 520 | return normalized_user_name |
518 | 521 |
|
| 522 | + @staticmethod |
| 523 | + def _is_keyword_stopword(term: str) -> bool: |
| 524 | + normalized = term.strip() |
| 525 | + return not normalized or StopwordManager.is_search_stopword(normalized) |
| 526 | + |
| 527 | + @staticmethod |
| 528 | + def _normalize_keyword_term(term: str) -> str: |
| 529 | + normalized = str(term).strip() |
| 530 | + if re.fullmatch(r"[A-Za-z][A-Za-z0-9]*(?:[._+\-/][A-Za-z0-9]+)*", normalized): |
| 531 | + return normalized.lower() |
| 532 | + return normalized |
| 533 | + |
| 534 | + @staticmethod |
| 535 | + def _keyword_extract_top_k(query: str, language: str) -> int: |
| 536 | + cleaned_query = query.strip() |
| 537 | + if not cleaned_query: |
| 538 | + return 0 |
| 539 | + if len(cleaned_query) <= 12: |
| 540 | + return 1 |
| 541 | + if language != "zh": |
| 542 | + token_count = len(re.findall(r"\b[a-zA-Z0-9]+\b", cleaned_query)) |
| 543 | + return 2 if token_count <= 8 else KEYWORD_EXTRACT_TOP_K |
| 544 | + if len(cleaned_query) <= 120: |
| 545 | + return 2 |
| 546 | + return KEYWORD_EXTRACT_TOP_K |
| 547 | + |
| 548 | + @classmethod |
| 549 | + def _rank_english_keyword_terms(cls, terms: list[str]) -> list[str]: |
| 550 | + term_stats: dict[str, dict[str, int | str]] = {} |
| 551 | + for index, term in enumerate(terms): |
| 552 | + normalized_term = cls._normalize_keyword_term(term) |
| 553 | + if cls._is_keyword_stopword(normalized_term): |
| 554 | + continue |
| 555 | + key = normalized_term.lower() |
| 556 | + if key not in term_stats: |
| 557 | + term_stats[key] = {"term": normalized_term, "index": index, "count": 0} |
| 558 | + term_stats[key]["count"] = int(term_stats[key]["count"]) + 1 |
| 559 | + |
| 560 | + def score(item: tuple[str, dict[str, int | str]]) -> tuple[float, int]: |
| 561 | + _, data = item |
| 562 | + term = str(data["term"]) |
| 563 | + count = int(data["count"]) |
| 564 | + term_score = count * 3.0 + min(len(term), 16) * 0.1 |
| 565 | + if any(ch.isdigit() for ch in term): |
| 566 | + term_score += 1.0 |
| 567 | + if len(term) <= 2: |
| 568 | + term_score -= 0.5 |
| 569 | + return (-term_score, int(data["index"])) |
| 570 | + |
| 571 | + return [str(data["term"]) for _, data in sorted(term_stats.items(), key=score)] |
| 572 | + |
519 | 573 | def _extract_weighted_keyword_terms(self, query: str) -> list[str]: |
520 | | - if detect_lang(query) == "zh": |
521 | | - import jieba.analyse |
| 574 | + language = detect_lang(query) |
| 575 | + keyword_top_k = self._keyword_extract_top_k(query, language) |
| 576 | + if keyword_top_k <= 0: |
| 577 | + return [] |
| 578 | + |
| 579 | + if language == "zh": |
| 580 | + jieba_analyse = importlib.import_module("jieba.analyse") |
522 | 581 |
|
523 | | - weighted_terms = jieba.analyse.extract_tags(query, topK=KEYWORD_EXTRACT_TOP_K) |
| 582 | + weighted_terms = jieba_analyse.extract_tags( |
| 583 | + query, |
| 584 | + topK=keyword_top_k, |
| 585 | + allowPOS=KEYWORD_ALLOW_POS, |
| 586 | + ) |
524 | 587 | else: |
525 | | - weighted_terms = [] |
526 | | - if self.tokenizer: |
527 | | - weighted_terms = self.tokenizer.tokenize_mixed(query) |
528 | | - else: |
529 | | - weighted_terms = re.findall(r"\b[a-zA-Z0-9]+\b", query.lower()) |
| 588 | + tokenizer = self.tokenizer or FastTokenizer() |
| 589 | + weighted_terms = self._rank_english_keyword_terms(tokenizer.tokenize_english(query)) |
530 | 590 |
|
531 | 591 | query_words: list[str] = [] |
532 | 592 | seen_words: set[str] = set() |
533 | 593 | for term in weighted_terms: |
534 | | - normalized_term = str(term).strip() |
535 | | - if not normalized_term or normalized_term in seen_words: |
| 594 | + normalized_term = self._normalize_keyword_term(term) |
| 595 | + dedupe_key = normalized_term.lower() |
| 596 | + if self._is_keyword_stopword(normalized_term) or dedupe_key in seen_words: |
536 | 597 | continue |
537 | | - seen_words.add(normalized_term) |
| 598 | + seen_words.add(dedupe_key) |
538 | 599 | query_words.append(normalized_term) |
539 | | - if len(query_words) >= KEYWORD_EXTRACT_TOP_K: |
| 600 | + if len(query_words) >= keyword_top_k: |
540 | 601 | break |
541 | 602 | return query_words |
542 | 603 |
|
|
0 commit comments