Spaces:
Running
Running
Upload 10 files
Browse files- README.md +39 -14
- app.py +124 -0
- config.json +10 -0
- counts.json +0 -0
- inference.py +93 -0
- model.py +251 -0
- model.safetensors +3 -0
- mutex.json +0 -0
- requirements.txt +4 -0
- vocab.json +0 -0
README.md
CHANGED
|
@@ -1,14 +1,39 @@
|
|
| 1 |
-
|
| 2 |
-
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Booru Smart Prompt Generator — Gradio UI
|
| 2 |
+
|
| 3 |
+
A web UI for the Booru tag prompt generator.
|
| 4 |
+
|
| 5 |
+
## Files
|
| 6 |
+
|
| 7 |
+
- `app.py` — Gradio application
|
| 8 |
+
- `model.py`, `model.safetensors`, `config.json`, `vocab.json`, `counts.json`, `mutex.json` — model assets
|
| 9 |
+
- `requirements.txt` — Python dependencies
|
| 10 |
+
|
| 11 |
+
## Usage on Hugging Face Spaces
|
| 12 |
+
|
| 13 |
+
1. Create a new **Gradio** Space.
|
| 14 |
+
2. Upload the contents of this folder (including `model.safetensors`).
|
| 15 |
+
3. Wait for the container to build and start.
|
| 16 |
+
4. Open the app and generate prompts.
|
| 17 |
+
|
| 18 |
+
## Parameters
|
| 19 |
+
|
| 20 |
+
- **Mode**: `Empirical` (alpha=0), `Diverse` (alpha=1), or `Custom alpha`.
|
| 21 |
+
- **Alpha**: distribution coefficient (active only in Custom mode).
|
| 22 |
+
- **Tags per prompt**: 1–128; 30 is recommended.
|
| 23 |
+
- **Anchor tags**: fixed starting tags the generator must include.
|
| 24 |
+
- **Blacklist tags**: tags the generator must never emit.
|
| 25 |
+
- **Min model probability**: rejects low-probability / conflicting candidates.
|
| 26 |
+
- **Temperature**: higher = more random.
|
| 27 |
+
- **Distribution bias weight**: strength of the empirical ↔ diverse reweighting.
|
| 28 |
+
|
| 29 |
+
## Local run
|
| 30 |
+
|
| 31 |
+
```bash
|
| 32 |
+
pip install -r requirements.txt
|
| 33 |
+
python app.py
|
| 34 |
+
```
|
| 35 |
+
|
| 36 |
+
> For an RTX 5090 / CUDA 13.x setup, install the matching PyTorch wheel before running, e.g.:
|
| 37 |
+
> ```bash
|
| 38 |
+
> pip install torch==2.12.1+cu130 --index-url https://download.pytorch.org/whl/cu130
|
| 39 |
+
> ```
|
app.py
ADDED
|
@@ -0,0 +1,124 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Gradio app for the Booru prompt generator."""
|
| 2 |
+
import json
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
|
| 5 |
+
import gradio as gr
|
| 6 |
+
import torch
|
| 7 |
+
from safetensors.torch import load_model
|
| 8 |
+
|
| 9 |
+
from model import Vocab, SimpleGraph, TagTransformer, NeuralPromptGenerator
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
_RELEASE_DIR = Path(__file__).parent
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
with open(_RELEASE_DIR / "config.json", "r", encoding="utf-8") as f:
|
| 16 |
+
_CONFIG = json.load(f)
|
| 17 |
+
with open(_RELEASE_DIR / "vocab.json", "r", encoding="utf-8") as f:
|
| 18 |
+
_TAGS = json.load(f)
|
| 19 |
+
with open(_RELEASE_DIR / "counts.json", "r", encoding="utf-8") as f:
|
| 20 |
+
_COUNTS = json.load(f)
|
| 21 |
+
with open(_RELEASE_DIR / "mutex.json", "r", encoding="utf-8") as f:
|
| 22 |
+
_MUTEX = json.load(f)
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
_VOCAB = Vocab(_TAGS, _COUNTS)
|
| 26 |
+
_GRAPH = SimpleGraph(_VOCAB, _MUTEX)
|
| 27 |
+
_DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 28 |
+
|
| 29 |
+
_MODEL = TagTransformer(
|
| 30 |
+
vocab_size=len(_VOCAB),
|
| 31 |
+
d_model=_CONFIG["d_model"],
|
| 32 |
+
nhead=_CONFIG["nhead"],
|
| 33 |
+
num_layers=_CONFIG["num_layers"],
|
| 34 |
+
dim_feedforward=_CONFIG["dim_feedforward"],
|
| 35 |
+
dropout=_CONFIG["dropout"],
|
| 36 |
+
max_len=_CONFIG["max_len"] + 1,
|
| 37 |
+
).to(_DEVICE)
|
| 38 |
+
load_model(_MODEL, str(_RELEASE_DIR / "model.safetensors"))
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def _parse_comma(text: str) -> list:
|
| 42 |
+
return [t.strip() for t in text.split(",") if t.strip()]
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def generate(
|
| 46 |
+
mode: str,
|
| 47 |
+
alpha: float,
|
| 48 |
+
count: int,
|
| 49 |
+
length: int,
|
| 50 |
+
anchor: str,
|
| 51 |
+
blacklist: str,
|
| 52 |
+
min_prob: float,
|
| 53 |
+
temperature: float,
|
| 54 |
+
distribution_weight: float,
|
| 55 |
+
seed: int,
|
| 56 |
+
):
|
| 57 |
+
if mode == "Empirical":
|
| 58 |
+
alpha = 0.0
|
| 59 |
+
elif mode == "Diverse":
|
| 60 |
+
alpha = 1.0
|
| 61 |
+
|
| 62 |
+
gen = NeuralPromptGenerator(
|
| 63 |
+
_MODEL,
|
| 64 |
+
_GRAPH,
|
| 65 |
+
device=_DEVICE,
|
| 66 |
+
seed=None if seed < 0 else seed,
|
| 67 |
+
distribution_weight=distribution_weight,
|
| 68 |
+
)
|
| 69 |
+
|
| 70 |
+
prompts = gen.generate(
|
| 71 |
+
alpha=alpha,
|
| 72 |
+
count=count,
|
| 73 |
+
length=length,
|
| 74 |
+
anchor=_parse_comma(anchor) or None,
|
| 75 |
+
blacklist=_parse_comma(blacklist) or None,
|
| 76 |
+
min_prob=min_prob,
|
| 77 |
+
temperature=temperature,
|
| 78 |
+
)
|
| 79 |
+
return "\n".join(", ".join(p) for p in prompts)
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
with gr.Blocks(title="Booru Smart Prompt Generator") as demo:
|
| 83 |
+
gr.Markdown("# Booru Smart Prompt Generator")
|
| 84 |
+
gr.Markdown(
|
| 85 |
+
"Generate semantically coherent Danbooru-style tag prompts. "
|
| 86 |
+
"Empirical mode follows the dataset distribution; Diverse mode boosts rare tags."
|
| 87 |
+
)
|
| 88 |
+
|
| 89 |
+
with gr.Row():
|
| 90 |
+
with gr.Column(scale=1):
|
| 91 |
+
mode = gr.Dropdown(
|
| 92 |
+
choices=["Empirical", "Diverse", "Custom alpha"],
|
| 93 |
+
value="Empirical",
|
| 94 |
+
label="Mode",
|
| 95 |
+
)
|
| 96 |
+
alpha = gr.Slider(0.0, 1.0, value=0.0, step=0.05, label="Alpha (custom)")
|
| 97 |
+
count = gr.Number(value=5, precision=0, label="Number of prompts", minimum=1, maximum=100)
|
| 98 |
+
length = gr.Slider(1, 128, value=30, step=1, label="Tags per prompt")
|
| 99 |
+
anchor = gr.Textbox(
|
| 100 |
+
label="Anchor tags (comma-separated)",
|
| 101 |
+
placeholder="e.g. 1girl, black_hair",
|
| 102 |
+
)
|
| 103 |
+
blacklist = gr.Textbox(
|
| 104 |
+
label="Blacklist tags (comma-separated)",
|
| 105 |
+
placeholder="e.g. 1boy, smile",
|
| 106 |
+
)
|
| 107 |
+
min_prob = gr.Slider(0.0, 0.1, value=0.0, step=0.001, label="Min model probability threshold")
|
| 108 |
+
temperature = gr.Slider(0.1, 2.0, value=1.0, step=0.05, label="Temperature")
|
| 109 |
+
distribution_weight = gr.Slider(0.0, 2.0, value=0.75, step=0.05, label="Distribution bias weight")
|
| 110 |
+
seed = gr.Number(value=-1, precision=0, label="Seed (-1 for random)")
|
| 111 |
+
generate_btn = gr.Button("Generate", variant="primary")
|
| 112 |
+
|
| 113 |
+
with gr.Column(scale=2):
|
| 114 |
+
output = gr.Textbox(label="Generated prompts", lines=20)
|
| 115 |
+
|
| 116 |
+
generate_btn.click(
|
| 117 |
+
fn=generate,
|
| 118 |
+
inputs=[mode, alpha, count, length, anchor, blacklist, min_prob, temperature, distribution_weight, seed],
|
| 119 |
+
outputs=output,
|
| 120 |
+
)
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
if __name__ == "__main__":
|
| 124 |
+
demo.launch()
|
config.json
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"vocab_size": 25871,
|
| 3 |
+
"d_model": 256,
|
| 4 |
+
"nhead": 4,
|
| 5 |
+
"num_layers": 4,
|
| 6 |
+
"dim_feedforward": 1024,
|
| 7 |
+
"dropout": 0.1,
|
| 8 |
+
"max_len": 128,
|
| 9 |
+
"distribution_weight": 0.75
|
| 10 |
+
}
|
counts.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
inference.py
ADDED
|
@@ -0,0 +1,93 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Standalone inference script for the Booru prompt generator release."""
|
| 2 |
+
import argparse
|
| 3 |
+
import json
|
| 4 |
+
import sys
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
from safetensors.torch import load_model
|
| 9 |
+
|
| 10 |
+
from model import Vocab, SimpleGraph, TagTransformer, NeuralPromptGenerator
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def load_generator(release_dir: str, seed: Optional[int] = None):
|
| 14 |
+
release_dir = Path(release_dir)
|
| 15 |
+
with open(release_dir / "config.json", "r", encoding="utf-8") as f:
|
| 16 |
+
config = json.load(f)
|
| 17 |
+
with open(release_dir / "vocab.json", "r", encoding="utf-8") as f:
|
| 18 |
+
tags = json.load(f)
|
| 19 |
+
with open(release_dir / "counts.json", "r", encoding="utf-8") as f:
|
| 20 |
+
counts = json.load(f)
|
| 21 |
+
with open(release_dir / "mutex.json", "r", encoding="utf-8") as f:
|
| 22 |
+
mutex = json.load(f)
|
| 23 |
+
|
| 24 |
+
vocab = Vocab(tags, counts)
|
| 25 |
+
graph = SimpleGraph(vocab, mutex)
|
| 26 |
+
|
| 27 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 28 |
+
model = TagTransformer(
|
| 29 |
+
vocab_size=len(vocab),
|
| 30 |
+
d_model=config["d_model"],
|
| 31 |
+
nhead=config["nhead"],
|
| 32 |
+
num_layers=config["num_layers"],
|
| 33 |
+
dim_feedforward=config["dim_feedforward"],
|
| 34 |
+
dropout=config["dropout"],
|
| 35 |
+
max_len=config["max_len"] + 1,
|
| 36 |
+
).to(device)
|
| 37 |
+
load_model(model, str(release_dir / "model.safetensors"))
|
| 38 |
+
|
| 39 |
+
return NeuralPromptGenerator(
|
| 40 |
+
model,
|
| 41 |
+
graph,
|
| 42 |
+
device=device,
|
| 43 |
+
seed=seed,
|
| 44 |
+
distribution_weight=config.get("distribution_weight", 0.75),
|
| 45 |
+
)
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def main():
|
| 49 |
+
parser = argparse.ArgumentParser(description="Generate Booru tag prompts.")
|
| 50 |
+
parser.add_argument("--release-dir", default=".", help="Path to the release folder")
|
| 51 |
+
parser.add_argument("--mode", choices=["empirical", "diverse"], default=None)
|
| 52 |
+
parser.add_argument("--alpha", type=float, default=None)
|
| 53 |
+
parser.add_argument("--count", type=int, default=10)
|
| 54 |
+
parser.add_argument("--length", type=int, default=30)
|
| 55 |
+
parser.add_argument("--anchor", default="", help="Comma-separated anchor tags")
|
| 56 |
+
parser.add_argument("--blacklist", default="", help="Comma-separated tags to forbid")
|
| 57 |
+
parser.add_argument("--min-prob", type=float, default=0.0)
|
| 58 |
+
parser.add_argument("--temperature", type=float, default=1.0)
|
| 59 |
+
parser.add_argument("--distribution-weight", type=float, default=None)
|
| 60 |
+
parser.add_argument("--seed", type=int, default=None)
|
| 61 |
+
args = parser.parse_args()
|
| 62 |
+
|
| 63 |
+
alpha = args.alpha
|
| 64 |
+
if args.mode == "empirical":
|
| 65 |
+
alpha = 0.0
|
| 66 |
+
elif args.mode == "diverse":
|
| 67 |
+
alpha = 1.0
|
| 68 |
+
if alpha is None:
|
| 69 |
+
alpha = 0.0
|
| 70 |
+
|
| 71 |
+
gen = load_generator(args.release_dir, seed=args.seed)
|
| 72 |
+
if args.distribution_weight is not None:
|
| 73 |
+
gen.distribution_weight = args.distribution_weight
|
| 74 |
+
|
| 75 |
+
anchor_tags = [t.strip() for t in args.anchor.split(",") if t.strip()]
|
| 76 |
+
blacklist_tags = [t.strip() for t in args.blacklist.split(",") if t.strip()]
|
| 77 |
+
|
| 78 |
+
prompts = gen.generate(
|
| 79 |
+
alpha=alpha,
|
| 80 |
+
count=args.count,
|
| 81 |
+
length=args.length,
|
| 82 |
+
anchor=anchor_tags or None,
|
| 83 |
+
blacklist=blacklist_tags or None,
|
| 84 |
+
min_prob=args.min_prob,
|
| 85 |
+
temperature=args.temperature,
|
| 86 |
+
)
|
| 87 |
+
|
| 88 |
+
for tags in prompts:
|
| 89 |
+
print(", ".join(tags))
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
if __name__ == "__main__":
|
| 93 |
+
main()
|
model.py
ADDED
|
@@ -0,0 +1,251 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Minimal self-contained inference model for the Booru prompt generator release."""
|
| 2 |
+
import math
|
| 3 |
+
import re
|
| 4 |
+
from typing import Dict, List, Optional, Set
|
| 5 |
+
|
| 6 |
+
import numpy as np
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn as nn
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
PAD_ID = 0
|
| 12 |
+
BOS_ID = 1
|
| 13 |
+
EOS_ID = 2
|
| 14 |
+
OFFSET = 3
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class Vocab:
|
| 18 |
+
def __init__(self, tags: List[str], counts: List[int]):
|
| 19 |
+
self.tags = tags
|
| 20 |
+
self.tag_to_idx = {tag: i for i, tag in enumerate(tags)}
|
| 21 |
+
self.counts = np.array(counts, dtype=np.int64)
|
| 22 |
+
self.total = int(self.counts.sum())
|
| 23 |
+
self.freqs = self.counts.astype(np.float64) / max(self.total, 1)
|
| 24 |
+
|
| 25 |
+
def __len__(self) -> int:
|
| 26 |
+
return len(self.tags)
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class SimpleGraph:
|
| 30 |
+
def __init__(self, vocab: Vocab, mutex: List[List[int]]):
|
| 31 |
+
self.vocab = vocab
|
| 32 |
+
self.mutex = mutex
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
class SinusoidalPositionalEncoding(nn.Module):
|
| 36 |
+
def __init__(self, d_model: int, max_len: int = 512):
|
| 37 |
+
super().__init__()
|
| 38 |
+
pe = torch.zeros(max_len, d_model)
|
| 39 |
+
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
|
| 40 |
+
div_term = torch.exp(
|
| 41 |
+
torch.arange(0, d_model, 2, dtype=torch.float) * (-math.log(10000.0) / d_model)
|
| 42 |
+
)
|
| 43 |
+
pe[:, 0::2] = torch.sin(position * div_term)
|
| 44 |
+
pe[:, 1::2] = torch.cos(position * div_term)
|
| 45 |
+
self.register_buffer("pe", pe.unsqueeze(0))
|
| 46 |
+
|
| 47 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 48 |
+
return x + self.pe[:, : x.size(1)]
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
class TagTransformer(nn.Module):
|
| 52 |
+
def __init__(
|
| 53 |
+
self,
|
| 54 |
+
vocab_size: int,
|
| 55 |
+
d_model: int = 256,
|
| 56 |
+
nhead: int = 4,
|
| 57 |
+
num_layers: int = 4,
|
| 58 |
+
dim_feedforward: int = 1024,
|
| 59 |
+
dropout: float = 0.1,
|
| 60 |
+
max_len: int = 256,
|
| 61 |
+
):
|
| 62 |
+
super().__init__()
|
| 63 |
+
self.vocab_size = vocab_size
|
| 64 |
+
self.d_model = d_model
|
| 65 |
+
self.embedding = nn.Embedding(vocab_size + OFFSET, d_model, padding_idx=PAD_ID)
|
| 66 |
+
self.pos_encoding = SinusoidalPositionalEncoding(d_model, max_len=max_len)
|
| 67 |
+
layer = nn.TransformerEncoderLayer(
|
| 68 |
+
d_model=d_model,
|
| 69 |
+
nhead=nhead,
|
| 70 |
+
dim_feedforward=dim_feedforward,
|
| 71 |
+
dropout=dropout,
|
| 72 |
+
batch_first=True,
|
| 73 |
+
norm_first=True,
|
| 74 |
+
)
|
| 75 |
+
self.transformer = nn.TransformerEncoder(layer, num_layers=num_layers)
|
| 76 |
+
self.output = nn.Linear(d_model, vocab_size + OFFSET)
|
| 77 |
+
self._init_weights()
|
| 78 |
+
|
| 79 |
+
def _init_weights(self):
|
| 80 |
+
for p in self.parameters():
|
| 81 |
+
if p.dim() > 1:
|
| 82 |
+
nn.init.xavier_uniform_(p)
|
| 83 |
+
|
| 84 |
+
def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
|
| 85 |
+
seq_len = input_ids.size(1)
|
| 86 |
+
mask = nn.Transformer.generate_square_subsequent_mask(seq_len, device=input_ids.device)
|
| 87 |
+
x = self.embedding(input_ids) * math.sqrt(self.d_model)
|
| 88 |
+
x = self.pos_encoding(x)
|
| 89 |
+
x = self.transformer(x, mask=mask, is_causal=True)
|
| 90 |
+
return self.output(x)
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
class NeuralPromptGenerator:
|
| 94 |
+
GROUP_SUFFIXES = ["_hair", "_eyes"]
|
| 95 |
+
|
| 96 |
+
def __init__(
|
| 97 |
+
self,
|
| 98 |
+
model: TagTransformer,
|
| 99 |
+
graph: SimpleGraph,
|
| 100 |
+
device: Optional[torch.device] = None,
|
| 101 |
+
seed: Optional[int] = None,
|
| 102 |
+
distribution_weight: float = 0.75,
|
| 103 |
+
fallback_top_k: int = 20,
|
| 104 |
+
):
|
| 105 |
+
self.model = model
|
| 106 |
+
self.graph = graph
|
| 107 |
+
self.vocab = graph.vocab
|
| 108 |
+
self.device = device or torch.device("cpu")
|
| 109 |
+
self.model.to(self.device)
|
| 110 |
+
self.model.eval()
|
| 111 |
+
self.rng = np.random.default_rng(seed)
|
| 112 |
+
self.distribution_weight = distribution_weight
|
| 113 |
+
self.fallback_top_k = fallback_top_k
|
| 114 |
+
|
| 115 |
+
self._n_real = len(self.vocab)
|
| 116 |
+
self._real_token_ids = torch.arange(self._n_real, device=self.device) + OFFSET
|
| 117 |
+
self._log_uniform = -np.log(max(self._n_real, 1))
|
| 118 |
+
self._special_ids = {PAD_ID, BOS_ID, EOS_ID}
|
| 119 |
+
|
| 120 |
+
self._group_ids = np.full(self._n_real, -1, dtype=np.int32)
|
| 121 |
+
for i, tag in enumerate(self.vocab.tags):
|
| 122 |
+
for gid, suffix in enumerate(self.GROUP_SUFFIXES):
|
| 123 |
+
if tag.endswith(suffix):
|
| 124 |
+
self._group_ids[i] = gid
|
| 125 |
+
break
|
| 126 |
+
self._group_ids_tensor = torch.from_numpy(self._group_ids).to(self.device)
|
| 127 |
+
|
| 128 |
+
def _target_log_prob(self, idx: int, alpha: float) -> float:
|
| 129 |
+
log_emp = np.log(max(self.vocab.freqs[idx], 1e-12))
|
| 130 |
+
return (1.0 - 2.0 * alpha) * log_emp
|
| 131 |
+
|
| 132 |
+
def _num_people(self, tag_idxs: Set[int]) -> int:
|
| 133 |
+
n = 0
|
| 134 |
+
for idx in tag_idxs:
|
| 135 |
+
tag = self.vocab.tags[idx]
|
| 136 |
+
if tag == "solo":
|
| 137 |
+
n = max(n, 1)
|
| 138 |
+
elif tag in ("1girl", "1boy", "1other"):
|
| 139 |
+
n = max(n, 1)
|
| 140 |
+
elif tag in ("2girls", "2boys", "2others"):
|
| 141 |
+
n = max(n, 2)
|
| 142 |
+
elif tag in ("3girls", "3boys", "3others"):
|
| 143 |
+
n = max(n, 3)
|
| 144 |
+
elif tag in ("4girls", "4boys", "4others"):
|
| 145 |
+
n = max(n, 4)
|
| 146 |
+
elif tag in ("5girls", "5boys", "5others"):
|
| 147 |
+
n = max(n, 5)
|
| 148 |
+
elif tag in (
|
| 149 |
+
"6+girls", "6+boys", "6+others",
|
| 150 |
+
"multiple_girls", "multiple_boys", "multiple_others",
|
| 151 |
+
):
|
| 152 |
+
n = max(n, 2)
|
| 153 |
+
return max(n, 1)
|
| 154 |
+
|
| 155 |
+
def generate(
|
| 156 |
+
self,
|
| 157 |
+
alpha: float,
|
| 158 |
+
count: int,
|
| 159 |
+
length: int = 30,
|
| 160 |
+
anchor: Optional[List[str]] = None,
|
| 161 |
+
blacklist: Optional[List[str]] = None,
|
| 162 |
+
min_prob: float = 0.0,
|
| 163 |
+
temperature: float = 1.0,
|
| 164 |
+
) -> List[List[str]]:
|
| 165 |
+
anchor = anchor or []
|
| 166 |
+
anchor_indices = [self.vocab.tag_to_idx[t] for t in anchor if t in self.vocab.tag_to_idx]
|
| 167 |
+
blacklist = blacklist or []
|
| 168 |
+
blacklist_indices = {self.vocab.tag_to_idx[t] for t in blacklist if t in self.vocab.tag_to_idx}
|
| 169 |
+
|
| 170 |
+
target_bias = torch.tensor(
|
| 171 |
+
[self._target_log_prob(i, alpha) for i in range(self._n_real)],
|
| 172 |
+
dtype=torch.float32,
|
| 173 |
+
device=self.device,
|
| 174 |
+
)
|
| 175 |
+
|
| 176 |
+
results: List[List[str]] = []
|
| 177 |
+
with torch.no_grad():
|
| 178 |
+
for _ in range(count):
|
| 179 |
+
prompt_tokens: List[int] = [BOS_ID] + [idx + OFFSET for idx in anchor_indices]
|
| 180 |
+
present_tag_idxs: Set[int] = set(anchor_indices)
|
| 181 |
+
excluded_tag_idxs: Set[int] = set(blacklist_indices)
|
| 182 |
+
group_counts: Dict[int, int] = {}
|
| 183 |
+
for idx in anchor_indices:
|
| 184 |
+
excluded_tag_idxs.update(self.graph.mutex[idx])
|
| 185 |
+
gid = self._group_ids[idx]
|
| 186 |
+
if gid >= 0:
|
| 187 |
+
group_counts[gid] = group_counts.get(gid, 0) + 1
|
| 188 |
+
|
| 189 |
+
target_len = max(length, len(anchor_indices))
|
| 190 |
+
if target_len > self.model.pos_encoding.pe.size(1) - 1:
|
| 191 |
+
target_len = self.model.pos_encoding.pe.size(1) - 1
|
| 192 |
+
|
| 193 |
+
while len(prompt_tokens) - 1 < target_len:
|
| 194 |
+
input_ids = torch.tensor([prompt_tokens], dtype=torch.long, device=self.device)
|
| 195 |
+
logits = self.model(input_ids)[:, -1, :]
|
| 196 |
+
|
| 197 |
+
model_probs = torch.softmax(logits / max(temperature, 1e-6), dim=-1).squeeze(0)
|
| 198 |
+
allowed = model_probs >= min_prob
|
| 199 |
+
|
| 200 |
+
real_allowed = allowed[self._real_token_ids]
|
| 201 |
+
if not real_allowed.any():
|
| 202 |
+
k = min(self.fallback_top_k, self._n_real)
|
| 203 |
+
topk = torch.topk(model_probs[self._real_token_ids], k=k).indices
|
| 204 |
+
real_allowed = torch.zeros(self._n_real, dtype=torch.bool, device=self.device)
|
| 205 |
+
real_allowed[topk] = True
|
| 206 |
+
allowed = allowed.clone()
|
| 207 |
+
allowed[self._real_token_ids] = real_allowed
|
| 208 |
+
|
| 209 |
+
max_people = self._num_people(present_tag_idxs)
|
| 210 |
+
full_group_ids = [gid for gid, c in group_counts.items() if c >= max_people]
|
| 211 |
+
|
| 212 |
+
biased_logits = logits.squeeze(0).clone()
|
| 213 |
+
biased_logits[self._real_token_ids] += self.distribution_weight * target_bias
|
| 214 |
+
|
| 215 |
+
biased_logits[~allowed] = -float("inf")
|
| 216 |
+
for tid in self._special_ids:
|
| 217 |
+
biased_logits[tid] = -float("inf")
|
| 218 |
+
for idx in present_tag_idxs:
|
| 219 |
+
biased_logits[idx + OFFSET] = -float("inf")
|
| 220 |
+
for idx in excluded_tag_idxs:
|
| 221 |
+
biased_logits[idx + OFFSET] = -float("inf")
|
| 222 |
+
if full_group_ids:
|
| 223 |
+
full_groups_tensor = torch.tensor(full_group_ids, dtype=torch.int32, device=self.device)
|
| 224 |
+
group_full_mask = torch.isin(self._group_ids_tensor, full_groups_tensor)
|
| 225 |
+
biased_logits[self._real_token_ids[group_full_mask]] = -float("inf")
|
| 226 |
+
|
| 227 |
+
probs = torch.softmax(biased_logits / max(temperature, 1e-6), dim=0)
|
| 228 |
+
if not torch.isfinite(probs).all() or probs.sum() == 0:
|
| 229 |
+
break
|
| 230 |
+
probs = probs.cpu().numpy()
|
| 231 |
+
probs = probs / probs.sum()
|
| 232 |
+
|
| 233 |
+
token_id = int(self.rng.choice(len(probs), p=probs))
|
| 234 |
+
if token_id in self._special_ids:
|
| 235 |
+
break
|
| 236 |
+
tag_idx = token_id - OFFSET
|
| 237 |
+
if tag_idx < 0 or tag_idx >= self._n_real:
|
| 238 |
+
break
|
| 239 |
+
if tag_idx in present_tag_idxs or tag_idx in excluded_tag_idxs:
|
| 240 |
+
break
|
| 241 |
+
|
| 242 |
+
prompt_tokens.append(token_id)
|
| 243 |
+
present_tag_idxs.add(tag_idx)
|
| 244 |
+
excluded_tag_idxs.update(self.graph.mutex[tag_idx])
|
| 245 |
+
gid = self._group_ids[tag_idx]
|
| 246 |
+
if gid >= 0:
|
| 247 |
+
group_counts[gid] = group_counts.get(gid, 0) + 1
|
| 248 |
+
|
| 249 |
+
results.append([self.vocab.tags[i - OFFSET] for i in prompt_tokens[1:]])
|
| 250 |
+
|
| 251 |
+
return results
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:2d36d43515b02420363c7bfde27edfe0499f9f438d6cc4ee62dfacf0175b474e
|
| 3 |
+
size 65867168
|
mutex.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
requirements.txt
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
torch>=2.12
|
| 2 |
+
numpy>=1.26
|
| 3 |
+
safetensors>=0.4
|
| 4 |
+
gradio>=5.0
|
vocab.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|