Gemma-4-26B-A4B-it-W4A16-G64-BF16Vision / eval /evaluate_gemma4_ppl_vllm.py
Mitchins's picture
Initial W4A16 G64 release
8830ced verified
Raw
History Blame Contribute Delete
4.66 kB
#!/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()