File size: 4,659 Bytes
8830ced
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
#!/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()