nanogpt-tr-v4 / 06_sample.py
musabc's picture
Upload 06_sample.py with huggingface_hub
ce562ac verified
Raw
History Blame Contribute Delete
4.98 kB
"""
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
# V3 modeli opsiyonel — bazı ortamlarda sadece V4 olabilir
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" # default: best, --latest ile latest
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")
# Checkpoint yukle
print(f"Checkpoint: {ckpt_path}")
ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
# V3 vs V4 ayrimi: V4 config'inde 'rope_theta' var
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 = Tokenizer.from_file(str(DATA_DIR / "tokenizer-tr-16k.json"))
# ChatML format (--chat veya version=v4-instruct otomatik)
auto_chat = version in ("v4-instruct", "v4-instruct-v2", "v4-dpo")
use_chat = args.chat or auto_chat
if use_chat:
# ChatML formatına çevir (SFT eğitimindeki format ile aynı)
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()