Download sample.py from nvidia/SDLLM-MDLM-1.7B-Base: direct link, hf CLI and curl.
- Browser
- Download file 3.41 kB
-
https://huggingface.co/nvidia/SDLLM-MDLM-1.7B-Base/resolve/main/sample.py
- Command line
-
hf download hf://nvidia/SDLLM-MDLM-1.7B-Base/sample.py
-
curl -L -o sample.py https://huggingface.co/nvidia/SDLLM-MDLM-1.7B-Base/resolve/main/sample.py
3.41 kB
| #!/usr/bin/env python3 | |
| """Sample an SDLLM release locally or directly from the Hugging Face Hub.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import torch | |
| from sdllm import load_model, sampling_defaults | |
| from sampling import sample | |
| def main() -> None: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("model", help="Local release directory or Hugging Face repo id") | |
| parser.add_argument("--prompt", default="", help="Optional text prefix") | |
| parser.add_argument("--num-samples", type=int, default=1) | |
| parser.add_argument("--max-new-tokens", type=int, | |
| help="Limit returned tokens; defaults to the full model canvas") | |
| parser.add_argument("--steps", type=int, help="Diffusion steps; defaults to the release config") | |
| parser.add_argument("--top-p", type=float, help="Override nucleus sampling probability") | |
| parser.add_argument("--temperature", "--token-temperature", dest="temperature", type=float, | |
| help="Token-sampling temperature (default: 1.0)") | |
| parser.add_argument("--noise-removal", choices=["none", "ancestral", "greedy"]) | |
| parser.add_argument("--device", default="cuda") | |
| parser.add_argument("--seed", type=int) | |
| parser.add_argument( | |
| "--verbose", nargs="?", const="full", choices=("none", "minimal", "full"), | |
| default="minimal", | |
| help="Diagnostic level: none, minimal (default), or full; --verbose alone means full", | |
| ) | |
| parser.add_argument("--show-defaults", action="store_true") | |
| args = parser.parse_args() | |
| if args.seed is not None: | |
| torch.manual_seed(args.seed) | |
| model, tokenizer, config = load_model(args.model, args.device, verbose=args.verbose == "full") | |
| if args.show_defaults: | |
| print(json.dumps(sampling_defaults(config), indent=2)) | |
| return | |
| if args.top_p is not None: | |
| config.sampling.p_nucleus = args.top_p | |
| if args.temperature is not None: | |
| if args.temperature < 0: | |
| parser.error("--temperature must be non-negative") | |
| config.sampling.temperature = args.temperature | |
| if args.noise_removal is not None: | |
| config.sampling.noise_removal = args.noise_removal | |
| if args.max_new_tokens is None: | |
| args.max_new_tokens = model.num_tokens | |
| if args.verbose == "full": | |
| print("Effective sampling parameters: " + json.dumps(sampling_defaults(config)), flush=True) | |
| # ``eos=False`` is essential: the prefix must not be terminated before | |
| # generation starts. | |
| try: | |
| prompt = tokenizer.encode(args.prompt, device=model.device, eos=False) | |
| except TypeError: | |
| # The legacy tiktoken tokenizer has a smaller encode API and does not | |
| # append EOS, which is precisely what prompted generation needs. | |
| prompt = tokenizer.encode(args.prompt).to(model.device) | |
| if args.verbose != "none": | |
| canvas_length = args.max_new_tokens if config.algo.name == "ar" else model.num_tokens | |
| print(f"Prompt: {prompt.numel()} tokens; canvas: {canvas_length} tokens; " | |
| f"returned output: {args.max_new_tokens} tokens", flush=True) | |
| outputs = sample(model, prompt, args.num_samples, args.max_new_tokens, args.steps, args.verbose) | |
| prefix_len = prompt.numel() | |
| for output in outputs: | |
| print(tokenizer.decode(output[prefix_len:prefix_len + args.max_new_tokens])) | |
| if __name__ == "__main__": | |
| main() | |