Repository navigation
Expand file tree
/
Copy pathutil.py
More file actions
64 lines (51 loc) · 2.32 KB
/
Copy pathutil.py
File metadata and controls
64 lines (51 loc) · 2.32 KB
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
import re
from typing import Any, Optional
from rich.console import Console
from rich.table import Table
from transformers import AutoTokenizer, LogitsProcessorList
from evaluate import build_trie_from_index, create_prefix_allowed_tokens_fn
from logit_processor import ConstrainedLogitsProcessor
def _format_config_value(val: Any, max_str_len: int = 120, max_seq_len: int = 5) -> str:
if isinstance(val, (list, dict)) and len(val) > max_seq_len:
return f"<{type(val).__name__} len={len(val)}>"
if isinstance(val, str) and len(val) > max_str_len:
return repr(val[:max_str_len] + "...")
return repr(val)
def print_config_table(title: str, attrs: dict[str, Any]) -> None:
if not attrs:
return
console = Console()
table = Table(title=title, show_lines=False)
table.add_column("config", style="cyan", no_wrap=True)
table.add_column("value", style="green")
for key in sorted(attrs.keys()):
table.add_row(key, _format_config_value(attrs[key]))
console.print(table)
def build_constrained_logits_processor(
index_path: str,
tokenizer: AutoTokenizer,
prefix: Optional[str] = None,
num_beams: int = 50,
) -> LogitsProcessorList:
"""Build Trie from index and return a LogitsProcessorList with ConstrainedLogitsProcessor."""
print(f"Building Trie from {index_path}...")
trie, prompt_suffix_ids, prefix_index = build_trie_from_index(index_path, tokenizer, prefix=prefix)
print(f"Trie built: prefix_index={prefix_index}, num_items={len(trie)}")
prefix_allowed_tokens_fn = create_prefix_allowed_tokens_fn(trie, prompt_suffix_ids)
logits_processor = LogitsProcessorList(
[
ConstrainedLogitsProcessor(
prefix_allowed_tokens_fn=prefix_allowed_tokens_fn,
num_beams=num_beams,
prefix_index=prefix_index,
prefix_ids=prompt_suffix_ids,
eos_token_id=tokenizer.eos_token_id,
)
]
)
return logits_processor
def _extract_number(text: str) -> list[int]:
return [int(m) for m in re.findall(r"_(\d+)", text)]
def _extract_all_tuples(text: str) -> list[tuple[int, int, int]]:
"""Extract all (a, b, c) tuples from text in <a_N><b_N><c_N> format."""
return [(int(a), int(b), int(c)) for a, b, c in re.findall(r"<a_(\d+)><b_(\d+)><c_(\d+)>", text)]