Spaces:
Running
Running
File size: 6,313 Bytes
c6dd74a a972f88 c6dd74a 1c0af3b c6dd74a a972f88 c6dd74a 1c0af3b c6dd74a a972f88 c6dd74a a972f88 c6dd74a a972f88 c6dd74a a972f88 c6dd74a a972f88 c6dd74a a972f88 c6dd74a a972f88 c6dd74a a972f88 c6dd74a 1c0af3b a972f88 1c0af3b c6dd74a | 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 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 | """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()
|