Spaces:
Running
Running
Download app.py from LoliRimuru/BooruPromptRNG: direct link, hf CLI and curl.
- Browser
- Download file 6.31 kB
-
https://huggingface.co/spaces/LoliRimuru/BooruPromptRNG/resolve/main/app.py
- Command line
-
hf download hf://spaces/LoliRimuru/BooruPromptRNG/app.py
-
curl -L -o app.py https://huggingface.co/spaces/LoliRimuru/BooruPromptRNG/resolve/main/app.py
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() | |