Repository navigation
Expand file tree
/
Copy pathchat_cli.py
More file actions
156 lines (130 loc) · 5.25 KB
/
Copy pathchat_cli.py
File metadata and controls
156 lines (130 loc) · 5.25 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
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
"""CLIチャット.
python src/chat_cli.py
python src/chat_cli.py --ckpt checkpoints/final --temperature 0.9
チャット中に使えるコマンド:
/reset 会話履歴を消す
/temp 0.9 ランダムさを変える (0に近いほど堅い)
/topk 40 候補を上位k個に絞る (0で無効)
/penalty 1.2 繰り返しへのペナルティ
/tokens 200 1回に生成する最大文字数
/config 現在の設定を表示
/exit 終了
"""
from __future__ import annotations
import argparse
import os
import sys
import time
from pathlib import Path
import torch
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from src.generate import DEFAULTS, chat_stream, load_bundle # noqa: E402
def enable_ansi() -> None:
"""古い Windows コンソールでは色指定がそのまま文字として出てしまう.
ENABLE_VIRTUAL_TERMINAL_PROCESSING を立てておく。
Windows Terminal では最初から有効だが、cmd.exe 直起動では効いていない。
"""
if os.name != "nt":
return
try:
import ctypes
kernel32 = ctypes.windll.kernel32
handle = kernel32.GetStdHandle(-11) # STD_OUTPUT_HANDLE
mode = ctypes.c_uint32()
if kernel32.GetConsoleMode(handle, ctypes.byref(mode)):
kernel32.SetConsoleMode(handle, mode.value | 0x0004)
except Exception:
pass
enable_ansi()
RESET = "\033[0m"
DIM = "\033[2m"
CYAN = "\033[36m"
GREEN = "\033[32m"
YELLOW = "\033[33m"
def parse_command(line: str, params: dict) -> bool | None:
"""コマンドなら処理して True/False を返す. 通常の発言なら None."""
if not line.startswith("/"):
return None
parts = line.split()
name = parts[0][1:]
arg = parts[1] if len(parts) > 1 else None
keys = {"temp": "temperature", "topk": "top_k", "penalty": "repetition_penalty",
"tokens": "max_new_tokens"}
if name in ("exit", "quit", "q"):
return False
if name == "reset":
print(f"{DIM}会話履歴を消しました{RESET}")
return True
if name == "config":
print(f"{DIM}" + " ".join(f"{k}={v}" for k, v in params.items()) + RESET)
return True
if name == "help":
print(__doc__)
return True
if name in keys and arg is not None:
key = keys[name]
params[key] = int(arg) if isinstance(DEFAULTS[key], int) else float(arg)
print(f"{DIM}{key} = {params[key]}{RESET}")
return True
print(f"{YELLOW}不明なコマンドです。/help を見てください。{RESET}")
return True
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--ckpt", default="checkpoints/final")
ap.add_argument("--temperature", type=float, default=DEFAULTS["temperature"])
ap.add_argument("--top-k", type=int, default=DEFAULTS["top_k"])
ap.add_argument("--repetition-penalty", type=float, default=DEFAULTS["repetition_penalty"])
ap.add_argument("--max-new-tokens", type=int, default=DEFAULTS["max_new_tokens"])
ap.add_argument("--history", type=int, default=2, help="何ターン前まで文脈に含めるか")
ap.add_argument("--seed", type=int, default=None)
args = ap.parse_args()
if args.seed is not None:
torch.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
model, tokenizer = load_bundle(args.ckpt)
params = {
"temperature": args.temperature,
"top_k": args.top_k,
"repetition_penalty": args.repetition_penalty,
"max_new_tokens": args.max_new_tokens,
}
print("=" * 60)
print(f" 1LM chat {model.n_params/1e6:.2f}M params / "
f"vocab {tokenizer.vocab_size} / context {model.cfg.block_size}")
print(f"{DIM} /help でコマンド一覧, /exit で終了{RESET}")
print("=" * 60)
history: list[tuple[str, str]] = []
while True:
try:
# PowerShell からパイプで流し込むと先頭行に BOM (U+FEFF) が付く。
# 残すと未知の文字として学習語彙から外れ、1文字目の予測が崩れる。
line = input(f"\n{CYAN}あなた>{RESET} ").lstrip("\ufeff").strip()
except (EOFError, KeyboardInterrupt):
print()
break
if not line:
continue
if not sys.stdin.isatty():
# パイプで流し込んだときは入力が画面に出ないため、自分で echo する
print(line)
result = parse_command(line, params)
if result is False:
break
if result is True:
if line.startswith("/reset"):
history.clear()
continue
print(f"{GREEN}1LM >{RESET} ", end="", flush=True)
start = time.time()
pieces = []
for piece in chat_stream(model, tokenizer, history[-args.history:], line, **params):
print(piece, end="", flush=True)
pieces.append(piece)
reply = "".join(pieces)
took = time.time() - start
speed = len(reply) / took if took > 0 else 0
print(f"\n{DIM}({len(reply)} 文字 / {took:.1f}秒 / {speed:.0f} 文字毎秒){RESET}")
history.append((line, reply))
print("またね。")
if __name__ == "__main__":
main()