217 lines
8.3 KiB
Python
217 lines
8.3 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""Rule-based hotword correction for ASR post-processing."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import re
|
||
from difflib import SequenceMatcher
|
||
from typing import Any
|
||
|
||
try:
|
||
from pypinyin import lazy_pinyin
|
||
except Exception: # pragma: no cover - optional runtime enhancement.
|
||
lazy_pinyin = None # type: ignore[assignment]
|
||
|
||
|
||
_HOTWORD_PROMPT_LEAK_PATTERNS = (
|
||
re.compile(r"^\s*Use this context when resolving named entities:\s*", re.IGNORECASE),
|
||
re.compile(r"^\s*(?:上下文信息[::]?\s*)?热词列表[::]\s*[[\[][^\]]]{0,500}[]\]]\s*[,。,::;;\s]*"),
|
||
)
|
||
|
||
|
||
def parse_hotwords(hotwords: str | list[dict[str, object]] | None) -> list[dict[str, Any]]:
|
||
if isinstance(hotwords, list):
|
||
return _extract_hotword_entries(hotwords)
|
||
|
||
raw_text = str(hotwords or "").strip()
|
||
if not raw_text:
|
||
return []
|
||
|
||
tokens = [part for part in re.split(r"[\s,,;;]+", raw_text) if part]
|
||
entries: list[dict[str, Any]] = []
|
||
index = 0
|
||
order = 0
|
||
while index < len(tokens):
|
||
word = tokens[index].strip()
|
||
if not word:
|
||
index += 1
|
||
continue
|
||
weight = 1.0
|
||
if index + 1 < len(tokens):
|
||
try:
|
||
weight = float(tokens[index + 1])
|
||
index += 2
|
||
except ValueError:
|
||
index += 1
|
||
else:
|
||
index += 1
|
||
entries.append({"text": word, "weight": weight, "order": order})
|
||
order += 1
|
||
return entries
|
||
|
||
|
||
def format_hotword_prompt_context(hotwords: str | list[dict[str, object]] | None) -> str:
|
||
entries = parse_hotwords(hotwords)
|
||
if not entries:
|
||
return ""
|
||
|
||
unique_words: list[str] = []
|
||
for entry in entries:
|
||
word = str(entry["text"]).strip()
|
||
if word and word not in unique_words:
|
||
unique_words.append(word)
|
||
if not unique_words:
|
||
return ""
|
||
return f"热词列表:[{', '.join(unique_words)}]"
|
||
|
||
|
||
def strip_hotword_prompt_leakage(text: str) -> str:
|
||
cleaned = str(text or "")
|
||
changed = True
|
||
while changed and cleaned:
|
||
changed = False
|
||
for pattern in _HOTWORD_PROMPT_LEAK_PATTERNS:
|
||
updated, count = pattern.subn("", cleaned, count=1)
|
||
if count:
|
||
cleaned = updated
|
||
changed = True
|
||
return cleaned.strip()
|
||
|
||
|
||
def _extract_hotword_entries(hotwords: list[dict[str, object]] | None) -> list[dict[str, Any]]:
|
||
entries: list[dict[str, Any]] = []
|
||
for index, item in enumerate(hotwords or []):
|
||
hotword_text = str(item.get("hotword", "")).strip()
|
||
if not hotword_text:
|
||
continue
|
||
try:
|
||
weight = float(item.get("weight", 1.0))
|
||
except Exception:
|
||
weight = 1.0
|
||
entries.append({"text": hotword_text, "weight": weight, "order": index})
|
||
return entries
|
||
|
||
|
||
def _normalize_ascii_token(text: str) -> str:
|
||
return re.sub(r"[^a-z0-9]+", "", str(text or "").lower())
|
||
|
||
|
||
def _common_prefix_length(left: str, right: str) -> int:
|
||
matched = 0
|
||
for left_char, right_char in zip(left, right):
|
||
if left_char != right_char:
|
||
break
|
||
matched += 1
|
||
return matched
|
||
|
||
|
||
def _common_suffix_length(left: str, right: str) -> int:
|
||
return _common_prefix_length(left[::-1], right[::-1])
|
||
|
||
|
||
def _normalize_cjk_pinyin(text: str) -> tuple[str, ...]:
|
||
if lazy_pinyin is None:
|
||
return ()
|
||
return tuple(part.strip().lower() for part in lazy_pinyin(str(text or ""), errors="ignore") if str(part).strip())
|
||
|
||
|
||
def _replace_ascii_hotwords(text: str, hotword_entries: list[dict[str, Any]]) -> tuple[str, bool, list[str]]:
|
||
updated_text = str(text or "")
|
||
changed = False
|
||
matched_hotwords: list[str] = []
|
||
ascii_pattern = re.compile(r"[A-Za-z][A-Za-z0-9\s._-]{0,40}")
|
||
for entry in hotword_entries:
|
||
hotword = str(entry["text"])
|
||
if not re.search(r"[A-Za-z]", hotword):
|
||
continue
|
||
normalized_hotword = _normalize_ascii_token(hotword)
|
||
if not normalized_hotword:
|
||
continue
|
||
|
||
def _replace_match(match: re.Match[str]) -> str:
|
||
nonlocal changed
|
||
candidate = match.group(0)
|
||
if _normalize_ascii_token(candidate) == normalized_hotword and candidate != hotword:
|
||
changed = True
|
||
if hotword not in matched_hotwords:
|
||
matched_hotwords.append(hotword)
|
||
return hotword
|
||
return candidate
|
||
|
||
updated_text = ascii_pattern.sub(_replace_match, updated_text)
|
||
return updated_text, changed, matched_hotwords
|
||
|
||
|
||
def _score_cjk_candidate(candidate: str, hotword_entry: dict[str, Any]) -> tuple[int, float, int, float, int] | None:
|
||
hotword = str(hotword_entry["text"])
|
||
if candidate == hotword or len(candidate) != len(hotword):
|
||
return None
|
||
candidate_pinyin = _normalize_cjk_pinyin(candidate)
|
||
hotword_pinyin = _normalize_cjk_pinyin(hotword)
|
||
if candidate_pinyin and candidate_pinyin == hotword_pinyin:
|
||
return (3, float(hotword_entry["weight"]), len(hotword), 1.0, -int(hotword_entry["order"]))
|
||
if len(hotword) <= 2:
|
||
return None
|
||
prefix_length = _common_prefix_length(candidate, hotword)
|
||
suffix_length = _common_suffix_length(candidate, hotword)
|
||
similarity = SequenceMatcher(None, candidate, hotword).ratio()
|
||
if prefix_length >= len(hotword) - 1 or suffix_length >= len(hotword) - 1:
|
||
return (2, float(hotword_entry["weight"]), prefix_length + suffix_length, similarity, -int(hotword_entry["order"]))
|
||
if similarity >= 0.67 and prefix_length >= 1 and suffix_length >= 1:
|
||
return (2, float(hotword_entry["weight"]), prefix_length + suffix_length, similarity, -int(hotword_entry["order"]))
|
||
return None
|
||
|
||
|
||
def _replace_cjk_hotwords(text: str, hotword_entries: list[dict[str, Any]]) -> tuple[str, bool, list[str]]:
|
||
updated_text = str(text or "")
|
||
changed = False
|
||
matched_hotwords: list[str] = []
|
||
chinese_entries = [entry for entry in hotword_entries if re.fullmatch(r"[\u4e00-\u9fff]{2,12}", str(entry["text"]))]
|
||
if not chinese_entries:
|
||
return updated_text, False, matched_hotwords
|
||
candidate_lengths = sorted({len(str(entry["text"])) for entry in chinese_entries})
|
||
token_pattern = re.compile(r"[\u4e00-\u9fff]{2,24}")
|
||
|
||
def _replace_match(match: re.Match[str]) -> str:
|
||
nonlocal changed
|
||
token = match.group(0)
|
||
best_choice: tuple[tuple[int, float, int, float, int], int, int, str] | None = None
|
||
for hotword_length in candidate_lengths:
|
||
if hotword_length > len(token):
|
||
continue
|
||
scoped_entries = [entry for entry in chinese_entries if len(str(entry["text"])) == hotword_length]
|
||
for index in range(0, len(token) - hotword_length + 1):
|
||
candidate = token[index:index + hotword_length]
|
||
for entry in scoped_entries:
|
||
score = _score_cjk_candidate(candidate, entry)
|
||
if score is None:
|
||
continue
|
||
current_choice = (score, index, hotword_length, str(entry["text"]))
|
||
if best_choice is None or current_choice > best_choice:
|
||
best_choice = current_choice
|
||
if best_choice is None:
|
||
return token
|
||
_, begin_index, hotword_length, hotword = best_choice
|
||
token_chars = list(token)
|
||
token_chars[begin_index:begin_index + hotword_length] = list(hotword)
|
||
changed = True
|
||
if hotword not in matched_hotwords:
|
||
matched_hotwords.append(hotword)
|
||
return "".join(token_chars)
|
||
|
||
updated_text = token_pattern.sub(_replace_match, updated_text)
|
||
return updated_text, changed, matched_hotwords
|
||
|
||
|
||
def apply_hotword_rules(text: str, hotwords: str | list[dict[str, object]] | None = None) -> tuple[str, bool, list[str]]:
|
||
hotword_entries = parse_hotwords(hotwords)
|
||
if not hotword_entries:
|
||
return str(text or ""), False, []
|
||
updated_text, ascii_changed, ascii_matches = _replace_ascii_hotwords(str(text or ""), hotword_entries)
|
||
updated_text, cjk_changed, cjk_matches = _replace_cjk_hotwords(updated_text, hotword_entries)
|
||
matched_hotwords: list[str] = []
|
||
for hotword_text in ascii_matches + cjk_matches:
|
||
if hotword_text not in matched_hotwords:
|
||
matched_hotwords.append(hotword_text)
|
||
return updated_text, ascii_changed or cjk_changed, matched_hotwords
|