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