Files
mini-llm/infer.py
T

199 lines
6.0 KiB
Python

import argparse
import json
import re
from pathlib import Path
import torch
from model import MiniLM, ModelConfig
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Run inference with a trained mini-10m checkpoint")
p.add_argument("--checkpoint", type=str, default=None, help="Path to checkpoint .pt file")
p.add_argument("--run-dir", type=str, default="runs/tiny10m", help="Run dir containing step_*.pt")
p.add_argument("--prompt", type=str, default="Once upon a time")
p.add_argument("--max-new-tokens", type=int, default=200)
p.add_argument("--temperature", type=float, default=0.8)
p.add_argument("--top-k", type=int, default=40)
p.add_argument("--top-p", type=float, default=0.95)
p.add_argument("--seed", type=int, default=1337)
p.add_argument("--device", type=str, default="auto")
p.add_argument("--dtype", type=str, default="auto", choices=["auto", "float32", "bfloat16"])
p.add_argument(
"--max-seq-len",
type=int,
default=None,
help="Override model max context window. If omitted, tries run config then defaults to 16384.",
)
return p.parse_args()
def detect_device(user_value: str) -> torch.device:
if user_value != "auto":
return torch.device(user_value)
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
def detect_dtype(user_value: str, device: torch.device) -> torch.dtype:
if user_value == "float32":
return torch.float32
if user_value == "bfloat16":
return torch.bfloat16
if device.type == "cuda" and torch.cuda.is_bf16_supported():
return torch.bfloat16
return torch.float32
def encode_byte_level(text: str) -> list[int]:
return list(text.encode("utf-8", errors="ignore"))
def decode_byte_level(tokens: list[int]) -> str:
return bytes(tokens).decode("utf-8", errors="replace")
def find_latest_checkpoint(run_dir: Path) -> Path:
ckpts = list(run_dir.glob("step_*.pt"))
if not ckpts:
raise FileNotFoundError(f"No checkpoints found in {run_dir}")
def step_num(path: Path) -> int:
m = re.search(r"step_(\d+)\.pt$", path.name)
return int(m.group(1)) if m else -1
return max(ckpts, key=step_num)
def maybe_load_run_config(run_dir: Path) -> dict:
cfg_path = run_dir / "config.json"
if not cfg_path.exists():
return {}
with open(cfg_path, "r", encoding="utf-8") as f:
return json.load(f)
def build_model(max_seq_len: int) -> MiniLM:
cfg = ModelConfig(
vocab_size=256,
d_model=352,
n_heads=8,
n_layers=7,
ffn_mult=4,
max_seq_len=max_seq_len,
dropout=0.0,
)
return MiniLM(cfg)
def sample_next_token(
logits: torch.Tensor,
temperature: float,
top_k: int,
top_p: float,
) -> int:
if temperature <= 0:
return int(torch.argmax(logits).item())
logits = logits / temperature
if top_k > 0:
k = min(top_k, logits.size(-1))
values, _ = torch.topk(logits, k)
cutoff = values[..., -1]
logits = torch.where(logits < cutoff, torch.full_like(logits, float("-inf")), logits)
if 0.0 < top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
sorted_probs = torch.softmax(sorted_logits, dim=-1)
cumulative_probs = torch.cumsum(sorted_probs, dim=-1)
sorted_mask = cumulative_probs > top_p
sorted_mask[..., 1:] = sorted_mask[..., :-1].clone()
sorted_mask[..., 0] = False
mask = torch.zeros_like(sorted_mask, dtype=torch.bool)
mask.scatter_(0, sorted_indices, sorted_mask)
logits = torch.where(mask, torch.full_like(logits, float("-inf")), logits)
probs = torch.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
return int(next_token.item())
@torch.no_grad()
def generate(
model: MiniLM,
prompt_ids: list[int],
max_new_tokens: int,
temperature: float,
top_k: int,
top_p: float,
device: torch.device,
amp_dtype: torch.dtype,
) -> list[int]:
out_ids = list(prompt_ids)
for _ in range(max_new_tokens):
context = out_ids[-model.cfg.max_seq_len :]
x = torch.tensor(context, dtype=torch.long, device=device).unsqueeze(0)
with torch.autocast(device_type=device.type, dtype=amp_dtype, enabled=device.type == "cuda"):
logits, _ = model(x)
next_logits = logits[0, -1]
next_id = sample_next_token(next_logits, temperature, top_k, top_p)
out_ids.append(next_id)
return out_ids
def main():
args = parse_args()
torch.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
device = detect_device(args.device)
amp_dtype = detect_dtype(args.dtype, device)
run_dir = Path(args.run_dir)
ckpt_path = Path(args.checkpoint) if args.checkpoint else find_latest_checkpoint(run_dir)
run_cfg = maybe_load_run_config(ckpt_path.parent)
max_seq_len = args.max_seq_len
if max_seq_len is None:
max_seq_len = int(run_cfg.get("max_seq_len", 16_384))
model = build_model(max_seq_len=max_seq_len).to(device)
ckpt = torch.load(ckpt_path, map_location=device)
state_dict = ckpt["model"] if isinstance(ckpt, dict) and "model" in ckpt else ckpt
model.load_state_dict(state_dict, strict=True)
model.eval()
prompt_ids = encode_byte_level(args.prompt)
if not prompt_ids:
raise ValueError("Prompt produced no byte tokens. Provide a non-empty prompt.")
out_ids = generate(
model=model,
prompt_ids=prompt_ids,
max_new_tokens=args.max_new_tokens,
temperature=args.temperature,
top_k=args.top_k,
top_p=args.top_p,
device=device,
amp_dtype=amp_dtype,
)
generated_suffix = out_ids[len(prompt_ids) :]
print(f"Device: {device}")
print(f"Checkpoint: {ckpt_path}")
print(f"Max context: {max_seq_len}")
print("--- PROMPT ---")
print(args.prompt)
print("--- GENERATED ---")
print(decode_byte_level(generated_suffix))
if __name__ == "__main__":
main()