Pragya / inference.py
ArushBuilds's picture
Update inference.py
07adfe3
Raw History Blame
21.9 kB
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()