import argparse import contextlib import sys from collections import Counter import torch from model import GPT, GPTConfig from tokenizer import TokenizerWrapper # ───────────────────────────────────────────────────────────────────────────── # Device selection # ───────────────────────────────────────────────────────────────────────────── def pick_device(requested: str | None) -> torch.device: if requested is not None: return torch.device(requested) if torch.cuda.is_available(): return torch.device("cuda") try: import torch_xla.core.xla_model as xm # noqa: F401 return torch.device("xla") except ImportError: pass if torch.backends.mps.is_available(): return torch.device("mps") return torch.device("cpu") # ───────────────────────────────────────────────────────────────────────────── # Checkpoint loading # ───────────────────────────────────────────────────────────────────────────── def load_model(checkpoint_path: str, device: torch.device) -> tuple[GPT, GPTConfig]: checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=True) if "model_config" not in checkpoint: raise RuntimeError( f"'{checkpoint_path}' has no saved model_config -- it was written " "before save_checkpoint() started saving the architecture config " "alongside the weights. Either retrain (new checkpoints save it " "automatically), or construct GPTConfig(...) by hand here to " "match whatever config.py looked like when this checkpoint was " "produced, and call GPT(that_config) instead of this function." ) model_config = GPTConfig(**checkpoint["model_config"]) model = GPT(model_config).to(device) state_dict = checkpoint["model_state_dict"] if any(k.startswith("_orig_mod.") for k in state_dict): state_dict = { (k[len("_orig_mod."):] if k.startswith("_orig_mod.") else k): v for k, v in state_dict.items() } print( "[Inference] Checkpoint was saved from a torch.compile()'d model " "(keys prefixed with '_orig_mod.') -- stripped the prefix to " "load onto this plain GPT instance.", file=sys.stderr, ) # Load fp32 weights into the plain model BEFORE quantizing. # Quantizing first replaces nn.Linear submodules with tensor-subclass # equivalents; load_state_dict then tries to copy float weights into them # which either raises or silently de-quantizes. Correct order: weights → quant. model.load_state_dict(state_dict) quant = checkpoint.get("quantization") if quant: import model as _m fn = {"int8": _m.apply_int8_to_model, "nvfp4": _m.apply_nvfp4_to_model}.get(quant) if fn is None: raise RuntimeError(f"checkpoint quantized with {quant!r} but this build can't re-apply it") fn(model) model._quant_mode = quant model.eval() n_params = model.get_num_params() / 1e6 step = checkpoint.get("step", "?") print(f"[Inference] Loaded checkpoint (step {step}, {n_params:.2f}M params) onto {device}", file=sys.stderr) return model, model_config # ───────────────────────────────────────────────────────────────────────────── # Sampling # ───────────────────────────────────────────────────────────────────────────── def _apply_repetition_penalty( logits: torch.Tensor, generated_ids: list[int], penalty: float, prompt_id_set: "list[int] | set[int] | None" = None, max_count: int = 6, ) -> torch.Tensor: """CTRL-style penalty, COUNT-SCALED: a token seen N times in the *generated* portion gets divided by penalty**N (if positive) / multiplied by penalty**N (if negative). Count-scaling vs flat penalty ------------------------------ The flat version (divide once regardless of occurrence count) is too weak on small/repetitive-domain models (e.g. TinyStories-style data where "so so so happy!" is a genuinely common training pattern): once the model locks onto a token, the penalty never gets any stronger, so a mildly favoured looping token never gets pushed low enough to escape. Scaling by occurrence count means each additional repeat compounds the penalty, so a 3-4 token-deep loop gets meaningfully suppressed even if a single application wouldn't have been enough. Why penalty counts are clamped to max_count -------------------------------------------- Without a cap, a legitimate word that happens to appear many times in a long generation (e.g. "the", "a", punctuation) accumulates a factor of penalty**N that grows without bound and can suppress even high-confidence logits into noise. max_count=6 lets count-scaling do its job on real loops (which repeat in short runs) while preventing runaway suppression of common tokens over a long context. Why prompt tokens are excluded from the count ----------------------------------------------- generated_ids is seeded with the full prompt before the first token is produced. Without the exclusion, common tokens in a 20-token prompt already arrive at the first generation step with count >= 2-3, giving them a standing penalty of penalty**3 before the model has emitted a single token. That's not repetition suppression -- it's prompt interference. prompt_id_set (built once per generate() call from the original prompt_ids) lets us count only within the generated portion. """ if penalty == 1.0 or not generated_ids: return logits if prompt_id_set: # Count only tokens that appear in the generated suffix (after the # prompt). generated_ids includes the prompt prefix, so we walk the # list and tally only tokens not in the prompt -- tokens that the # model itself produced. Prompt tokens still receive the flat # penalty**1 if they recur in the generated portion (they're not # immune; we just don't pre-charge them for being in the prompt). gen_counts: Counter[int] = Counter() prompt_remaining = _build_prompt_remaining(list(prompt_id_set)) for tid in generated_ids: if prompt_remaining.get(tid, 0) > 0: prompt_remaining[tid] -= 1 continue gen_counts[tid] += 1 counts = gen_counts else: counts = Counter(generated_ids) if not counts: return logits sorted_tokens = sorted(counts.keys()) seen = torch.tensor(sorted_tokens, device=logits.device, dtype=torch.long) # Clamp counts to max_count so common tokens don't accumulate a # runaway factor over a long generation. raw_reps = [min(counts[t], max_count) for t in sorted_tokens] reps = torch.tensor(raw_reps, device=logits.device, dtype=logits.dtype) vals = logits[0, seen] factor = penalty ** reps logits[0, seen] = torch.where(vals > 0, vals / factor, vals * factor) return logits def _block_repeated_ngrams(logits: torch.Tensor, generated_ids: list[int], ngram_size: int) -> torch.Tensor: """Hard n-gram repeat block (HF's `no_repeat_ngram_size`, same idea). Look at the last (ngram_size - 1) generated tokens. If that same prefix has appeared earlier in generated_ids, whatever token followed it there is banned from being chosen again right now (logit -> -inf). This is a HARD constraint, unlike repetition_penalty which only nudges probabilities down -- it's what actually stops "so so so so so..." dead rather than just making it less likely each step. Only kicks in once at least ngram_size-1 tokens have been generated. """ if ngram_size <= 0 or len(generated_ids) < ngram_size - 1: return logits prefix_len = ngram_size - 1 if prefix_len == 0: return logits current_prefix = tuple(generated_ids[-prefix_len:]) banned = set() for i in range(len(generated_ids) - prefix_len): if tuple(generated_ids[i:i + prefix_len]) == current_prefix: banned.add(generated_ids[i + prefix_len]) if banned: banned_idx = torch.tensor(sorted(banned), device=logits.device, dtype=torch.long) logits[0, banned_idx] = float("-inf") return logits def _top_k_filter(logits: torch.Tensor, top_k: int) -> torch.Tensor: k = min(top_k, logits.size(-1)) values, _ = torch.topk(logits, k) logits[logits < values[:, [-1]]] = float("-inf") return logits def _top_p_filter(logits: torch.Tensor, top_p: float) -> torch.Tensor: """Nucleus sampling: keep the smallest prefix of sorted tokens whose cumulative probability crosses top_p, drop the rest.""" sorted_logits, sorted_idx = torch.sort(logits, descending=True) probs = torch.softmax(sorted_logits, dim=-1) cum_probs = torch.cumsum(probs, dim=-1) # remove token i if the cumulative prob BEFORE it already exceeds top_p # (i.e. it wasn't needed to cross the threshold) -- matches the standard # HF TopPLogitsWarper convention. remove = (cum_probs - probs) > top_p sorted_logits = sorted_logits.masked_fill(remove, float("-inf")) return torch.full_like(logits, float("-inf")).scatter(1, sorted_idx, sorted_logits) def _build_prompt_remaining(prompt_ids: list[int]) -> Counter[int]: """Return the prompt token multiplicities needed by the repetition penalty walker. A plain set hides duplicate counts, which underestimates how many times a repeated prompt token should be available to the counter. """ return Counter(prompt_ids) def _clip_max_new_tokens(max_new_tokens: int, prompt_ids: list[int], block_size: int) -> int: """Prevent the generation loop from silently overwriting the prompt by clipping the available completion budget to the remaining context window. """ if block_size <= 0: return 0 available = block_size - len(prompt_ids) if available <= 0: raise ValueError( f"prompt is {len(prompt_ids)} tokens >= block_size {block_size}; nothing " f"left to generate. Shorten the prompt (this build has no sliding-window " f"generation)." ) if max_new_tokens > available: print(f"[generate] NOTE: clipped max_new_tokens {max_new_tokens} -> {available} " f"(block_size={block_size}, prompt={len(prompt_ids)}).", file=sys.stderr) return min(max_new_tokens, available) @torch.no_grad() def generate( model: GPT, tokenizer: TokenizerWrapper, prompt: str, device: torch.device, max_new_tokens: int = 200, temperature: float = 1.0, top_k: int | None = None, top_p: float | None = None, repetition_penalty: float = 1.0, no_repeat_ngram_size: int = 0, stream: bool = True, seed: int | None = None, num_samples: int = 1, ) -> list[str]: if seed is not None: torch.manual_seed(seed) prompt_ids = tokenizer.encode(prompt, add_bos=True, add_eos=False) if len(prompt_ids) > model.config.block_size: prompt_ids = prompt_ids[-model.config.block_size:] idx = torch.tensor([prompt_ids] * num_samples, dtype=torch.long, device=device) block_size = model.config.block_size eos_id = getattr(tokenizer, "eos_id", None) if eos_id is None: eos_id = getattr(tokenizer, "eos_token_id", None) if eos_id is None: print("[generate] WARNING: tokenizer exposes no eos_id/eos_token_id -- " "early stopping disabled, generations run to budget.", file=sys.stderr) # Built once from the original prompt and kept as a multiset so # repetition penalty excludes prompt tokens by duplicate-aware count. prompt_id_set: list[int] = list(prompt_ids) # Clip the requested budget to the remaining context budget so the # context window never drops the prompt silently when max_new_tokens # exceeds the model's terminal block_size choice. max_new_tokens = _clip_max_new_tokens(max_new_tokens, prompt_ids, block_size) generated_ids = [list(prompt_ids) for _ in range(num_samples)] finished = [False] * num_samples # Decode the FULL sequence each step and print only the new suffix, # rather than decoding one new token at a time: BPE/subword tokens don't # each cleanly map to standalone text (a token can be half a multi-byte # character or mid-word piece), so decoding incrementally token-by-token # can emit garbled text at token boundaries. Re-decoding the whole # sequence and diffing against what's already been printed sidesteps # that regardless of the tokenizer's internals. O(T) decode cost per # step, which is fine at CLI-generation scale. prev_text = [tokenizer.decode(g) for g in generated_ids] if stream else None if stream and num_samples == 1: print(prev_text[0], end="", flush=True) for _ in range(max_new_tokens): if all(finished): break idx_cond = idx if idx.size(1) <= block_size else idx[:, -block_size:] logits, _ = model(idx_cond) logits = logits[:, -1, :].float() # last position, fp32 for stable sampling math regardless of training dtype for b in range(num_samples): if finished[b]: continue row = logits[b:b + 1] row = _apply_repetition_penalty(row, generated_ids[b], repetition_penalty, prompt_id_set=prompt_id_set) row = _block_repeated_ngrams(row, generated_ids[b], no_repeat_ngram_size) row = row / max(temperature, 1e-5) if top_k is not None: row = _top_k_filter(row, top_k) if top_p is not None: row = _top_p_filter(row, top_p) logits[b:b + 1] = row probs = torch.softmax(logits, dim=-1) dead = ~torch.isfinite(probs).any(dim=-1) | (probs.sum(dim=-1) <= 0) if dead.any(): logits[dead] = 0.0 probs = torch.softmax(logits, dim=-1) next_ids = torch.multinomial(probs, num_samples=1) # (num_samples, 1) # Once a sequence hits EOS, keep feeding it its own last token so the # batch stays rectangular, but stop appending to its recorded output. for b in range(num_samples): if finished[b]: next_ids[b, 0] = idx[b, -1] continue tid = next_ids[b, 0].item() generated_ids[b].append(tid) if eos_id is not None and tid == eos_id: finished[b] = True idx = torch.cat([idx, next_ids], dim=1) if stream and num_samples == 1: new_text = tokenizer.decode(generated_ids[0]) print(new_text[len(prev_text[0]):], end="", flush=True) prev_text[0] = new_text if stream and num_samples == 1: print() return [tokenizer.decode(g) for g in generated_ids] # ───────────────────────────────────────────────────────────────────────────── # CLI # ───────────────────────────────────────────────────────────────────────────── def build_arg_parser() -> argparse.ArgumentParser: p = argparse.ArgumentParser(description="Generate text from a trained Arya/Veylon checkpoint.") p.add_argument("--checkpoint", required=True, help="Path to a checkpoint saved by train.py") p.add_argument("--tokenizer", default="./tokenizer", help="Path passed to TokenizerWrapper") p.add_argument("--prompt", default=None, help="Prompt text (omit if using --interactive)") p.add_argument("--interactive", action="store_true", help="Drop into a REPL: one prompt per line") 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=50) p.add_argument("--top-p", type=float, default=None, help="Nucleus sampling threshold, e.g. 0.9. Combine freely with --top-k.") p.add_argument("--repetition-penalty", type=float, default=1.3, help="CTRL-style penalty on already-generated tokens, now count-scaled " "(penalty**occurrences). 1.0 disables it -- small models loop into " "degenerate repetition ('so so so so...') without some penalty here, " "so this defaults ON.") p.add_argument("--no-repeat-ngram-size", type=int, default=3, help="Hard-ban repeating any n-gram of this size (HF-style). 0 disables it. " "This is what actually stops infinite word loops -- repetition-penalty " "alone only makes them less likely, this makes them impossible.") p.add_argument("--num-samples", type=int, default=1, help="Generate N completions per prompt (batched)") p.add_argument("--seed", type=int, default=None) p.add_argument("--device", default=None, help="cuda / cpu / mps / xla -- auto-detected if omitted") p.add_argument("--dtype", default="auto", choices=["auto", "fp32", "fp16", "bf16"], help="Inference precision. 'auto' uses fp16 on CUDA, fp32 elsewhere.") return p def _resolve_dtype(dtype_arg: str, device: torch.device) -> torch.dtype | None: if device.type not in ("cuda", "cpu"): # torch.autocast is only exercised here against cuda/cpu; MPS and XLA # handle precision differently (XLA in particular usually wants # XLA_USE_BF16 or torch_xla's own autocast, not torch.autocast). # Guessing at that without real hardware to test against would be # worse than just running those in fp32 and saying so. if dtype_arg != "auto": print(f"[Inference] --dtype={dtype_arg} requested but autocast isn't wired up for device " f"'{device.type}' -- running in fp32.", file=sys.stderr) return None if dtype_arg == "fp32": return None # no autocast if dtype_arg == "fp16": if device.type == "cpu": print("[Inference] fp16 autocast unsupported on CPU -- using bf16 or fp32.", file=sys.stderr) return torch.bfloat16 return torch.float16 if dtype_arg == "bf16": return torch.bfloat16 # auto return torch.float16 if device.type == "cuda" else None def main() -> None: args = build_arg_parser().parse_args() if not args.prompt and not args.interactive: print("Provide --prompt \"...\" or pass --interactive.", file=sys.stderr) sys.exit(1) device = pick_device(args.device) model, model_config = load_model(args.checkpoint, device) tokenizer = TokenizerWrapper(args.tokenizer) autocast_dtype = _resolve_dtype(args.dtype, device) def run(prompt: str, stream: bool) -> list[str]: ctx = torch.autocast(device_type=device.type, dtype=autocast_dtype) if autocast_dtype else contextlib.nullcontext() with ctx: return generate( model, tokenizer, prompt, device, max_new_tokens=args.max_new_tokens, temperature=args.temperature, top_k=args.top_k, top_p=args.top_p, repetition_penalty=args.repetition_penalty, no_repeat_ngram_size=args.no_repeat_ngram_size, stream=stream, seed=args.seed, num_samples=args.num_samples, ) if args.interactive: print("[Inference] Interactive mode. Ctrl-C or empty line to exit.", file=sys.stderr) while True: try: prompt = input("\n>>> ") except (EOFError, KeyboardInterrupt): break if not prompt.strip(): break outputs = run(prompt, stream=(args.num_samples == 1)) if args.num_samples > 1: for i, text in enumerate(outputs): print(f"\n--- sample {i} ---\n{text}") else: outputs = run(args.prompt, stream=(args.num_samples == 1)) if args.num_samples > 1: for i, text in enumerate(outputs): print(f"\n--- sample {i} ---\n{text}") if __name__ == "__main__": main()