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