musabc commited on
Commit
e2e6604
·
verified ·
1 Parent(s): 04d9347

upload 06_sample.py

Browse files
Files changed (1) hide show
  1. 06_sample.py +248 -0
06_sample.py ADDED
@@ -0,0 +1,248 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Egitilmis modelden ornek metin uret. V3 / V4 / V5 ile uyumlu.
3
+
4
+ Otomatik tespit:
5
+ - ckpt['version'] = 'v5' -> V5 (200M, 32K vocab, T=2048, theta=100K)
6
+ - ckpt['version'] = 'v4*' -> V4 (50M)
7
+ - ckpt['config'] icinde 'rope_theta' yoksa -> V3
8
+ - tokenizer auto-detect (v5 -> tokenizer-tr-v5, v4 -> tokenizer-tr-16k)
9
+
10
+ Kullanim:
11
+ python 06_sample.py # V5 default (best ckpt)
12
+ python 06_sample.py --version v4 # V4'ten sample
13
+ python 06_sample.py --prompt "İstanbul" --max-tokens 200
14
+ python 06_sample.py --latest # latest_ckpt.pt
15
+ python 06_sample.py --ckpt runs/tr-200m-v5/best_ckpt.pt
16
+ python 06_sample.py --num-samples 5 --temperature 0.7
17
+ python 06_sample.py --chat --prompt "Türkiye nedir?" # SFT modeli için
18
+ """
19
+
20
+ import argparse
21
+ import os
22
+ from pathlib import Path
23
+
24
+ import torch
25
+ from tokenizers import Tokenizer
26
+
27
+ # Liger Kernel'i sample sirasinda kapat — kosul: tek-token forward Liger'in
28
+ # chunked CE'sine ihtiyac duymaz, fused kernel JIT compile overhead'i sample'i yavaslatir
29
+ os.environ.setdefault("NANOGPT_NO_LIGER", "1")
30
+
31
+ # Model import'lari — opsiyonel
32
+ HAS_V3 = HAS_V4 = HAS_V5 = False
33
+ try:
34
+ from model import GPT, GPTConfig
35
+ HAS_V3 = True
36
+ except ImportError:
37
+ GPT = GPTConfig = None
38
+
39
+ try:
40
+ from model_v4 import GPTV4, GPTConfigV4
41
+ HAS_V4 = True
42
+ except ImportError:
43
+ GPTV4 = GPTConfigV4 = None
44
+
45
+ try:
46
+ from model_v5 import GPTV5, GPTConfigV5
47
+ HAS_V5 = True
48
+ except ImportError:
49
+ GPTV5 = GPTConfigV5 = None
50
+
51
+
52
+ DATA_DIR = Path(__file__).parent / "data"
53
+ RUN_DIRS = {
54
+ "v3": Path(__file__).parent / "runs" / "tr-50m-v3",
55
+ "v4": Path(__file__).parent / "runs" / "tr-50m-v4",
56
+ "v5": Path(__file__).parent / "runs" / "tr-200m-v5",
57
+ }
58
+ TOKENIZERS = {
59
+ "v3": "tokenizer-tr-16k.json",
60
+ "v4": "tokenizer-tr-16k.json",
61
+ "v5": "tokenizer-tr-v5.json",
62
+ }
63
+
64
+
65
+ def detect_version(ckpt: dict) -> str:
66
+ """Checkpoint icinden version tespit et."""
67
+ v = ckpt.get("version", "")
68
+ if isinstance(v, str):
69
+ if v.startswith("v5"):
70
+ return "v5"
71
+ if v.startswith("v4"):
72
+ return "v4"
73
+ # Config bazli fallback
74
+ cfg = ckpt.get("config", {})
75
+ if "rope_theta" not in cfg:
76
+ return "v3"
77
+ # V5 vs V4: vocab_size farkli (V5=32000, V4=16000)
78
+ vs = cfg.get("vocab_size", 0)
79
+ if vs >= 24000:
80
+ return "v5"
81
+ return "v4"
82
+
83
+
84
+ def build_model(version: str, ckpt: dict, device: str):
85
+ cfg_dict = ckpt["config"]
86
+ if version == "v5":
87
+ if not HAS_V5:
88
+ raise ImportError("V5 checkpoint ama model_v5.py yok.")
89
+ cfg = GPTConfigV5(**cfg_dict)
90
+ model = GPTV5(cfg).to(device)
91
+ return model, cfg, "V5 (RoPE+RMSNorm+SwiGLU+QK-norm+softcap, 200M)"
92
+ if version == "v4":
93
+ if not HAS_V4:
94
+ raise ImportError("V4 checkpoint ama model_v4.py yok.")
95
+ cfg = GPTConfigV4(**cfg_dict)
96
+ model = GPTV4(cfg).to(device)
97
+ return model, cfg, "V4 (RoPE+RMSNorm+SwiGLU+QK-norm, 50M)"
98
+ if version == "v3":
99
+ if not HAS_V3:
100
+ raise ImportError("V3 checkpoint ama model.py yok.")
101
+ cfg = GPTConfig(**cfg_dict)
102
+ model = GPT(cfg).to(device)
103
+ return model, cfg, "V3 (LayerNorm+GELU+learned PE)"
104
+ raise ValueError(f"Bilinmeyen version: {version}")
105
+
106
+
107
+ def main():
108
+ parser = argparse.ArgumentParser()
109
+ parser.add_argument("--prompt", type=str, default="Türkiye")
110
+ parser.add_argument("--max-tokens", type=int, default=200)
111
+ parser.add_argument("--temperature", type=float, default=0.8)
112
+ parser.add_argument("--top-k", type=int, default=50)
113
+ parser.add_argument("--repetition-penalty", type=float, default=1.15)
114
+ parser.add_argument("--no-repeat-ngram", type=int, default=3)
115
+ parser.add_argument("--num-samples", type=int, default=3)
116
+ parser.add_argument("--ckpt", type=str, default=None,
117
+ help="Tam checkpoint yolu (yoksa --version + best/latest)")
118
+ parser.add_argument("--version", type=str, default="v5",
119
+ choices=["v3", "v4", "v5"],
120
+ help="Hangi run dizini (--ckpt verilmediyse)")
121
+ parser.add_argument("--latest", action="store_true",
122
+ help="best yerine latest checkpoint'i kullan")
123
+ parser.add_argument("--chat", action="store_true",
124
+ help="SFT/Instruct ChatML formatı uygula (otomatik tespit de var)")
125
+ parser.add_argument("--instruction", type=str, default=None,
126
+ help="ChatML için ayrı instruction (input ile birlikte)")
127
+ parser.add_argument("--tokenizer", type=str, default=None,
128
+ help="Tokenizer dosya yolu (yoksa version'a göre seç)")
129
+ parser.add_argument("--seed", type=int, default=None)
130
+ parser.add_argument("--device", type=str, default=None,
131
+ choices=["cuda", "cpu"])
132
+ args = parser.parse_args()
133
+
134
+ if args.seed is not None:
135
+ torch.manual_seed(args.seed)
136
+
137
+ device = args.device or ("cuda" if torch.cuda.is_available() else "cpu")
138
+ print(f"Device: {device}")
139
+
140
+ # Checkpoint yolunu belirle
141
+ if args.ckpt:
142
+ ckpt_path = Path(args.ckpt)
143
+ else:
144
+ run_dir = RUN_DIRS[args.version]
145
+ name = "latest_ckpt.pt" if args.latest else "best_ckpt.pt"
146
+ ckpt_path = run_dir / name
147
+ if not ckpt_path.exists():
148
+ # Fallback: best yoksa latest dene
149
+ alt = run_dir / ("best_ckpt.pt" if args.latest else "latest_ckpt.pt")
150
+ if alt.exists():
151
+ ckpt_path = alt
152
+ print(f" ({name} yok, {alt.name} kullanılıyor)")
153
+ else:
154
+ raise FileNotFoundError(
155
+ f"Checkpoint yok: {ckpt_path} (run_dir={run_dir})"
156
+ )
157
+
158
+ print(f"Checkpoint: {ckpt_path}")
159
+ ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
160
+
161
+ # Version tespit
162
+ version = detect_version(ckpt)
163
+ if version != args.version and not args.ckpt:
164
+ print(f" ! Algılanan version={version}, --version={args.version}")
165
+ print(f"Version: {version}")
166
+
167
+ # Model
168
+ model, cfg, desc = build_model(version, ckpt, device)
169
+ # State dict — torch.compile prefix temizle
170
+ state = ckpt["model"]
171
+ state = {k.replace("_orig_mod.", ""): v for k, v in state.items()}
172
+ model.load_state_dict(state)
173
+ model.eval()
174
+
175
+ step = ckpt.get("step", "?")
176
+ val = ckpt.get("best_val", None)
177
+ val_str = f", val={val:.4f}" if val is not None and val != float("inf") else ""
178
+ n_params = model.num_params() if hasattr(model, "num_params") else \
179
+ sum(p.numel() for p in model.parameters())
180
+ print(f"Model: {desc}")
181
+ print(f" {n_params/1e6:.2f}M param (step={step}, "
182
+ f"version={ckpt.get('version','?')}{val_str})")
183
+
184
+ # Tokenizer
185
+ tok_path = args.tokenizer or str(DATA_DIR / TOKENIZERS[version])
186
+ if not Path(tok_path).exists():
187
+ raise FileNotFoundError(f"Tokenizer yok: {tok_path}")
188
+ tokenizer = Tokenizer.from_file(tok_path)
189
+ print(f"Tokenizer: {Path(tok_path).name} "
190
+ f"(vocab={tokenizer.get_vocab_size()})")
191
+
192
+ # ChatML format (--chat veya version=*-instruct otomatik)
193
+ raw_version = ckpt.get("version", "")
194
+ auto_chat = any(t in str(raw_version) for t in ("instruct", "sft", "dpo", "chat"))
195
+ use_chat = args.chat or auto_chat
196
+
197
+ if use_chat:
198
+ user_msg = (f"{args.instruction}\n{args.prompt}"
199
+ if args.instruction else args.prompt)
200
+ formatted = f"<|user|>\n{user_msg}\n<|assistant|>\n"
201
+ print(f"\nChatML format AKTİF (version={raw_version})")
202
+ print(f"User prompt: {user_msg!r}")
203
+ else:
204
+ formatted = args.prompt
205
+ print(f"\nRaw prompt: {args.prompt!r}")
206
+
207
+ print(f"Settings: max={args.max_tokens}, temp={args.temperature}, "
208
+ f"top_k={args.top_k}, rep_pen={args.repetition_penalty}, "
209
+ f"no_rep_ngram={args.no_repeat_ngram}")
210
+ print("=" * 70)
211
+
212
+ ids = tokenizer.encode(formatted).ids
213
+ x = torch.tensor([ids], dtype=torch.long, device=device)
214
+
215
+ use_bf16 = device == "cuda" and torch.cuda.is_bf16_supported()
216
+ dtype = torch.bfloat16 if use_bf16 else torch.float32
217
+
218
+ # Context window — V5 block_size=2048 (RoPE buffer x2 = 4096'ya kadar uzar)
219
+ max_ctx = cfg.block_size
220
+
221
+ for i in range(args.num_samples):
222
+ # Context overflow koruması
223
+ cur_ids = ids
224
+ if len(cur_ids) >= max_ctx:
225
+ print(f" ! Prompt {len(cur_ids)} token, "
226
+ f"context {max_ctx} → kırpılıyor")
227
+ cur_ids = cur_ids[-(max_ctx - args.max_tokens):]
228
+ x_i = torch.tensor([cur_ids], dtype=torch.long, device=device)
229
+
230
+ amp_ctx = (torch.amp.autocast(device_type="cuda", dtype=dtype)
231
+ if device == "cuda" else torch.no_grad())
232
+ with amp_ctx, torch.no_grad():
233
+ out = model.generate(
234
+ x_i,
235
+ max_new_tokens=args.max_tokens,
236
+ temperature=args.temperature,
237
+ top_k=args.top_k,
238
+ repetition_penalty=args.repetition_penalty,
239
+ no_repeat_ngram_size=args.no_repeat_ngram,
240
+ )
241
+ text = tokenizer.decode(out[0].tolist())
242
+ print(f"\n--- Sample {i+1} ---")
243
+ print(text)
244
+ print()
245
+
246
+
247
+ if __name__ == "__main__":
248
+ main()