-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsample.py
More file actions
182 lines (154 loc) · 7.57 KB
/
Copy pathsample.py
File metadata and controls
182 lines (154 loc) · 7.57 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
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
"""학습된 kopus 체크포인트에서 텍스트를 생성한다.
python sample.py --out_dir=out-debug
python sample.py --out_dir=out-nano_110m --prompt="오늘 아침" --temperature=0.7
python sample.py --out_dir=out-debug --check_kv # KV캐시 등가 검증
"""
import os
import sys
from contextlib import nullcontext
import torch
from model import MODEL_KEYS, GPT
from tokenizer import Tokenizer
# WHY: train.py의 인자 파서와 사실상 중복이다. 그래도 복사한다 — sample.py를 train.py에
# import로 묶어 두면 "생성만 해 보고 싶은 사람"이 학습 코드 전체를 읽어야 한다.
# 12줄 중복이 결합보다 싸다.
DEFAULTS = dict(
out_dir="",
prompt="안녕하세요",
num_samples=3,
max_new_tokens=200,
temperature=0.8,
top_k=50,
seed=1337,
device="auto",
check_kv=False,
)
# --check_kv에서 프롬프트가 너무 짧을 때 대신 쓰는 문장 (prefill 분기를 태우려면 4토큰 이상 필요).
FALLBACK_PROMPT = "오늘은 날씨가 참 좋습니다."
def parse_args(argv):
cfg = dict(DEFAULTS)
for arg in argv:
if not arg.startswith("--"):
raise SystemExit(f"인자는 --key=value 형태여야 합니다: {arg}")
key, _, val = arg[2:].partition("=")
if key not in DEFAULTS:
raise SystemExit(f"모르는 인자입니다: --{key} (가능: {', '.join(DEFAULTS)})")
default = DEFAULTS[key]
if isinstance(default, bool):
cfg[key] = val.lower() not in ("false", "0") if val else True
elif isinstance(default, int):
cfg[key] = int(val)
elif isinstance(default, float):
cfg[key] = float(val)
else:
cfg[key] = val
return cfg
def load_checkpoint(out_dir, device):
ckpt_path = os.path.join(out_dir, "ckpt.pt")
if not os.path.exists(ckpt_path):
raise SystemExit(f"{ckpt_path}가 없습니다. 먼저 train.py로 학습하세요.")
# weights_only=False: 우리가 직접 저장한 체크포인트라 신뢰한다(config 딕셔너리 포함).
ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
full_cfg = ckpt["config"]
model = GPT({k: full_cfg[k] for k in MODEL_KEYS})
# compile을 쓰지 않는 게 기본이지만, 쓴 체크포인트도 열리게 접두사를 벗겨 둔다.
state = {k.replace("_orig_mod.", "", 1): v for k, v in ckpt["model"].items()}
model.load_state_dict(state)
model.eval().to(device)
print(f"{ckpt_path} 로드 완료 — iter {ckpt.get('iter_num', '?')}, "
f"val loss {ckpt.get('best_val_loss', float('nan')):.4f}, "
f"파라미터 {model.num_params():,}")
return model, full_cfg
def load_tokenizer(vocab_size):
path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "data", "ko", "tokenizer.json")
if os.path.exists(path):
# merges[:vocab_size-256] 슬라이스. vocab_size가 256이면 merges가 빈 리스트가 되어
# 자연스럽게 byte-level이 된다 — 특수 분기가 필요 없다.
return Tokenizer.load(path, vocab_size=vocab_size)
print(f"[경고] {path}가 없어 byte-level 토크나이저로 폴백합니다. "
f"제대로 된 결과를 보려면 python data/ko/prepare.py를 먼저 실행하세요.")
return Tokenizer.bytes()
def amp_ctx(device):
"""CUDA면 fp16 autocast, CPU면 그냥 fp32. 단일 코드 경로를 유지한다."""
if device.startswith("cuda"):
return torch.autocast("cuda", dtype=torch.float16)
return nullcontext()
def encode_prompt(tok, text, min_tokens=1):
ids = tok.encode(text)
if len(ids) < min_tokens:
ids = tok.encode(FALLBACK_PROMPT)
while 0 < len(ids) < min_tokens: # 그래도 짧으면 반복해서 채운다
ids = ids * 2
if not ids:
raise SystemExit("프롬프트가 비어 있습니다.")
return ids
def run_check_kv(model, tok, cfg, device):
"""KV캐시 경로의 logits가 전체 재-forward의 logits와 같은지 단언한다."""
# 프롬프트를 4토큰 이상으로 맞춰 prefill(T>1) 분기와 디코드(T=1) 분기를 모두 태운다.
ids = encode_prompt(tok, cfg["prompt"], min_tokens=4)
prompt_len = len(ids)
seq = torch.tensor([ids], dtype=torch.long, device=device)
steps = 8
cached = []
caches = model.alloc_kv_cache(1, device=device)
pos = 0
with torch.no_grad(), amp_ctx(device):
for _ in range(steps):
chunk = seq if pos == 0 else seq[:, -1:]
logits = model.forward_cached(chunk, caches, pos)
pos += chunk.size(1)
cached.append(logits.float())
seq = torch.cat((seq, torch.argmax(logits, dim=-1, keepdim=True)), dim=1)
# 마지막에 뽑은 토큰은 아직 모델에 먹인 적이 없으므로 뺀다.
full_logits, _ = model(seq[:, :-1])
full_logits = full_logits.float()
# GOTCHA: fp16 tolerance 2e-2는 T4 실측 전이라 아직 검증되지 않은 값이다.
# 커널 폴백/누산 순서가 달라지면 조정이 필요할 수 있다.
fp16 = device.startswith("cuda")
atol = 2e-2 if fp16 else 1e-4
ok = True
for j, logits in enumerate(cached):
ref = full_logits[:, prompt_len - 1 + j]
close = torch.allclose(logits, ref, atol=atol)
print(f" step {j}: allclose={close} maxdiff={(logits - ref).abs().max().item():.3e}")
ok = ok and close
label = "fp16" if fp16 else "fp32"
if ok:
print(f"[PASS] KV캐시 경로가 전체 재-forward와 일치합니다 ({label}, atol={atol}).")
else:
print(f"[FAIL] KV캐시 등가 검증 실패 ({label}, atol={atol}). "
f"RoPE 위치 오프셋이나 인과마스크를 의심하세요.")
sys.exit(1)
def main():
cfg = parse_args(sys.argv[1:])
if not cfg["out_dir"]:
raise SystemExit("--out_dir=<체크포인트 디렉토리>는 필수입니다.")
torch.manual_seed(cfg["seed"])
device = cfg["device"]
if device == "auto":
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"device={device}")
model, full_cfg = load_checkpoint(cfg["out_dir"], device)
tok = load_tokenizer(full_cfg["vocab_size"])
if cfg["check_kv"]:
run_check_kv(model, tok, cfg, device)
return
# GOTCHA: byte-level(vocab 256) 체크포인트는 한글 한 글자가 UTF-8 3바이트라,
# 온도가 높으면 3연속 추첨을 모두 통과해야 글자가 완성된다. 하나만 어긋나면 �다.
# BPE(vocab 32768)는 한 번 뽑으면 음절 이상이 나와 같은 온도에서도 멀쩡하다.
# fertility가 생성 품질로 바로 이어지는 지점이라 기본값을 낮추지 않고 남겨 뒀다.
if full_cfg["vocab_size"] == 256 and cfg["temperature"] > 0.3:
print(f"[힌트] byte-level 체크포인트입니다. --temperature=0 으로 돌리면 "
f"모델이 실제로 외운 문장이 보입니다 (현재 {cfg['temperature']}).")
ids = encode_prompt(tok, cfg["prompt"])
x = torch.tensor([ids], dtype=torch.long, device=device)
top_k = cfg["top_k"] if cfg["top_k"] > 0 else None
for i in range(cfg["num_samples"]):
with torch.no_grad(), amp_ctx(device):
out = model.generate(x, cfg["max_new_tokens"], cfg["temperature"], top_k)
print(f"\n--- 샘플 {i + 1}/{cfg['num_samples']} ---")
# decode는 errors="replace"를 내장한다 — 학습 안 된 모델이 뱉는 깨진 바이트에도
# UnicodeDecodeError로 죽지 않는다.
print(tok.decode(out[0].tolist()))
if __name__ == "__main__":
main()