Download inference.py from LoliRimuru/BooruPromptGenerator: direct link, hf CLI and curl.
- Browser
- Download file 3.7 kB
-
https://huggingface.co/LoliRimuru/BooruPromptGenerator/resolve/8b27df3f7c6d96d6e29b039d92351e7d227c2ad5/inference.py
- Command line
-
hf download hf://LoliRimuru/BooruPromptGenerator@8b27df3f7c6d96d6e29b039d92351e7d227c2ad5/inference.py
-
curl -L -o inference.py https://huggingface.co/LoliRimuru/BooruPromptGenerator/resolve/8b27df3f7c6d96d6e29b039d92351e7d227c2ad5/inference.py
3.7 kB
| """Standalone inference script for the Booru prompt generator release.""" | |
| import argparse | |
| import json | |
| import sys | |
| from pathlib import Path | |
| import torch | |
| from safetensors.torch import load_model | |
| from model import Vocab, SimpleGraph, TagTransformer, NeuralPromptGenerator | |
| def load_generator(release_dir: str, seed: Optional[int] = None): | |
| release_dir = Path(release_dir) | |
| with open(release_dir / "config.json", "r", encoding="utf-8") as f: | |
| config = json.load(f) | |
| with open(release_dir / "vocab.json", "r", encoding="utf-8") as f: | |
| tags = json.load(f) | |
| with open(release_dir / "counts.json", "r", encoding="utf-8") as f: | |
| counts = json.load(f) | |
| with open(release_dir / "mutex.json", "r", encoding="utf-8") as f: | |
| mutex = json.load(f) | |
| vocab = Vocab(tags, counts) | |
| graph = SimpleGraph(vocab, mutex) | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| model = TagTransformer( | |
| vocab_size=len(vocab), | |
| d_model=config["d_model"], | |
| nhead=config["nhead"], | |
| num_layers=config["num_layers"], | |
| dim_feedforward=config["dim_feedforward"], | |
| dropout=config["dropout"], | |
| max_len=config["max_len"] + 2, | |
| ).to(device) | |
| load_model(model, str(release_dir / "model.safetensors")) | |
| return NeuralPromptGenerator( | |
| model, | |
| graph, | |
| device=device, | |
| seed=seed, | |
| distribution_weight=config.get("distribution_weight", 0.75), | |
| ) | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Generate Booru tag prompts.") | |
| parser.add_argument("--release-dir", default=".", help="Path to the release folder") | |
| parser.add_argument("--mode", choices=["empirical", "diverse"], default=None) | |
| parser.add_argument("--alpha", type=float, default=None) | |
| parser.add_argument("--count", type=int, default=10) | |
| parser.add_argument("--length", type=int, default=30) | |
| parser.add_argument("--anchor", default="", help="Comma-separated anchor tags") | |
| parser.add_argument("--blacklist", default="", help="Comma-separated tags to forbid") | |
| parser.add_argument("--rating", default="g", choices=["g", "s", "q", "e"], | |
| help="Content rating token to condition on") | |
| parser.add_argument("--min-prob", type=float, default=0.0005) | |
| parser.add_argument("--temperature", type=float, default=1.0) | |
| parser.add_argument("--top-k", type=int, default=0, help="Top-k sampling (0 = disabled)") | |
| parser.add_argument("--top-p", type=float, default=1.0, help="Nucleus/top-p sampling (1.0 = disabled)") | |
| parser.add_argument("--distribution-weight", type=float, default=None) | |
| parser.add_argument("--seed", type=int, default=None) | |
| args = parser.parse_args() | |
| alpha = args.alpha | |
| if args.mode == "empirical": | |
| alpha = 0.0 | |
| elif args.mode == "diverse": | |
| alpha = 1.0 | |
| if alpha is None: | |
| alpha = 0.0 | |
| gen = load_generator(args.release_dir, seed=args.seed) | |
| if args.distribution_weight is not None: | |
| gen.distribution_weight = args.distribution_weight | |
| anchor_tags = [t.strip() for t in args.anchor.split(",") if t.strip()] | |
| blacklist_tags = [t.strip() for t in args.blacklist.split(",") if t.strip()] | |
| prompts = gen.generate( | |
| alpha=alpha, | |
| count=args.count, | |
| length=args.length, | |
| anchor=anchor_tags or None, | |
| blacklist=blacklist_tags or None, | |
| min_prob=args.min_prob, | |
| temperature=args.temperature, | |
| top_k=getattr(args, "top_k", 0), | |
| top_p=getattr(args, "top_p", 1.0), | |
| rating=args.rating, | |
| ) | |
| for tags in prompts: | |
| print(", ".join(tags)) | |
| if __name__ == "__main__": | |
| main() | |