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.
- text_rewrite/__init__.py +3 -0
- text_rewrite/filters/__init__.py +9 -0
- text_rewrite/filters/base.py +15 -0
- text_rewrite/filters/entity_fuzzy.py +187 -0
- text_rewrite/filters/fuzzy_phoneme.py +130 -0
- text_rewrite/filters/hotword.py +50 -0
- text_rewrite/filters/jieba_ner.py +65 -0
- text_rewrite/filters/ner.py +72 -0
- text_rewrite/filters/regex.py +42 -0
- text_rewrite/pipeline.py +44 -0
- text_rewrite/utils/__init__.py +0 -0
- text_rewrite/utils/algo_calc.py +521 -0
- text_rewrite/utils/algo_phoneme.py +330 -0
- text_rewrite/utils/rag_fast.py +805 -0
- text_rewrite-0.1.0.dist-info/METADATA +116 -0
- text_rewrite-0.1.0.dist-info/RECORD +18 -0
- text_rewrite-0.1.0.dist-info/WHEEL +5 -0
- text_rewrite-0.1.0.dist-info/top_level.txt +1 -0
text_rewrite/__init__.py
ADDED
|
@@ -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
|
text_rewrite/pipeline.py
ADDED
|
@@ -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
|