text-rewrite 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,3 @@
1
+ from .pipeline import Pipeline
2
+
3
+ __all__ = ["Pipeline"]
@@ -0,0 +1,9 @@
1
+ from .base import BaseFilter
2
+ from .hotword import HotwordFilter
3
+ from .regex import RegexFilter
4
+ from .ner import NERFilter
5
+ from .jieba_ner import JiebaNERFilter
6
+ from .fuzzy_phoneme import FuzzyPhonemeFilter
7
+ from .entity_fuzzy import EntityAwareFuzzyFilter
8
+
9
+ __all__ = ["BaseFilter", "HotwordFilter", "RegexFilter", "NERFilter", "JiebaNERFilter", "FuzzyPhonemeFilter", "EntityAwareFuzzyFilter"]
@@ -0,0 +1,15 @@
1
+ from abc import ABC, abstractmethod
2
+
3
+ class BaseFilter(ABC):
4
+ """
5
+ Abstract base class for all text filters.
6
+ """
7
+ def __init__(self, name: str = None):
8
+ self.name = name or self.__class__.__name__
9
+
10
+ @abstractmethod
11
+ def process(self, text: str) -> str:
12
+ """
13
+ Process the input text and return the filtered text.
14
+ """
15
+ pass
@@ -0,0 +1,187 @@
1
+ import logging
2
+ import re
3
+ from typing import List
4
+ from collections import defaultdict
5
+ from .base import BaseFilter
6
+ from .fuzzy_phoneme import FuzzyPhonemeFilter
7
+ logger = logging.getLogger(__name__)
8
+
9
+ class EntityAwareFuzzyFilter(BaseFilter):
10
+ """
11
+ SOTA Entity-Aware Fuzzy Phoneme Filter.
12
+ Combines high-speed global phoneme matching with Jieba-based NER tagging
13
+ to correct ASR homophone errors accurately and efficiently.
14
+ """
15
+ def __init__(self, name: str = "EntityAwareFuzzyFilter", rules: List[str] = None):
16
+ """
17
+ :param rules: A list of rule strings.
18
+ Format: `[tag]word1:threshold1|word2:replacement2:threshold2`
19
+ tag, replacement, and threshold are optional.
20
+ Example: "[nr]叶开:0.8|张三:0.7", "李四:0.7", "欧阳锋"
21
+ """
22
+ super().__init__(name=name)
23
+ self.global_hotwords = {}
24
+ self.global_thresholds = {}
25
+ self.tagged_hotwords = defaultdict(dict)
26
+ self.tagged_thresholds = defaultdict(dict)
27
+
28
+ self.global_filter = None
29
+ self.tagged_filters = {}
30
+ self.ner_filter = None
31
+
32
+ if rules:
33
+ self.add_rules(rules)
34
+
35
+ def add_rules(self, rules: List[str]):
36
+ tag_pattern = re.compile(r'^\[([^\]]+)\]')
37
+ for rule in rules:
38
+ self._parse_rule(rule, tag_pattern)
39
+
40
+ self._rebuild_filters()
41
+ self._setup_ner_filter()
42
+
43
+ def _parse_rule(self, rule: str, tag_pattern: re.Pattern):
44
+ rule = rule.strip()
45
+ if not rule:
46
+ return
47
+
48
+ tag_match = tag_pattern.match(rule)
49
+ tag = tag_match.group(1) if tag_match else None
50
+
51
+ if tag_match:
52
+ rule = rule[tag_match.end():]
53
+
54
+ parts = [p.strip() for p in rule.split('|')]
55
+ for part in parts:
56
+ if not part:
57
+ continue
58
+
59
+ subparts = [sp.strip() for sp in part.split(':')]
60
+ word = subparts[0]
61
+ replacement = word
62
+ threshold = 0.6
63
+
64
+ if len(subparts) == 2:
65
+ try:
66
+ threshold = float(subparts[1])
67
+ except ValueError:
68
+ replacement = subparts[1]
69
+ elif len(subparts) >= 3:
70
+ replacement = subparts[1]
71
+ try:
72
+ threshold = float(subparts[2])
73
+ except ValueError:
74
+ pass
75
+
76
+ if tag:
77
+ self.tagged_hotwords[tag][word] = replacement
78
+ self.tagged_thresholds[tag][word] = threshold
79
+ else:
80
+ self.global_hotwords[word] = replacement
81
+ self.global_thresholds[word] = threshold
82
+
83
+ def _rebuild_filters(self):
84
+ # Rebuild global filter
85
+ if self.global_hotwords:
86
+ self.global_filter = FuzzyPhonemeFilter(hotwords=self.global_hotwords, threshold=0.6)
87
+ self.global_filter.custom_thresholds = self.global_thresholds
88
+
89
+ # Rebuild tagged filters
90
+ for tag, hw_dict in self.tagged_hotwords.items():
91
+ f = FuzzyPhonemeFilter(hotwords=hw_dict, threshold=0.6)
92
+ f.custom_thresholds = self.tagged_thresholds[tag]
93
+ self.tagged_filters[tag] = f
94
+
95
+ def _setup_ner_filter(self):
96
+ # Init JiebaNERFilter if tagged rules exist
97
+ if not self.tagged_filters:
98
+ return
99
+
100
+ import functools
101
+
102
+ # Cache the processing of individual words to avoid redundant phoneme DP searches
103
+ # Since tagged filters don't change state during processing, caching per tag is safe.
104
+ word_cache = {}
105
+ for tag in self.tagged_filters:
106
+ word_cache[tag] = functools.lru_cache(maxsize=4096)(self.tagged_filters[tag].process)
107
+
108
+ def jieba_callback(terms, text):
109
+ new_text = text
110
+
111
+ # Keep track of character offsets manually since jieba
112
+ # splits continuously.
113
+ offset = 0
114
+ entities = []
115
+ for term in terms:
116
+ word = term.word
117
+ tag_str = term.flag
118
+
119
+ if tag_str in self.tagged_filters:
120
+ entities.append({
121
+ 'word': word,
122
+ 'tag': tag_str,
123
+ 'start': offset,
124
+ 'end': offset + len(word)
125
+ })
126
+ offset += len(word)
127
+
128
+ # Process from back to front to avoid index shifting after string replacement
129
+ for ent in reversed(entities):
130
+ word = ent['word']
131
+ tag = ent['tag']
132
+ start = ent['start']
133
+ end = ent['end']
134
+
135
+ # Fuzzy phoneme search on the isolated entity word (Cached)
136
+ corrected = word_cache[tag](word)
137
+ if corrected != word:
138
+ new_text = new_text[:start] + corrected + new_text[end:]
139
+
140
+ return new_text
141
+
142
+ from .jieba_ner import JiebaNERFilter
143
+ self.ner_filter = JiebaNERFilter(replace_callback=jieba_callback)
144
+ self.ner_filter._load_model() # Pre-load it
145
+
146
+ # Build prechecker
147
+ all_tagged_hw = {}
148
+ all_tagged_thresholds = {}
149
+ for tag, hw_dict in self.tagged_hotwords.items():
150
+ all_tagged_hw.update(hw_dict)
151
+ all_tagged_thresholds.update(self.tagged_thresholds[tag])
152
+
153
+ if all_tagged_hw:
154
+ self.ner_prechecker = FuzzyPhonemeFilter(hotwords=all_tagged_hw, threshold=0.6)
155
+ self.ner_prechecker.custom_thresholds = all_tagged_thresholds
156
+ else:
157
+ self.ner_prechecker = None
158
+
159
+ def process(self, text: str) -> str:
160
+ """
161
+ Execute SOTA cascading replacement.
162
+ 1. Global Numba match (fastest, for long words or safe thresholds).
163
+ 2. Jieba NER constraint match (for tags like 'nr').
164
+ """
165
+ if not text:
166
+ return text
167
+
168
+
169
+ # Step 1: Global substitution
170
+ if self.global_filter:
171
+ text = self.global_filter.process(text)
172
+
173
+ # Step 2: Tag-constrained substitution
174
+ if self.ner_filter:
175
+ # 预检:如果不可能有带标签的热词出现,直接跳过 jieba 分词
176
+ skip_ner = False
177
+ if hasattr(self, 'ner_prechecker') and self.ner_prechecker:
178
+ from text_rewrite.utils.algo_phoneme import get_phoneme_info
179
+ inp = get_phoneme_info(text)
180
+ cands = self.ner_prechecker.rag.index.get_candidates(inp)
181
+ if not cands:
182
+ skip_ner = True
183
+
184
+ if not skip_ner:
185
+ text = self.ner_filter.process(text)
186
+
187
+ return text
@@ -0,0 +1,130 @@
1
+ import logging
2
+ from typing import Dict
3
+ from .base import BaseFilter
4
+ from ..utils.algo_phoneme import get_phoneme_info, get_phoneme_seq
5
+ from ..utils.rag_fast import FastRAG
6
+
7
+ logger = logging.getLogger(__name__)
8
+
9
+ class FuzzyPhonemeFilter(BaseFilter):
10
+ """
11
+ Fuzzy Phoneme Filter using Numba JIT DP acceleration.
12
+ Handles exact and fuzzy phoneme matches to correct ASR homophone errors.
13
+ """
14
+ def __init__(self, name: str = "FuzzyPhonemeFilter", hotwords: Dict[str, str] = None, threshold: float = 0.6):
15
+ """
16
+ :param threshold: The similarity threshold (0.0 to 1.0) for a phoneme sequence to match.
17
+ 0.6 is a good default for allowing some ASR errors.
18
+ """
19
+ super().__init__(name=name)
20
+ self.rag = FastRAG(threshold=threshold)
21
+ self.hotword_replacements = {}
22
+ # Cache hw_info to avoid recomputing
23
+ self._hw_info_cache = {}
24
+
25
+ # Optional mapping of original word -> specific threshold
26
+ self.custom_thresholds = {}
27
+
28
+ if hotwords:
29
+ self.add_hotwords(hotwords)
30
+
31
+ def add_hotwords(self, hotwords: Dict[str, str], custom_thresholds: Dict[str, float] = None):
32
+ """
33
+ Add hotwords dynamically.
34
+ :param hotwords: Dictionary mapping the phonetic target (e.g., "张三") to the replacement string (e.g., "张三").
35
+ :param custom_thresholds: Optional dictionary mapping target word to its specific threshold.
36
+ """
37
+ if custom_thresholds:
38
+ self.custom_thresholds.update(custom_thresholds)
39
+
40
+ hw_dict = {}
41
+ for original, replacement in hotwords.items():
42
+ phonemes = get_phoneme_seq(original)
43
+ hw_dict[original] = phonemes
44
+ self.hotword_replacements[original] = replacement
45
+ self._hw_info_cache[original] = [p.info for p in phonemes]
46
+
47
+ self.rag.add_hotwords(hw_dict)
48
+ self.rag.numba_searcher.build_cache(self._hw_info_cache)
49
+
50
+ def process(self, text: str) -> str:
51
+ """
52
+ Process text by converting to phonemes, searching for fuzzy matches,
53
+ mapping indices back, and replacing.
54
+ """
55
+ if not text:
56
+ return text
57
+
58
+ # 1. Get input phonemes with char_start/char_end
59
+ input_phonemes = get_phoneme_info(text)
60
+ if not input_phonemes:
61
+ return text
62
+
63
+ # 2. Get candidates from inverted index to avoid full scan
64
+ candidates = self.rag.index.get_candidates(input_phonemes)
65
+ if not candidates:
66
+ return text
67
+
68
+ input_info = [p.info for p in input_phonemes]
69
+
70
+ # Find the absolute minimum threshold to query Numba, post-filter later
71
+ min_thresh = self.rag.threshold
72
+ if self.custom_thresholds:
73
+ min_thresh = min(min_thresh, min(self.custom_thresholds.values()))
74
+
75
+ # 3. Perform fine-grained Numba Substring DP search
76
+ replacements = []
77
+
78
+ # Pre-encode input_info ONCE to avoid O(N_candidates) encoding overhead
79
+ inp_codes, inp_langs, _, inp_is_ws, inp_is_we = self.rag.numba_searcher.encode_input_vecs(input_info)
80
+
81
+ candidate_keys = [hw_key for hw_key, _ in candidates]
82
+
83
+ batch_res = self.rag.numba_searcher.search_batch_with_encoded_input(
84
+ candidate_keys, inp_codes, inp_langs, inp_is_ws, inp_is_we, threshold=min_thresh
85
+ )
86
+
87
+ for hw_key, res_list in batch_res.items():
88
+ for score, start_idx, end_idx in res_list:
89
+ # Apply word-specific threshold
90
+ if score >= self.custom_thresholds.get(hw_key, self.rag.threshold):
91
+ replacements.append({
92
+ 'score': score,
93
+ 'start_idx': start_idx,
94
+ 'end_idx': end_idx,
95
+ 'replacement': self.hotword_replacements[hw_key]
96
+ })
97
+
98
+ if not replacements:
99
+ return text
100
+
101
+ # 4. Resolve overlaps (greedy: highest score, then longest span)
102
+ replacements.sort(key=lambda x: (-x['score'], -(x['end_idx'] - x['start_idx'])))
103
+
104
+ final_reps = []
105
+ used_indices = set()
106
+ for r in replacements:
107
+ r_set = set(range(r['start_idx'], r['end_idx']))
108
+ if not r_set.intersection(used_indices):
109
+ final_reps.append(r)
110
+ used_indices.update(r_set)
111
+
112
+ # 5. Execute replacements from back to front to avoid index shifting issues
113
+ final_reps.sort(key=lambda x: x['start_idx'], reverse=True)
114
+
115
+ result_text = text
116
+ for r in final_reps:
117
+ # Map phoneme index to char index
118
+ # end_idx is exclusive in phoneme array
119
+ start_p = r['start_idx']
120
+ end_p = r['end_idx'] - 1
121
+
122
+ if start_p >= len(input_phonemes) or end_p >= len(input_phonemes) or start_p > end_p:
123
+ continue
124
+
125
+ char_start = input_phonemes[start_p].char_start
126
+ char_end = input_phonemes[end_p].char_end
127
+
128
+ result_text = result_text[:char_start] + r['replacement'] + result_text[char_end:]
129
+
130
+ return result_text
@@ -0,0 +1,50 @@
1
+ from typing import Dict, List
2
+ from .base import BaseFilter
3
+
4
+ try:
5
+ from flashtext import KeywordProcessor
6
+ except ImportError:
7
+ KeywordProcessor = None
8
+
9
+
10
+ class HotwordFilter(BaseFilter):
11
+ """
12
+ Hotword Filter for high-performance exact word replacements.
13
+ Uses Aho-Corasick algorithm via flashtext library.
14
+ Ideal for massive dictionaries (100k+ words).
15
+ """
16
+ def __init__(self, name: str = "HotwordFilter", hotwords: Dict[str, str] = None, case_sensitive: bool = False):
17
+ super().__init__(name=name)
18
+ if KeywordProcessor is None:
19
+ raise ImportError("flashtext library is required for HotwordFilter. Run `pip install flashtext`.")
20
+
21
+ self.keyword_processor = KeywordProcessor(case_sensitive=case_sensitive)
22
+ if hotwords:
23
+ self.add_hotwords(hotwords)
24
+
25
+ def add_hotwords(self, hotwords: Dict[str, str]):
26
+ """
27
+ Add hotword replacements dynamically.
28
+ :param hotwords: Dictionary mapping original words to target replacement words.
29
+ Note: flashtext format is { 'replacement': ['word1', 'word2'] } for add_keywords_from_dict.
30
+ Since we get { 'word': 'replacement' }, we iterate or transform it.
31
+ """
32
+ for original_word, replacement_word in hotwords.items():
33
+ self.keyword_processor.add_keyword(original_word, replacement_word)
34
+
35
+ def remove_hotwords(self, words: List[str]):
36
+ """
37
+ Remove hotwords from the trie dynamically.
38
+ """
39
+ for word in words:
40
+ self.keyword_processor.remove_keyword(word)
41
+
42
+ def process(self, text: str) -> str:
43
+ """
44
+ Process text by efficiently finding and replacing all matching hotwords.
45
+ """
46
+ if not text:
47
+ return text
48
+
49
+ # flashtext operates in O(N) where N is text length, extremely fast.
50
+ return self.keyword_processor.replace_keywords(text)
@@ -0,0 +1,65 @@
1
+ import logging
2
+ from typing import Callable, Any
3
+ from .base import BaseFilter
4
+
5
+ logger = logging.getLogger(__name__)
6
+
7
+ class JiebaNERFilter(BaseFilter):
8
+ """
9
+ Extremely fast POS/NER Filter using Jieba.
10
+ Uses `jieba.posseg` to perform POS tagging.
11
+ """
12
+ def __init__(self, name: str = "JiebaNERFilter", replace_callback: Callable[[Any, str], str] = None):
13
+ """
14
+ :param replace_callback: A custom function that takes the jieba.posseg output (generator of pairs)
15
+ and returns the modified text string.
16
+ """
17
+ super().__init__(name=name)
18
+ self.replace_callback = replace_callback
19
+ self._posseg = None
20
+ self._loaded = False
21
+
22
+ def _load_model(self):
23
+ """
24
+ Lazy load jieba to avoid overhead if not used.
25
+ """
26
+ if self._loaded:
27
+ return
28
+
29
+ logger.info("Lazy loading Jieba POS tagger...")
30
+ try:
31
+ import jieba.posseg as pseg
32
+ self._posseg = pseg
33
+ # Warm up jieba to load dictionary
34
+ list(pseg.cut("预热"))
35
+ self._loaded = True
36
+ logger.info("Jieba loaded successfully.")
37
+ except ImportError:
38
+ raise ImportError("jieba is required for JiebaNERFilter. Run `pip install jieba`.")
39
+ except Exception as e:
40
+ logger.error(f"Failed to load jieba: {e}")
41
+ raise
42
+
43
+ def process(self, text: str) -> str:
44
+ if not text:
45
+ return text
46
+
47
+ if not self._loaded:
48
+ self._load_model()
49
+
50
+ if not self._posseg:
51
+ return text
52
+
53
+ try:
54
+ # pseg.cut returns a generator of pair(word, flag)
55
+ # flag corresponds to POS tag, e.g., 'nr' for person name
56
+ terms = list(self._posseg.cut(text, HMM=False))
57
+
58
+ if self.replace_callback:
59
+ return self.replace_callback(terms, text)
60
+ else:
61
+ return text
62
+
63
+ except Exception as e:
64
+ logger.error(f"Jieba NER processing failed: {e}")
65
+ return text
@@ -0,0 +1,72 @@
1
+ import logging
2
+ from typing import Callable, Any
3
+ from .base import BaseFilter
4
+
5
+ logger = logging.getLogger(__name__)
6
+
7
+ class NERFilter(BaseFilter):
8
+ """
9
+ NER Filter for replacing or normalizing entities based on Named Entity Recognition.
10
+ Uses HanLP. Supports lazy loading of the model to save memory until first use.
11
+ """
12
+ def __init__(self, name: str = "NERFilter", model_name: str = "LARGE_ALBERT_BASE", replace_callback: Callable[[Any], str] = None):
13
+ """
14
+ :param model_name: HanLP model identifier for NER/POS.
15
+ :param replace_callback: A custom function that takes the HanLP output (e.g., list of words/tags)
16
+ and returns the modified text string.
17
+ """
18
+ super().__init__(name=name)
19
+ self.model_name = model_name
20
+ self.replace_callback = replace_callback
21
+ self._model = None
22
+ self._loaded = False
23
+
24
+ def _load_model(self):
25
+ """
26
+ Lazy load the HanLP model.
27
+ """
28
+ if self._loaded:
29
+ return
30
+
31
+ logger.info(f"Lazy loading HanLP NER model '{self.model_name}'...")
32
+ try:
33
+ import hanlp
34
+ # In a real scenario, you'd load the specific NER/POS model you need.
35
+ # Example: tok/pos/ner pipeline
36
+ self._model = hanlp.load(hanlp.pretrained.mtl.CLOSE_TOK_POS_NER_SRL_DEP_SDP_CON_ELECTRA_SMALL_ZH)
37
+ self._loaded = True
38
+ logger.info("HanLP NER model loaded successfully.")
39
+ except ImportError:
40
+ raise ImportError("hanlp is required for NERFilter. Run `pip install hanlp`.")
41
+ except Exception as e:
42
+ logger.error(f"Failed to load HanLP model: {e}")
43
+ raise
44
+
45
+ def process(self, text: str) -> str:
46
+ """
47
+ Process text by extracting entities and applying the replacement logic.
48
+ """
49
+ if not text:
50
+ return text
51
+
52
+ if not self._loaded:
53
+ self._load_model()
54
+
55
+ if not self._model:
56
+ return text # Fail-safe, return original if model failed to load but no exception was raised
57
+
58
+ try:
59
+ # Note: actual HanLP return format depends on the specific model loaded.
60
+ # Here we assume a MTL (multi-task learning) model that returns a dict.
61
+ doc = self._model(text)
62
+
63
+ if self.replace_callback:
64
+ return self.replace_callback(doc, text)
65
+ else:
66
+ # Default behavior: do nothing if no callback is provided, or implement a basic rule
67
+ # (e.g., wrap entities in brackets). Here we just return the text.
68
+ return text
69
+
70
+ except Exception as e:
71
+ logger.error(f"NER processing failed: {e}")
72
+ return text
@@ -0,0 +1,42 @@
1
+ import re
2
+ from typing import Dict, List, Tuple
3
+ from .base import BaseFilter
4
+
5
+ class RegexFilter(BaseFilter):
6
+ """
7
+ Regex Filter for pattern-based text replacements.
8
+ Patterns are pre-compiled for performance.
9
+ """
10
+ def __init__(self, name: str = "RegexFilter", rules: Dict[str, str] = None):
11
+ super().__init__(name=name)
12
+ self.compiled_rules: List[Tuple[re.Pattern, str]] = []
13
+ if rules:
14
+ self.add_rules(rules)
15
+
16
+ def add_rules(self, rules: Dict[str, str]):
17
+ """
18
+ Add new regex replacement rules.
19
+ :param rules: A dictionary where key is the regex pattern and value is the replacement string.
20
+ """
21
+ for pattern, replacement in rules.items():
22
+ compiled_pattern = re.compile(pattern)
23
+ self.compiled_rules.append((compiled_pattern, replacement))
24
+
25
+ def clear_rules(self):
26
+ """
27
+ Clear all existing regex rules.
28
+ """
29
+ self.compiled_rules.clear()
30
+
31
+ def process(self, text: str) -> str:
32
+ """
33
+ Process text by sequentially applying all regex replacements.
34
+ """
35
+ if not text:
36
+ return text
37
+
38
+ current_text = text
39
+ for compiled_pattern, replacement in self.compiled_rules:
40
+ current_text = compiled_pattern.sub(replacement, current_text)
41
+
42
+ return current_text
@@ -0,0 +1,44 @@
1
+ from typing import List
2
+ import logging
3
+ from .filters.base import BaseFilter
4
+
5
+ logger = logging.getLogger(__name__)
6
+
7
+ class Pipeline:
8
+ """
9
+ TextRewrite Pipeline.
10
+ Manages a sequence of filters and applies them to input text in order.
11
+ """
12
+ def __init__(self):
13
+ self.filters: List[BaseFilter] = []
14
+
15
+ def add_filter(self, text_filter: BaseFilter):
16
+ """
17
+ Register a filter to the pipeline.
18
+ """
19
+ if not isinstance(text_filter, BaseFilter):
20
+ raise TypeError("Filter must be an instance of BaseFilter")
21
+ self.filters.append(text_filter)
22
+ return self # For method chaining
23
+
24
+ def process(self, text: str) -> str:
25
+ """
26
+ Process the input text through all registered filters sequentially.
27
+ """
28
+ if not text:
29
+ return text
30
+
31
+ current_text = text
32
+ for f in self.filters:
33
+ try:
34
+ current_text = f.process(current_text)
35
+ except Exception as e:
36
+ logger.error(f"Error processing text in filter {f.name}: {e}")
37
+ # Depending on strictness, we might want to re-raise or continue.
38
+ # For ASR/OCR robustness, continuing with the unmodified text of this step is often preferred.
39
+ continue
40
+
41
+ return current_text
42
+
43
+ def __call__(self, text: str) -> str:
44
+ return self.process(text)
File without changes