BooruPromptRNG / app.py
LoliRimuru's picture
Upload 10 files
a972f88 verified
Raw History Blame Contribute Delete
6.31 kB
"""Gradio app for the Booru prompt generator."""
import json
from pathlib import Path
import gradio as gr
import torch
from safetensors.torch import load_model
from model import Vocab, SimpleGraph, TagTransformer, NeuralPromptGenerator
_RELEASE_DIR = Path(__file__).parent
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"] + 1,
).to(_DEVICE)
load_model(_MODEL, str(_RELEASE_DIR / "model.safetensors"))
def _parse_comma(text: str) -> list:
return [t.strip() for t in text.split(",") if t.strip()]
def generate(
mode: str,
alpha: float,
count: int,
length: int,
rating: str,
anchor: str,
blacklist: str,
min_prob: float,
temperature: float,
top_k: int,
top_p: float,
distribution_weight: float,
seed: int,
):
if mode == "Empirical":
alpha = 0.0
elif mode == "Diverse":
alpha = 1.0
gen = NeuralPromptGenerator(
_MODEL,
_GRAPH,
device=_DEVICE,
seed=None if seed < 0 else seed,
distribution_weight=distribution_weight,
)
prompts = gen.generate(
alpha=alpha,
count=count,
length=length,
rating=rating,
anchor=_parse_comma(anchor) or None,
blacklist=_parse_comma(blacklist) or None,
min_prob=min_prob,
temperature=temperature,
top_k=top_k,
top_p=top_p,
)
return "\n".join(", ".join(p) for p in prompts)
with gr.Blocks(title="Booru Smart Prompt Generator") as demo:
gr.Markdown("# Booru Smart Prompt Generator")
gr.Markdown(
"Generate semantically coherent Danbooru-style tag prompts. "
"Pick a generation mode and optional content rating, then tune sampling controls."
)
with gr.Row():
with gr.Column(scale=1):
gr.Markdown("### Prompt setup")
mode = gr.Dropdown(
choices=["Empirical", "Diverse", "Custom alpha"],
value="Empirical",
label="Mode",
info="Empirical = follow dataset frequencies (alpha=0). Diverse = boost rare tags (alpha=1). Custom alpha lets you set the slider below.",
)
alpha = gr.Slider(
0.0, 1.0, value=0.0, step=0.05,
label="Alpha (custom)",
info="0 = empirical distribution, 1 = maximum diversity. Only used when Mode is 'Custom alpha'.",
)
rating = gr.Dropdown(
choices=["g", "s", "q", "e"],
value="g",
label="Content rating",
info="g = general, s = sensitive, q = questionable, e = explicit. Conditions the model on the rating token.",
)
count = gr.Number(
value=5, precision=0,
label="Number of prompts",
minimum=1, maximum=100,
info="How many independent prompts to generate.",
)
length = gr.Slider(
1, 128, value=30, step=1,
label="Tags per prompt",
info="Target number of tags in each prompt (excluding anchor tags).",
)
anchor = gr.Textbox(
label="Anchor tags",
placeholder="e.g. 1girl, black_hair",
info="Comma-separated tags that must appear in every prompt. Unknown tags are ignored.",
)
blacklist = gr.Textbox(
label="Blacklist tags",
placeholder="e.g. 1boy, smile",
info="Comma-separated tags the model must not generate. Unknown tags are ignored.",
)
gr.Markdown("### Sampling controls")
min_prob = gr.Slider(
0.0, 0.1, value=0.0005, step=0.0001,
label="Min model probability",
info="Tags whose raw model probability is below this threshold are dropped. If no tag passes, a top-k fallback is used.",
)
temperature = gr.Slider(
0.1, 2.0, value=1.0, step=0.05,
label="Temperature",
info="Lower = more deterministic / focused. Higher = more random.",
)
top_k = gr.Slider(
0, 200, value=0, step=1,
label="Top-k",
info="Keep only the k most likely tags at each step. 0 disables top-k.",
)
top_p = gr.Slider(
0.0, 1.0, value=1.0, step=0.01,
label="Top-p (nucleus)",
info="Keep the smallest set of tags whose cumulative probability exceeds this value. 1.0 disables nucleus sampling.",
)
distribution_weight = gr.Slider(
0.0, 2.0, value=0.5, step=0.05,
label="Distribution bias weight",
info="How strongly the empirical/diverse bias is applied on top of the model's scores.",
)
seed = gr.Number(
value=-1, precision=0,
label="Seed (-1 for random)",
info="Set a non-negative seed for reproducible generation.",
)
generate_btn = gr.Button("Generate", variant="primary")
with gr.Column(scale=2):
output = gr.Textbox(label="Generated prompts", lines=24)
generate_btn.click(
fn=generate,
inputs=[
mode, alpha, count, length, rating, anchor, blacklist,
min_prob, temperature, top_k, top_p, distribution_weight, seed,
],
outputs=output,
)
if __name__ == "__main__":
demo.launch()