diffusion-51M-base / evaluation /run_continuation.py
d0rj's picture
Publish evaluated diffusion v2 with PLL intervals and TensorBoard traces
80aea5b verified
Raw
History Blame Contribute Delete
3.23 kB
"""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()