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