Other
Transformers
TensorBoard
Safetensors
English
diffusion_lm
fill-mask
custom_code
tiny-llm-ablation
from-scratch
diffusion
masked-language-modeling
Eval Results (legacy)
Instructions to use d0rj/diffusion-51M-base with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use d0rj/diffusion-51M-base with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModelForMaskedLM model = AutoModelForMaskedLM.from_pretrained("d0rj/diffusion-51M-base", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
File size: 11,608 Bytes
80aea5b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 | """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)
|