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:
# 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
| """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):] | |
| 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.""" | |
| 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. | |
| """ | |
| 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") | |
| 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) | |