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