| """ |
| Egitilmis modelden ornek metin uret. |
| |
| Kullanim: |
| python 06_sample.py |
| python 06_sample.py --prompt "Istanbul" --max-tokens 200 --temperature 0.7 |
| python 06_sample.py --num-samples 5 |
| """ |
|
|
| import argparse |
| from pathlib import Path |
|
|
| import torch |
| from tokenizers import Tokenizer |
|
|
| |
| try: |
| from model import GPT, GPTConfig |
| HAS_V3 = True |
| except ImportError: |
| HAS_V3 = False |
| GPT = GPTConfig = None |
|
|
| from model_v4 import GPTV4, GPTConfigV4 |
|
|
| DATA_DIR = Path(__file__).parent / "data" |
| RUN_DIR = Path(__file__).parent / "runs" / "tr-50m-v4" |
| CKPT_PATH = RUN_DIR / "best_ckpt.pt" |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--prompt", type=str, default="Türkiye") |
| parser.add_argument("--max-tokens", type=int, default=200) |
| parser.add_argument("--temperature", type=float, default=0.8) |
| parser.add_argument("--top-k", type=int, default=50) |
| parser.add_argument("--repetition-penalty", type=float, default=1.15) |
| parser.add_argument("--no-repeat-ngram", type=int, default=3) |
| parser.add_argument("--num-samples", type=int, default=3) |
| parser.add_argument("--ckpt", type=str, default=str(CKPT_PATH)) |
| parser.add_argument("--latest", action="store_true", |
| help="best yerine latest checkpoint'i kullan") |
| parser.add_argument("--chat", action="store_true", |
| help="SFT/Instruct ChatML formatı uygula") |
| parser.add_argument("--instruction", type=str, default=None, |
| help="ChatML için ayrı instruction (input ile birlikte)") |
| parser.add_argument("--seed", type=int, default=None) |
| args = parser.parse_args() |
|
|
| if args.seed is not None: |
| torch.manual_seed(args.seed) |
|
|
| device = "cuda" if torch.cuda.is_available() else "cpu" |
| print(f"Device: {device}") |
|
|
| ckpt_path = args.ckpt |
| if args.latest: |
| ckpt_path = str(RUN_DIR / "latest_ckpt.pt") |
|
|
| |
| print(f"Checkpoint: {ckpt_path}") |
| ckpt = torch.load(ckpt_path, map_location=device, weights_only=False) |
| |
| is_v4 = "rope_theta" in ckpt["config"] |
| if is_v4: |
| cfg = GPTConfigV4(**ckpt["config"]) |
| model = GPTV4(cfg).to(device) |
| print("Model: V4 (RoPE + RMSNorm + SwiGLU + QK-norm)") |
| else: |
| if not HAS_V3: |
| raise ImportError( |
| "V3 checkpoint ama model.py yok. V3 için model.py'yi de kopyala." |
| ) |
| cfg = GPTConfig(**ckpt["config"]) |
| model = GPT(cfg).to(device) |
| print("Model: V3 (LayerNorm + GELU + learned PE)") |
| model.load_state_dict(ckpt["model"]) |
| model.eval() |
| step = ckpt.get("step", "?") |
| val = ckpt.get("best_val", None) |
| version = ckpt.get("version", "base") |
| val_str = f", val={val:.4f}" if val is not None else "" |
| n_params = model.num_params() if hasattr(model, "num_params") else \ |
| sum(p.numel() for p in model.parameters()) |
| print(f"Model: {n_params/1e6:.2f}M param " |
| f"(step={step}, version={version}{val_str})") |
|
|
| |
| tokenizer = Tokenizer.from_file(str(DATA_DIR / "tokenizer-tr-16k.json")) |
|
|
| |
| auto_chat = version in ("v4-instruct", "v4-instruct-v2", "v4-dpo") |
| use_chat = args.chat or auto_chat |
|
|
| if use_chat: |
| |
| if args.instruction: |
| user_msg = f"{args.instruction}\n{args.prompt}" |
| else: |
| user_msg = args.prompt |
| formatted = f"<|user|>\n{user_msg}\n<|assistant|>\n" |
| print(f"\nChatML format AKTİF (version={version})") |
| print(f"User prompt: {user_msg!r}") |
| else: |
| formatted = args.prompt |
| print(f"\nRaw prompt: {args.prompt!r}") |
|
|
| print(f"Settings: max={args.max_tokens}, temp={args.temperature}, top_k={args.top_k}") |
| print("=" * 70) |
|
|
| ids = tokenizer.encode(formatted).ids |
| x = torch.tensor([ids], dtype=torch.long, device=device) |
|
|
| use_bf16 = device == "cuda" and torch.cuda.is_bf16_supported() |
| dtype = torch.bfloat16 if use_bf16 else torch.float32 |
|
|
| for i in range(args.num_samples): |
| with torch.amp.autocast(device_type="cuda", dtype=dtype) \ |
| if device == "cuda" else torch.no_grad(): |
| with torch.no_grad(): |
| out = model.generate( |
| x.clone(), |
| max_new_tokens=args.max_tokens, |
| temperature=args.temperature, |
| top_k=args.top_k, |
| repetition_penalty=args.repetition_penalty, |
| no_repeat_ngram_size=args.no_repeat_ngram, |
| ) |
| text = tokenizer.decode(out[0].tolist()) |
| print(f"\n--- Sample {i+1} ---") |
| print(text) |
| print() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|