#!/usr/bin/env python3 """Perplexity via vLLM prompt logprobs for a Gemma 4 text-only path. This is intentionally shared by BF16 and compressed-tensors W4A16 runs. It scores the observed next token at every noninitial position in deterministic, contiguous held-out WikiText-2 test windows. ``--cpu-offload-gb`` makes the otherwise too-large BF16 parent testable on a 24 GB card without changing the model or scoring implementation. """ from __future__ import annotations import argparse import json import math from datetime import datetime, timezone from pathlib import Path from datasets import load_dataset from transformers import AutoTokenizer from vllm import LLM, SamplingParams from vllm.inputs import TokensPrompt def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--model", type=Path, required=True) parser.add_argument("--tokenizer", type=Path, required=True) parser.add_argument("--label", required=True) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--cache-dir", type=Path, required=True) parser.add_argument("--num-windows", type=int, default=4) parser.add_argument("--window-tokens", type=int, default=512) parser.add_argument("--quantization", default=None) parser.add_argument("--cpu-offload-gb", type=float, default=0.0) return parser.parse_args() def held_out_windows(tokenizer, cache_dir: Path, num_windows: int, size: int): dataset = load_dataset( "Salesforce/wikitext", "wikitext-2-raw-v1", split="test", cache_dir=str(cache_dir), ) text = "\n\n".join(row["text"] for row in dataset if row["text"].strip()) ids = tokenizer(text, add_special_tokens=False)["input_ids"] needed = num_windows * size if len(ids) < needed: raise RuntimeError(f"Need {needed} tokens, corpus yielded {len(ids)}") return dataset, [ids[i * size : (i + 1) * size] for i in range(num_windows)] def main() -> None: args = parse_args() if args.window_tokens < 2: raise ValueError("--window-tokens must be at least 2") tokenizer = AutoTokenizer.from_pretrained(args.tokenizer, local_files_only=True) dataset, windows = held_out_windows( tokenizer, args.cache_dir, args.num_windows, args.window_tokens ) llm = LLM( model=str(args.model), tokenizer=str(args.tokenizer), dtype="bfloat16", quantization=args.quantization, max_model_len=args.window_tokens + 1, max_num_seqs=1, max_num_batched_tokens=args.window_tokens + 1, gpu_memory_utilization=0.80, cpu_offload_gb=args.cpu_offload_gb, language_model_only=True, limit_mm_per_prompt={"image": 0, "video": 0}, ) params = SamplingParams( temperature=0.0, max_tokens=1, ignore_eos=True, prompt_logprobs=1, detokenize=False, ) outputs = llm.generate( [TokensPrompt(prompt_token_ids=ids) for ids in windows], params, use_tqdm=False ) nll = 0.0 token_count = 0 for window, output in zip(windows, outputs, strict=True): values = output.prompt_logprobs if values is None or len(values) != len(window): raise RuntimeError("vLLM did not return one prompt-logprob entry per prompt token") for token_id, entry in zip(window[1:], values[1:], strict=True): if entry is None or token_id not in entry: raise RuntimeError("observed token missing from prompt-logprob response") nll -= entry[token_id].logprob token_count += 1 result = { "label": args.label, "model": args.model.name, "tokenizer": args.tokenizer.name, "dataset": "Salesforce/wikitext", "dataset_config": "wikitext-2-raw-v1", "split": "test", "dataset_fingerprint": dataset._fingerprint, "window_selection": "first contiguous non-empty test-corpus token windows", "num_windows": args.num_windows, "window_tokens": args.window_tokens, "evaluated_next_tokens": token_count, "nll_sum": nll, "mean_nll": nll / token_count, "perplexity": math.exp(nll / token_count), "engine": "vLLM prompt_logprobs=1", "quantization": args.quantization, "cpu_offload_gb": args.cpu_offload_gb, "utc": datetime.now(timezone.utc).isoformat(), } args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(result, indent=2) + "\n") print(json.dumps(result, indent=2)) if __name__ == "__main__": main()