| """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() |
|
|