Text Generation
PyTorch
English
diffusion-language-modeling
SDLLM-Duo-1.7B-Base / sample.py
jlemercier's picture
Release SDLLM inference package
a5d1368 verified
Raw
History Blame Contribute Delete
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()