d0rj's picture
Publish evaluated diffusion v2 with PLL intervals and TensorBoard traces
80aea5b verified
Raw
History Blame Contribute Delete
11.6 kB
"""HF loading and an explicitly experimental masked-model adapter.
Only imported for an actual evaluation, never for planning.
"""
import torch
from lm_eval.models.huggingface import HFLM
from transformers import AutoTokenizer
from tqdm import tqdm
class UL2HFLM(HFLM):
"""S-denoiser continuation scoring; only answer text contributes to NLL.
Source: S + context + sentinel_0 + EOS.
Decoder input: BOS + sentinel_0 + answer[:-1]. Controls are conditioned
on, never scored; the softmax still includes the complete vocabulary.
"""
def _encode_pair(self, context, continuation):
# Match causal HFLM's text boundary, without tokenizer-added controls.
spaces = len(context) - len(context.rstrip())
if spaces:
continuation = context[-spaces:] + continuation
context = context[:-spaces]
whole = self.tok_encode(context + continuation, add_special_tokens=False)
prefix = self.tok_encode(context, add_special_tokens=False)
return prefix, whole[len(prefix):]
@torch.inference_mode()
def _loglikelihood_tokens(self, requests, disable_tqdm=False, override_bs=None):
contract = self.model.config.ul2
span = contract["sentinel_ids"][0]
bos = self.model.config.decoder_start_token_id
batch_size = override_bs or self.batch_size
results = [None] * len(requests)
ordered = sorted(enumerate(requests), key=lambda x: -(len(x[1][1]) + len(x[1][2])))
for start in tqdm(range(0, len(ordered), batch_size), disable=disable_tqdm, desc="UL2 likelihood"):
batch = ordered[start:start + batch_size]
sources, decoders = [], []
for _, (_, context, answer) in batch:
if not answer or len(answer) >= self.max_length:
raise ValueError("UL2 scoring requires a nonempty answer shorter than the context limit")
budget = min(self.max_length + 1 - len(answer), self.model.config.max_position_embeddings - 3)
context = context[-budget:]
sources.append([contract["mode_ids"]["S"], *context, span, contract["eos_id"]])
decoders.append([bos, span, *answer[:-1]])
def pad(rows):
ids = torch.full((len(rows), max(map(len, rows))), contract["pad_id"], dtype=torch.long, device=self.device)
mask = torch.zeros_like(ids)
for i, row in enumerate(rows):
ids[i, :len(row)] = torch.tensor(row, device=self.device)
mask[i, :len(row)] = 1
return ids, mask
source, source_mask = pad(sources)
decoder, decoder_mask = pad(decoders)
hidden = self.model.model(input_ids=source, attention_mask=source_mask,
decoder_input_ids=decoder, decoder_attention_mask=decoder_mask,
use_cache=False).last_hidden_state
# Project only scored positions and bound temporary vocabulary logits.
for row, (index, (key, _, answer)) in enumerate(batch):
score, greedy = 0.0, True
for offset in range(0, len(answer), 256):
target = torch.tensor(answer[offset:offset + 256], device=self.device)
logits = self.model.lm_head(hidden[row, 1 + offset:1 + offset + len(target)]).float()
score += torch.log_softmax(logits, -1).gather(1, target[:, None]).sum().item()
greedy = greedy and bool((logits.argmax(-1) == target).all())
results[index] = (score, greedy)
if key is not None:
self.cache_hook.add_partial("loglikelihood", key, results[index])
return results
def loglikelihood_rolling(self, requests, disable_tqdm=False):
raise NotImplementedError("UL2 continuation likelihood is not rolling autoregressive perplexity")
def generate_until(self, requests, disable_tqdm=False):
raise NotImplementedError("UL2 benchmark adapter currently supports likelihood tasks only")
class PrefixHFLM(HFLM):
"""Score causal answers conditioned on a fully bidirectional text prefix."""
@torch.inference_mode()
def _loglikelihood_tokens(self, requests, disable_tqdm=False, override_bs=None):
results = [None] * len(requests)
ordered = sorted(enumerate(requests), key=lambda x: -(len(x[1][1]) + len(x[1][2])))
size = override_bs or self.batch_size
for start in tqdm(range(0, len(ordered), size), disable=disable_tqdm, desc="Prefix likelihood"):
batch = ordered[start:start + size]
rows, prefixes = [], []
for _, (_, context, answer) in batch:
if not answer or len(answer) > self.max_length:
raise ValueError("Require nonempty answer no longer than context limit")
context = context[-(self.max_length + 1 - len(answer)):]
rows.append(context + answer[:-1])
prefixes.append(len(context))
ids = torch.full((len(rows), max(map(len, rows))), self.tokenizer.pad_token_id or 0,
device=self.device, dtype=torch.long)
mask = torch.zeros_like(ids)
for i, row in enumerate(rows):
ids[i, :len(row)] = torch.tensor(row, device=self.device)
mask[i, :len(row)] = 1
hidden = self.model.model(ids, attention_mask=mask,
prefix_lengths=torch.tensor(prefixes, device=self.device),
use_cache=False).last_hidden_state
for i, (index, (key, _, answer)) in enumerate(batch):
score, greedy = 0., True
for offset in range(0, len(answer), 256):
target = torch.tensor(answer[offset:offset+256], device=self.device)
logits = self.model.lm_head(hidden[i, prefixes[i]-1+offset:prefixes[i]-1+offset+len(target)]).float()
score += torch.log_softmax(logits, -1).gather(1, target[:, None]).sum().item()
greedy = greedy and bool((logits.argmax(-1) == target).all())
results[index] = (score, greedy)
if key is not None:
self.cache_hook.add_partial("loglikelihood", key, results[index])
return results
def loglikelihood_rolling(self, requests, disable_tqdm=False):
raise NotImplementedError("Use an explicitly defined prefix-continuation protocol")
class DiffusionHFLM(HFLM):
"""Continuation PLL, with other answer tokens visible.
The bool is full-token reconstruction accuracy, NOT AR exact match.
Generation uses one masked next-token slot at a time, not the model's
unconditional parallel denoiser. Both are experimental protocols.
"""
@torch.inference_mode()
def _loglikelihood_tokens(self, requests, disable_tqdm=False, override_bs=None):
results = []
for cache_key, context, continuation in tqdm(requests, disable=disable_tqdm, desc="Diffusion PLL"):
if not continuation:
results.append((0.0, True))
continue
if len(continuation) >= self.max_length:
raise ValueError("Diffusion PLL needs room for context and the entire continuation")
context = context[-(self.max_length - len(continuation)):]
tokens = context + continuation
base = torch.tensor([tokens], dtype=torch.long, device=self.device)
timestep = torch.tensor([1 / len(tokens)], device=self.device)
score, greedy = 0.0, True
# Independent single-mask replicas preserve the PLL definition.
# Project only masked positions instead of every vocabulary logit.
for start in range(len(context), len(tokens), 16):
positions = torch.arange(start, min(start + 16, len(tokens)), device=self.device)
rows = torch.arange(len(positions), device=self.device)
masked = base.expand(len(positions), -1).clone()
masked[rows, positions] = self.model.config.mask_token_id
hidden = self.model.model(input_ids=masked, timesteps=timestep.expand(len(positions)))
logits = self.model.lm_head(hidden[rows, positions]).float()
targets = base[0, positions]
score += torch.log_softmax(logits, dim=-1).gather(1, targets[:, None]).sum().item()
greedy = greedy and bool((logits.argmax(-1) == targets).all())
result = (score, greedy)
results.append(result)
if cache_key is not None:
self.cache_hook.add_partial("loglikelihood", cache_key, result)
return results
def loglikelihood_rolling(self, requests, disable_tqdm=False):
raise NotImplementedError("Diffusion PLL must not be reported as autoregressive perplexity")
@torch.inference_mode()
def _model_generate(self, context, max_length, stop, **generation_kwargs):
if context.shape[0] != 1:
raise ValueError("Diffusion generation requires batch size 1")
if generation_kwargs.get("do_sample", False):
raise ValueError("Diffusion adapter implements greedy generation only")
sequence = context
prompt_length = context.shape[1]
while sequence.shape[1] < min(max_length, self.max_length):
slot = torch.full((1, 1), self.model.config.mask_token_id,
dtype=torch.long, device=self.device)
masked = torch.cat((sequence, slot), dim=1)
timestep = torch.tensor([1 / masked.shape[1]], device=self.device)
logits = self.model(input_ids=masked, timesteps=timestep).logits[:, -1]
predicted = logits.argmax(dim=-1, keepdim=True)
sequence = torch.cat((sequence, predicted), dim=1)
text = self.tok_decode(sequence[0, prompt_length:].tolist())
if predicted.item() == self.eot_token_id or any(s and s in text for s in stop):
break
return sequence
def load_model(plan, args):
common = dict(batch_size=args.batch_size, device=args.device, dtype=args.dtype,
max_length=args.max_length, trust_remote_code=args.trust_remote_code,
revision=args.revision)
tokenizer = AutoTokenizer.from_pretrained(
plan["tokenizer"], trust_remote_code=args.trust_remote_code,
revision=args.tokenizer_revision or args.revision,
)
if plan["local_model"] is None:
return HFLM(pretrained=plan["checkpoint"], tokenizer=tokenizer, **common)
from tiny_llm.models import MODEL_REGISTRY
entry = MODEL_REGISTRY[plan["local_model"]]
dtype = args.dtype if args.dtype == "auto" else getattr(torch, args.dtype)
# Checkpoints have no auto_map or tokenizer. Never load train_state.pt.
config = entry["config_class"].from_pretrained(plan["checkpoint"])
model = entry["model_class"].from_pretrained(plan["checkpoint"], config=config, dtype=dtype)
model.to(args.device).eval()
backend = "seq2seq" if plan["kind"] == "seq2seq" else "causal"
adapter = DiffusionHFLM if plan["kind"] == "diffusion" else HFLM
if plan['local_model'] == 'prefixlm':
adapter = PrefixHFLM
if plan["kind"] == "seq2seq" and getattr(config, "ul2", None):
adapter = UL2HFLM
return adapter(pretrained=model, tokenizer=tokenizer, backend=backend, **common)