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()