File size: 5,419 Bytes
bc8e2a9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 | """Backend-agnostic hotword trie logit-boosting core."""
from __future__ import annotations
import re
from typing import Any, Callable, Dict, List, Sequence, Set
_CONTROL_TOKEN_RE = re.compile(r"<\|[^>]+?\|>")
_BARE_TAG_RE = re.compile(r"</?[^>\s]+>")
_CJK_KANA_HANGUL_RE = re.compile("[\u4e00-\u9fff\u3040-\u30ff\uac00-\ud7af]")
class HotwordTrie:
"""Prefix trie over hotword token sequences with per-step boost lookup."""
def __init__(
self,
token_sequences: Sequence[Sequence[int]],
*,
start_boost: float,
continuation_boost: float,
) -> None:
self.start_boost = float(start_boost)
self.continuation_boost = float(continuation_boost)
self.trie: Dict[int, Dict[int, Any]] = {}
self.max_sequence_len = 0
for seq in token_sequences:
ids = [int(token_id) for token_id in seq]
if not ids:
continue
node = self.trie
for token_id in ids:
node = node.setdefault(token_id, {})
self.max_sequence_len = max(self.max_sequence_len, len(ids))
self.start_token_ids = sorted(self.trie.keys())
def __bool__(self) -> bool:
return bool(self.trie)
def boosts_for_generated(self, generated_ids: Sequence[int]) -> Dict[int, float]:
boosts: Dict[int, float] = {}
if self.start_boost:
for token_id in self.start_token_ids:
boosts[token_id] = max(boosts.get(token_id, 0.0), self.start_boost)
if not generated_ids or not self.continuation_boost or self.max_sequence_len <= 1:
return boosts
max_prefix_len = min(len(generated_ids), self.max_sequence_len - 1)
for prefix_len in range(1, max_prefix_len + 1):
node: Dict[int, Any] = self.trie
matched = True
for token_id in generated_ids[-prefix_len:]:
next_node = node.get(int(token_id))
if next_node is None:
matched = False
break
node = next_node
if not matched:
continue
for next_token_id in node.keys():
boosts[int(next_token_id)] = max(
boosts.get(int(next_token_id), 0.0), self.continuation_boost
)
return boosts
def _has_cjk_or_kana_or_hangul(text: str) -> bool:
return bool(_CJK_KANA_HANGUL_RE.search(str(text or "")))
def _hotword_text_variants(word: str) -> List[str]:
word = str(word or "").strip()
if not word:
return []
variants = [word]
if not _has_cjk_or_kana_or_hangul(word) and re.search(r"[A-Za-z0-9_]", word):
variants.append(" " + word)
out: List[str] = []
seen: Set[str] = set()
for value in variants:
if value not in seen:
seen.add(value)
out.append(value)
return out
def token_is_control_or_special(token: str, token_id: int, special_ids: Set[int]) -> bool:
if int(token_id) in special_ids:
return True
token = str(token)
return bool(_CONTROL_TOKEN_RE.fullmatch(token) or _BARE_TAG_RE.fullmatch(token))
def build_hotword_sequences(
hotwords: Sequence[str],
*,
encode: Callable[[str], List[int]],
id_to_token: Callable[[int], str],
special_ids: Set[int],
) -> Dict[str, List[List[int]]]:
special_ids = set(int(x) for x in special_ids if x is not None)
sequences: Dict[str, List[List[int]]] = {}
seen_global: Set[tuple[int, ...]] = set()
for word in hotwords:
word = str(word or "").strip()
if not word:
continue
variants: List[List[int]] = []
for text in _hotword_text_variants(word):
ids = [
int(token_id)
for token_id in encode(text)
if not token_is_control_or_special(id_to_token(int(token_id)), int(token_id), special_ids)
]
key = tuple(ids)
if not key or key in seen_global:
continue
seen_global.add(key)
variants.append(ids)
if variants:
sequences[word] = variants
return sequences
def flatten_sequences(sequences_by_word: Dict[str, List[List[int]]]) -> List[List[int]]:
return [ids for variants in sequences_by_word.values() for ids in variants]
def parse_hotwords(raw: Any) -> List[str]:
values: List[str] = []
if isinstance(raw, (list, tuple)):
values = [str(x).strip() for x in raw]
elif raw:
values = [x.strip() for x in re.split(r"[,,]", str(raw))]
out: List[str] = []
seen: Set[str] = set()
for value in values:
if value and value not in seen:
seen.add(value)
out.append(value)
return out
def build_trie_from_hotwords(
hotwords: Sequence[str],
*,
encode: Callable[[str], List[int]],
id_to_token: Callable[[int], str],
special_ids: Set[int],
start_boost: float,
continuation_boost: float,
) -> tuple[HotwordTrie, Dict[str, List[List[int]]]]:
sequences_by_word = build_hotword_sequences(
hotwords, encode=encode, id_to_token=id_to_token, special_ids=special_ids
)
trie = HotwordTrie(
flatten_sequences(sequences_by_word),
start_boost=start_boost,
continuation_boost=continuation_boost,
)
return trie, sequences_by_word
|