| |
| """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) |
| |
| |
| try: |
| prompt = tokenizer.encode(args.prompt, device=model.device, eos=False) |
| except TypeError: |
| |
| |
| 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() |
|
|