Text Generation
PyTorch
English
diffusion-language-modeling
SDLLM-AR-1.7B-Base / eval /lm_eval.py
jlemercier's picture
Release SDLLM inference package
8901f3e verified
Raw
History Blame Contribute Delete
6.22 kB
"""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()