BooruPromptGenerator / inference.py
LoliRimuru's picture
Upload 10 files
8b27df3 verified
Raw
History Blame Contribute Delete
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()