Spaces:
Running
Running
Download inference.py from ArushBuilds/Pragya: direct link, hf CLI and curl.
- Browser
- Download file 21.9 kB
-
https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/main/inference.py
- Command line
-
hf download hf://spaces/ArushBuilds/Pragya/inference.py
-
curl -L -o inference.py https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/main/inference.py
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) | |
| 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() |