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