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