File size: 3,226 Bytes
80aea5b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Standalone WikiText continuation PLL reproduction for the exported repository."""
import argparse
import hashlib
import json
from pathlib import Path
import numpy as np
import torch
import torch.nn.functional as F
from datasets import load_dataset
from transformers import AutoModelForMaskedLM, AutoTokenizer


@torch.inference_mode()
def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output', type=Path, required=True)
    parser.add_argument('--limit-blocks', type=int, help='Smoke only; never report as a full benchmark')
    args = parser.parse_args()
    if args.limit_blocks is not None and args.limit_blocks < 1:
        parser.error('--limit-blocks must be positive')
    if args.output.exists():
        raise FileExistsError(args.output)
    root = Path(__file__).resolve().parents[1]
    torch.set_num_threads(4)
    torch.backends.cuda.matmul.allow_tf32 = False
    tokenizer = AutoTokenizer.from_pretrained(root, trust_remote_code=True)
    tokenizer.model_max_length = 10**9
    data = load_dataset('Salesforce/wikitext', 'wikitext-2-raw-v1', split='test')
    text = '\n'.join(data['text'])
    ids = tokenizer.encode(text, add_special_tokens=False)
    blocks = torch.tensor(ids[:len(ids)//1024*1024]).reshape(-1, 1024)
    if args.limit_blocks is not None:
        blocks = blocks[:args.limit_blocks]
    # Match the measured protocol, including dtype conversion of RoPE buffers.
    model = AutoModelForMaskedLM.from_pretrained(root, trust_remote_code=True, dtype=torch.float32).cuda().eval()
    model.to(torch.bfloat16)
    values = []
    for i, block in enumerate(blocks):
        base = block.cuda()
        total = 0.
        for start in range(512, 1024, 16):
            positions = torch.arange(start, start+16, device='cuda')
            rows = torch.arange(16, device='cuda')
            masked = base[None].expand(16, -1).clone()
            masked[rows, positions] = model.config.mask_token_id
            hidden = model.model(masked, timesteps=torch.full((16,), 1/1024, device='cuda'))
            logits = model.lm_head(hidden[rows, positions]).float()
            total += F.cross_entropy(logits, base[positions], reduction='sum').item()
        values.append(total/512)
        if i % 10 == 0:
            print('PLL block', i, flush=True)
    array = np.asarray(values)
    rng = np.random.default_rng(2026)
    lo, hi = np.quantile(array[rng.integers(len(array), size=(10000, len(array)))].mean(1), [.025, .975])
    result = dict(nll=float(array.mean()), nll_ci95=[float(lo), float(hi)],
                  ppl=float(np.exp(array.mean())), ppl_ci95=[float(np.exp(lo)), float(np.exp(hi))],
                  block_nll=values, blocks=len(blocks), scored_tokens=len(blocks)*512,
                  dropped_tail_tokens=len(ids)%1024, corpus_sha256=hashlib.sha256(text.encode()).hexdigest(),
                  protocol='Single-mask continuation PLL; pseudo-perplexity is not AR PPL',
                  smoke_only=args.limit_blocks is not None,
                  dtype='bfloat16', device='cuda', bootstrap_samples=10000, bootstrap_seed=2026)
    args.output.write_text(json.dumps(result, indent=2) + '\n')


if __name__ == '__main__':
    main()