File size: 3,696 Bytes
8b27df3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 | """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()
|