Download generate_tiny_digits.py from shibatch/tinydigitdiffusion3m: direct link, hf CLI and curl.
- Browser
- Download file 2.39 kB
-
https://huggingface.co/shibatch/tinydigitdiffusion3m/resolve/main/generate_tiny_digits.py
- Command line
-
hf download hf://shibatch/tinydigitdiffusion3m/generate_tiny_digits.py
-
curl -L -o generate_tiny_digits.py https://huggingface.co/shibatch/tinydigitdiffusion3m/resolve/main/generate_tiny_digits.py
2.39 kB
| #!/usr/bin/env python3 | |
| """Generate a handwritten 1-8 digit image from a numeric prompt.""" | |
| from __future__ import annotations | |
| import argparse | |
| from pathlib import Path | |
| import torch | |
| from PIL import Image | |
| from tiny_digit_diffusion import ddim_sample, load_model | |
| DEFAULT_REPO_ID = "shibatch/tinydigitdiffusion3m" | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("prompt", help="A numeric prompt containing 1-8 digits, for example 2026.") | |
| parser.add_argument( | |
| "--model-dir", | |
| default=None, | |
| help="Local model directory. If omitted, download the model from Hugging Face.", | |
| ) | |
| parser.add_argument("--repo-id", default=DEFAULT_REPO_ID) | |
| parser.add_argument("--revision", default="main") | |
| parser.add_argument("--output", default="generated_digits.png") | |
| parser.add_argument("--steps", type=int, default=50) | |
| parser.add_argument("--guidance-scale", type=float, default=1.0) | |
| parser.add_argument("--seed", type=int, default=0) | |
| parser.add_argument("--device", choices=["auto", "cpu", "cuda"], default="auto") | |
| return parser.parse_args() | |
| def resolve_model_dir( | |
| model_dir: str | None, | |
| repo_id: str, | |
| revision: str, | |
| ) -> Path: | |
| if model_dir is not None: | |
| return Path(model_dir).expanduser().resolve() | |
| from huggingface_hub import snapshot_download | |
| snapshot_dir = snapshot_download( | |
| repo_id=repo_id, | |
| revision=revision, | |
| allow_patterns=["model/config.json", "model/model.safetensors"], | |
| ) | |
| return Path(snapshot_dir) / "model" | |
| def main() -> None: | |
| args = parse_args() | |
| device = torch.device( | |
| "cuda" if args.device == "auto" and torch.cuda.is_available() else | |
| "cpu" if args.device == "auto" else args.device | |
| ) | |
| model_dir = resolve_model_dir(args.model_dir, args.repo_id, args.revision) | |
| model = load_model(model_dir, device) | |
| image = ddim_sample( | |
| model, | |
| [args.prompt], | |
| device, | |
| sampling_steps=args.steps, | |
| guidance_scale=args.guidance_scale, | |
| seed=args.seed, | |
| )[0, 0] | |
| pixels = ((image.cpu() + 1) * 127.5).round().clamp(0, 255).byte().numpy() | |
| output = Path(args.output) | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| Image.fromarray(pixels, mode="L").save(output) | |
| print(output.resolve()) | |
| if __name__ == "__main__": | |
| main() | |