"""Conditional likelihood evaluation for AR, MDLM, EsoLM and Duo.""" from __future__ import annotations import math import torch import torch.nn.functional as F from lm_eval.api.model import LM from lm_eval.api.registry import register_model from sampling import sample from sdllm import load_model @register_model("dLLM") class SDLLMEvalHarness(LM): """The release-paper Monte Carlo conditional likelihood estimator.""" def __init__(self, model_path: str, device: str = "cuda", batch_size: int | None = None, likelihood_batch_size: int | None = None, likelihood_mc_num: int = 32, max_gen_toks: int = 256, **_: object): super().__init__() self.model, self.tokenizer, self.config = load_model(model_path, device) self._device = torch.device(device) # This controls only the Monte Carlo replicas for diffusion # likelihoods. Keep it independent of lm-eval's request batch size. self.batch_size = int(likelihood_batch_size or 8) self.likelihood_mc_num = int(likelihood_mc_num) self.max_gen_toks = int(max_gen_toks) @property def rank(self): return 0 @property def world_size(self): return 1 def _encode_pair(self, context: str, continuation: str): spaces = len(context) - len(context.rstrip()) if spaces: continuation, context = context[-spaces:] + continuation, context[:-spaces] whole = self.tokenizer.encode(context + continuation) prefix = self.tokenizer.encode(context) return prefix, whole[len(prefix):] def _perturb(self, sequence, prefix_length): t = self.model._sample_t(sequence.shape[0], None) _, alpha = self.model.noise(t) alpha = alpha.unsqueeze(-1) sigma = self.model._sigma_from_alphat(alpha) noisy = self.model.q_xt(sequence, alpha) noisy[:, :prefix_length] = sequence[:, :prefix_length] return noisy, 1 - alpha, sigma @torch.no_grad() def _mdlm_ll(self, prefix, target): sequence = torch.cat((prefix, target))[None].repeat(self.batch_size, 1).to(self._device) values = [] for _ in range(max(1, math.ceil(self.likelihood_mc_num / self.batch_size))): noisy, probability, sigma = self._perturb(sequence, len(prefix)) mask = noisy == self.model.mask_index logits = self.model(noisy, sigma) loss = F.cross_entropy(logits[mask], sequence[mask], reduction="none") / probability.expand_as(sequence)[mask] values.append(-loss.sum().item() / self.batch_size) return sum(values) / len(values) @torch.no_grad() def _esolm_ll(self, prefix, target): sequence = torch.cat((prefix, target))[None].repeat(self.batch_size, 1).to(self._device) values = [] for _ in range(max(1, math.ceil(self.likelihood_mc_num / self.batch_size))): noisy, probability, sigma = self._perturb(sequence, len(prefix)) order = self.model._sort_indices(noisy[:, len(prefix):], shuffle=self.config.algo.diffusion_shuffle) fixed = torch.arange(len(prefix), device=self._device).repeat(self.batch_size, 1) order = torch.cat((fixed, len(prefix) + order), dim=1) noisy, clean = torch.gather(noisy, 1, order), torch.gather(sequence, 1, order) mask = noisy == self.model.mask_index logits = self.model(noisy, sigma, sort_idx=order) loss = F.cross_entropy(logits[mask], clean[mask], reduction="none") / probability.expand_as(noisy)[mask] values.append(-loss.sum().item() / self.batch_size) return sum(values) / len(values) @torch.no_grad() def _duo_ll(self, prefix, target): sequence = torch.cat((prefix, target))[None].repeat(self.batch_size, 1).to(self._device) values = [] for _ in range(max(1, math.ceil(self.likelihood_mc_num / self.batch_size))): noisy, probability, sigma = self._perturb(sequence, len(prefix)) logits = self.model(noisy, sigma) loss = self.model.nll_per_token(logits, noisy, sequence, 1 - probability, -1, low_var=False) values.append(-loss[:, len(prefix):].sum().item() / self.batch_size) return sum(values) / len(values) @torch.no_grad() def _ar_ll(self, prefix, target): sequence = torch.cat((prefix, target))[None].to(self._device) if sequence.shape[1] < 2: return 0.0 logits = self.model.backbone(sequence[:, :-1], torch.zeros(1, device=self._device)) logits[:, :, self.model.mask_index] = self.model.neg_infinity loss = F.cross_entropy(logits[0], sequence[0, 1:], reduction="none") return -loss[max(len(prefix) - 1, 0):].sum().item() def loglikelihood(self, requests): result = [] for request in requests: prefix, target = self._encode_pair(request.args[0], request.args[1]) if len(prefix) + len(target) > self.model.num_tokens: raise ValueError("Example exceeds context length; rolling evaluation is not enabled.") family = self.config.algo.name method = {"ar": self._ar_ll, "mdlm": self._mdlm_ll, "esolm": self._esolm_ll, "duo_base": self._duo_ll}[family] result.append((method(prefix, target), False)) return result def loglikelihood_rolling(self, requests): raise NotImplementedError @torch.no_grad() def generate_until(self, requests): result = [] for request in requests: context, options = request.args try: prefix = self.tokenizer.encode(context, device=self._device, eos=False) except TypeError: prefix = self.tokenizer.encode(context).to(self._device) text = self.tokenizer.decode( sample(self.model, prefix, 1, self.max_gen_toks, verbosity="none")[0, len(prefix):] ) for stop in options.get("until", []): text = text.split(stop, 1)[0] result.append(text) return result if __name__ == "__main__": from lm_eval.__main__ import cli_evaluate cli_evaluate()