| |
| """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() |
|
|