Download sampling.py from nvidia/SDLLM-MDLM-1.7B-Base: direct link, hf CLI and curl.
- Browser
- Download file 5.64 kB
-
https://huggingface.co/nvidia/SDLLM-MDLM-1.7B-Base/resolve/main/sampling.py
- Command line
-
hf download hf://nvidia/SDLLM-MDLM-1.7B-Base/sampling.py
-
curl -L -o sampling.py https://huggingface.co/nvidia/SDLLM-MDLM-1.7B-Base/resolve/main/sampling.py
5.64 kB
| """Prompted and unconditional samplers for SDLLM release checkpoints.""" | |
| from __future__ import annotations | |
| import sys | |
| import torch | |
| from tqdm.auto import tqdm | |
| def _enable_compiled_attention(verbosity: str) -> None: | |
| """Use compiled FlexAttention, the release default for full canvases.""" | |
| import models.dit | |
| if models.dit.flex_attention_compiled is models.dit.flex_attention: | |
| models.dit.flex_attention_compiled = torch.compile( | |
| models.dit.flex_attention, dynamic=True) | |
| if verbosity == "full": | |
| print( | |
| "Using compiled FlexAttention. The first sampling run for a new canvas " | |
| "shape includes compilation warm-up.", | |
| file=sys.stderr, | |
| flush=True, | |
| ) | |
| def sample_autoregressive(model, prompt: torch.Tensor, num_samples: int, | |
| max_new_tokens: int, verbosity: str) -> torch.Tensor: | |
| """Gumbel-max (temperature 1) ancestral sampling conditioned on ``prompt``.""" | |
| if prompt.numel() == 0: | |
| prompt = torch.tensor([model.tokenizer.bos_token_id], device=model.device) | |
| prompt = prompt.to(model.device, dtype=torch.long) | |
| if prompt.numel() >= model.num_tokens: | |
| raise ValueError(f"Prompt has {prompt.numel()} tokens; limit is {model.num_tokens - 1}.") | |
| output_length = min(model.num_tokens, prompt.numel() + max_new_tokens) | |
| samples = prompt.repeat(num_samples, 1) | |
| sigma = torch.zeros(num_samples, dtype=model.dtype, device=model.device) | |
| model.backbone.reset_kv_cache() | |
| temperature = float(model.config.sampling.temperature) | |
| for _ in tqdm(range(prompt.numel(), output_length), desc="Sampling", | |
| disable=verbosity == "none"): | |
| logits = model.backbone(samples, sigma=sigma, x0=None, kv_cache=False)[:, -1] | |
| logits[:, model.mask_index] = model.neg_infinity | |
| if temperature == 0: | |
| token = logits.argmax(-1, keepdim=True) | |
| else: | |
| gumbel = torch.rand_like(logits).log().neg().log().neg() | |
| token = (logits / temperature + gumbel).argmax(-1, keepdim=True) | |
| samples = torch.cat((samples, token), dim=1) | |
| model.backbone.reset_kv_cache() | |
| return samples | |
| def sample_diffusion(model, prompt: torch.Tensor, num_samples: int, | |
| max_new_tokens: int, steps: int | None, | |
| verbosity: str) -> torch.Tensor: | |
| """Conditional ancestral diffusion sampling with the prompt clamped.""" | |
| prompt = prompt.to(model.device, dtype=torch.long) | |
| # The checkpoint-compatible legacy runtime owns its sampler in | |
| # ``samplers.py`` rather than implementing private update methods on the | |
| # model. Its native sampler already supports prompt conditioning. | |
| if str(model.config.algo.backbone).endswith("_legacy"): | |
| model.config.sampling.verbose_progress = verbosity != "none" | |
| if steps is not None: | |
| model.config.sampling.steps = steps | |
| condition = None if prompt.numel() == 0 else [prompt.unsqueeze(0)] * num_samples | |
| if verbosity == "full": | |
| print(f"Sampling a {model.num_tokens}-token continuation with " | |
| f"{model.config.sampling.predictor} for {model.config.sampling.steps} " | |
| f"reverse steps; returning its first {max_new_tokens} tokens", flush=True) | |
| print("Native legacy sampler does not expose per-step callbacks; waiting for sampling to finish...", flush=True) | |
| samples = model.generate_samples(num_samples=num_samples, condition=condition) | |
| # Legacy greedy-tail implementations may argmax every canvas position | |
| # in their final step. Preserve the user-provided condition exactly. | |
| if prompt.numel(): | |
| samples[:, :prompt.numel()] = prompt | |
| return samples | |
| if prompt.numel() + max_new_tokens > model.num_tokens: | |
| max_new_tokens = model.num_tokens - prompt.numel() | |
| if max_new_tokens <= 0: | |
| return prompt.repeat(num_samples, 1) | |
| steps = model.config.sampling.steps if steps is None else steps | |
| x = model.prior_sample(num_samples, prompt.numel() + max_new_tokens) | |
| if prompt.numel(): | |
| x[:, :prompt.numel()] = prompt | |
| timesteps = torch.linspace(1, 1e-5, steps + 1, device=model.device) | |
| dt = (1 - 1e-5) / steps | |
| cache = None | |
| for t in tqdm(timesteps[:-1], desc="Sampling", disable=verbosity == "none"): | |
| time = t.expand(num_samples, 1) | |
| if model.sampler == "ancestral_cache": | |
| cache, x_next = model._ancestral_update(x, time, dt, cache, False) | |
| cache = cache if torch.equal(x_next, x) and not model.time_conditioning else None | |
| x = x_next | |
| elif model.sampler == "ancestral": | |
| _, x = model._ancestral_update(x, time, dt, None, False) | |
| else: | |
| x = model._analytic_update(x, time, dt) | |
| if prompt.numel(): | |
| x[:, :prompt.numel()] = prompt | |
| if model.config.sampling.noise_removal == "ancestral": | |
| time = timesteps[-1].expand(num_samples, 1) | |
| _, x = model._ancestral_update(x, time, None, cache, noise_removal_step=True) | |
| if prompt.numel(): | |
| x[:, :prompt.numel()] = prompt | |
| return x | |
| def sample(model, prompt: torch.Tensor, num_samples: int, max_new_tokens: int, | |
| steps: int | None = None, verbosity: str = "minimal") -> torch.Tensor: | |
| _enable_compiled_attention(verbosity) | |
| if model.config.algo.name == "ar": | |
| return sample_autoregressive(model, prompt, num_samples, max_new_tokens, verbosity) | |
| return sample_diffusion(model, prompt, num_samples, max_new_tokens, steps, verbosity) | |