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