Upload 10 files
Browse files- README.md +68 -3
- config.json +18 -0
- counts.json +0 -0
- inference.py +100 -0
- model.pt +3 -0
- model.py +277 -0
- model.safetensors +3 -0
- mutex.json +0 -0
- requirements.txt +3 -0
- vocab.json +0 -0
README.md
CHANGED
|
@@ -1,3 +1,68 @@
|
|
| 1 |
-
|
| 2 |
-
|
| 3 |
-
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Booru Smart Prompt Generator
|
| 2 |
+
|
| 3 |
+
A Transformer tag language model trained on Danbooru-style metadata that generates coherent tag prompts.
|
| 4 |
+
|
| 5 |
+
## Files
|
| 6 |
+
|
| 7 |
+
- `model.safetensors` — model weights
|
| 8 |
+
- `config.json` — model architecture / generation defaults
|
| 9 |
+
- `vocab.json` — tag vocabulary
|
| 10 |
+
- `counts.json` — per-tag occurrence counts
|
| 11 |
+
- `mutex.json` — mutually exclusive tag pairs
|
| 12 |
+
- `model.py` — self-contained model + sampler
|
| 13 |
+
- `inference.py` — command-line generator
|
| 14 |
+
|
| 15 |
+
## Install
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
pip install -r requirements.txt
|
| 19 |
+
```
|
| 20 |
+
|
| 21 |
+
For the RTX 5090 / CUDA 13.x setup used during training, install the matching PyTorch wheel, e.g.:
|
| 22 |
+
|
| 23 |
+
```bash
|
| 24 |
+
pip install torch==2.12.1+cu130 --index-url https://download.pytorch.org/whl/cu130
|
| 25 |
+
```
|
| 26 |
+
|
| 27 |
+
## Generate prompts
|
| 28 |
+
|
| 29 |
+
```bash
|
| 30 |
+
# Empirical mode (follow Booru distribution)
|
| 31 |
+
python inference.py --mode empirical --count 10 --length 30
|
| 32 |
+
|
| 33 |
+
# Diverse mode (boost rare tags)
|
| 34 |
+
python inference.py --mode diverse --count 10 --length 30
|
| 35 |
+
|
| 36 |
+
# Anchor + blacklist + content rating
|
| 37 |
+
python inference.py \
|
| 38 |
+
--mode empirical --count 5 --length 30 --rating s \
|
| 39 |
+
--anchor "1girl,black_hair" --blacklist "1boy,smile"
|
| 40 |
+
|
| 41 |
+
# Constrained sampling (top-k / nucleus)
|
| 42 |
+
python inference.py \
|
| 43 |
+
--mode empirical --count 5 --length 30 \
|
| 44 |
+
--temperature 0.8 --top-k 50 --top-p 0.95
|
| 45 |
+
```
|
| 46 |
+
|
| 47 |
+
## Parameters
|
| 48 |
+
|
| 49 |
+
| Flag | Default | Description |
|
| 50 |
+
|------|---------|-------------|
|
| 51 |
+
| `--mode` | — | `empirical` (alpha=0) or `diverse` (alpha=1) |
|
| 52 |
+
| `--alpha` | 0.0 | Fine-grained distribution coefficient (0=empirical, 1=diverse) |
|
| 53 |
+
| `--rating` | `g` | Content rating token: `g` (general), `s` (sensitive), `q` (questionable), `e` (explicit) |
|
| 54 |
+
| `--count` | 10 | Number of prompts to generate |
|
| 55 |
+
| `--length` | 30 | Tags per prompt, maximum 128 |
|
| 56 |
+
| `--anchor` | — | Comma-separated tags that must appear in every prompt |
|
| 57 |
+
| `--blacklist` | — | Comma-separated tags the model must not generate |
|
| 58 |
+
| `--min-prob` | 0.0005 | Minimum raw model probability for a candidate tag |
|
| 59 |
+
| `--temperature` | 1.0 | Sampling temperature: lower = more focused, higher = more random |
|
| 60 |
+
| `--top-k` | 0 | Top-k sampling: keep only the k most likely tags. 0 disables it |
|
| 61 |
+
| `--top-p` | 1.0 | Nucleus / top-p sampling: keep the smallest set whose cumulative probability exceeds p. 1.0 disables it |
|
| 62 |
+
| `--distribution-weight` | 0.75 | Strength of the empirical/diverse bias applied to model scores |
|
| 63 |
+
| `--seed` | — | Random seed for reproducible generation |
|
| 64 |
+
|
| 65 |
+
## Notes
|
| 66 |
+
|
| 67 |
+
- Maximum prompt length is **128 tags** (including anchor), but quality is best around 30.
|
| 68 |
+
- Tags not in `vocab.json` are ignored for anchor/blacklist.
|
config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
+
"pad_id": 0,
|
| 10 |
+
"bos_id": 1,
|
| 11 |
+
"eos_id": 2,
|
| 12 |
+
"rating_g": 3,
|
| 13 |
+
"rating_s": 4,
|
| 14 |
+
"rating_q": 5,
|
| 15 |
+
"rating_e": 6,
|
| 16 |
+
"offset": 7,
|
| 17 |
+
"distribution_weight": 0.5
|
| 18 |
+
}
|
counts.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
inference.py
ADDED
|
@@ -0,0 +1,100 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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"] + 2,
|
| 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("--rating", default="g", choices=["g", "s", "q", "e"],
|
| 58 |
+
help="Content rating token to condition on")
|
| 59 |
+
parser.add_argument("--min-prob", type=float, default=0.0005)
|
| 60 |
+
parser.add_argument("--temperature", type=float, default=1.0)
|
| 61 |
+
parser.add_argument("--top-k", type=int, default=0, help="Top-k sampling (0 = disabled)")
|
| 62 |
+
parser.add_argument("--top-p", type=float, default=1.0, help="Nucleus/top-p sampling (1.0 = disabled)")
|
| 63 |
+
parser.add_argument("--distribution-weight", type=float, default=None)
|
| 64 |
+
parser.add_argument("--seed", type=int, default=None)
|
| 65 |
+
args = parser.parse_args()
|
| 66 |
+
|
| 67 |
+
alpha = args.alpha
|
| 68 |
+
if args.mode == "empirical":
|
| 69 |
+
alpha = 0.0
|
| 70 |
+
elif args.mode == "diverse":
|
| 71 |
+
alpha = 1.0
|
| 72 |
+
if alpha is None:
|
| 73 |
+
alpha = 0.0
|
| 74 |
+
|
| 75 |
+
gen = load_generator(args.release_dir, seed=args.seed)
|
| 76 |
+
if args.distribution_weight is not None:
|
| 77 |
+
gen.distribution_weight = args.distribution_weight
|
| 78 |
+
|
| 79 |
+
anchor_tags = [t.strip() for t in args.anchor.split(",") if t.strip()]
|
| 80 |
+
blacklist_tags = [t.strip() for t in args.blacklist.split(",") if t.strip()]
|
| 81 |
+
|
| 82 |
+
prompts = gen.generate(
|
| 83 |
+
alpha=alpha,
|
| 84 |
+
count=args.count,
|
| 85 |
+
length=args.length,
|
| 86 |
+
anchor=anchor_tags or None,
|
| 87 |
+
blacklist=blacklist_tags or None,
|
| 88 |
+
min_prob=args.min_prob,
|
| 89 |
+
temperature=args.temperature,
|
| 90 |
+
top_k=getattr(args, "top_k", 0),
|
| 91 |
+
top_p=getattr(args, "top_p", 1.0),
|
| 92 |
+
rating=args.rating,
|
| 93 |
+
)
|
| 94 |
+
|
| 95 |
+
for tags in prompts:
|
| 96 |
+
print(", ".join(tags))
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
if __name__ == "__main__":
|
| 100 |
+
main()
|
model.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:df1d367f5d3a6fa475f05ac24c0716fb256c987e1ab003350ecb3c52331cbf06
|
| 3 |
+
size 65754008
|
model.py
ADDED
|
@@ -0,0 +1,277 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
+
RATING_G = 3
|
| 15 |
+
RATING_S = 4
|
| 16 |
+
RATING_Q = 5
|
| 17 |
+
RATING_E = 6
|
| 18 |
+
OFFSET = 7
|
| 19 |
+
|
| 20 |
+
_RATING_TOKENS = {"g": RATING_G, "s": RATING_S, "q": RATING_Q, "e": RATING_E}
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def _apply_top_k_top_p(probs: torch.Tensor, top_k: int, top_p: float) -> torch.Tensor:
|
| 24 |
+
"""Filter a probability distribution with top-k and/or nucleus (top-p) sampling."""
|
| 25 |
+
if top_k > 0:
|
| 26 |
+
k = min(top_k, probs.size(0))
|
| 27 |
+
threshold = torch.topk(probs, k).values[-1]
|
| 28 |
+
probs = probs.where(probs >= threshold, torch.zeros_like(probs))
|
| 29 |
+
if top_p < 1.0:
|
| 30 |
+
sorted_probs, sorted_idx = torch.sort(probs, descending=True)
|
| 31 |
+
cumsum = torch.cumsum(sorted_probs, dim=0)
|
| 32 |
+
nucleus_mask = cumsum <= top_p
|
| 33 |
+
if nucleus_mask.any():
|
| 34 |
+
nucleus_mask[0] = True
|
| 35 |
+
kept = torch.zeros_like(probs, dtype=torch.bool)
|
| 36 |
+
kept.scatter_(0, sorted_idx, nucleus_mask)
|
| 37 |
+
probs = probs.where(kept, torch.zeros_like(probs))
|
| 38 |
+
return probs
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
class Vocab:
|
| 42 |
+
def __init__(self, tags: List[str], counts: List[int]):
|
| 43 |
+
self.tags = tags
|
| 44 |
+
self.tag_to_idx = {tag: i for i, tag in enumerate(tags)}
|
| 45 |
+
self.counts = np.array(counts, dtype=np.int64)
|
| 46 |
+
self.total = int(self.counts.sum())
|
| 47 |
+
self.freqs = self.counts.astype(np.float64) / max(self.total, 1)
|
| 48 |
+
|
| 49 |
+
def __len__(self) -> int:
|
| 50 |
+
return len(self.tags)
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
class SimpleGraph:
|
| 54 |
+
def __init__(self, vocab: Vocab, mutex: List[List[int]]):
|
| 55 |
+
self.vocab = vocab
|
| 56 |
+
self.mutex = mutex
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
class TagTransformer(nn.Module):
|
| 60 |
+
"""Permutation-equivariant Transformer for tag-set generation."""
|
| 61 |
+
|
| 62 |
+
def __init__(
|
| 63 |
+
self,
|
| 64 |
+
vocab_size: int,
|
| 65 |
+
d_model: int = 256,
|
| 66 |
+
nhead: int = 4,
|
| 67 |
+
num_layers: int = 4,
|
| 68 |
+
dim_feedforward: int = 1024,
|
| 69 |
+
dropout: float = 0.1,
|
| 70 |
+
max_len: int = 256,
|
| 71 |
+
):
|
| 72 |
+
super().__init__()
|
| 73 |
+
self.vocab_size = vocab_size
|
| 74 |
+
self.d_model = d_model
|
| 75 |
+
self.max_len = max_len
|
| 76 |
+
self.embedding = nn.Embedding(vocab_size + OFFSET, d_model, padding_idx=PAD_ID)
|
| 77 |
+
layer = nn.TransformerEncoderLayer(
|
| 78 |
+
d_model=d_model,
|
| 79 |
+
nhead=nhead,
|
| 80 |
+
dim_feedforward=dim_feedforward,
|
| 81 |
+
dropout=dropout,
|
| 82 |
+
batch_first=True,
|
| 83 |
+
norm_first=True,
|
| 84 |
+
)
|
| 85 |
+
self.transformer = nn.TransformerEncoder(layer, num_layers=num_layers, enable_nested_tensor=False)
|
| 86 |
+
self.output = nn.Linear(d_model, vocab_size + OFFSET)
|
| 87 |
+
self._init_weights()
|
| 88 |
+
|
| 89 |
+
def _init_weights(self):
|
| 90 |
+
for p in self.parameters():
|
| 91 |
+
if p.dim() > 1:
|
| 92 |
+
nn.init.xavier_uniform_(p)
|
| 93 |
+
|
| 94 |
+
def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
|
| 95 |
+
seq_len = input_ids.size(1)
|
| 96 |
+
mask = nn.Transformer.generate_square_subsequent_mask(seq_len, device=input_ids.device).bool()
|
| 97 |
+
x = self.embedding(input_ids)
|
| 98 |
+
padding_mask = input_ids == PAD_ID
|
| 99 |
+
x = self.transformer(
|
| 100 |
+
x,
|
| 101 |
+
mask=mask,
|
| 102 |
+
src_key_padding_mask=padding_mask,
|
| 103 |
+
is_causal=True,
|
| 104 |
+
)
|
| 105 |
+
return self.output(x)
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
class NeuralPromptGenerator:
|
| 109 |
+
GROUP_SUFFIXES = ["_hair", "_eyes"]
|
| 110 |
+
|
| 111 |
+
def __init__(
|
| 112 |
+
self,
|
| 113 |
+
model: TagTransformer,
|
| 114 |
+
graph: SimpleGraph,
|
| 115 |
+
device: Optional[torch.device] = None,
|
| 116 |
+
seed: Optional[int] = None,
|
| 117 |
+
distribution_weight: float = 0.75,
|
| 118 |
+
fallback_top_k: int = 20,
|
| 119 |
+
):
|
| 120 |
+
self.model = model
|
| 121 |
+
self.graph = graph
|
| 122 |
+
self.vocab = graph.vocab
|
| 123 |
+
self.device = device or torch.device("cpu")
|
| 124 |
+
self.model.to(self.device)
|
| 125 |
+
self.model.eval()
|
| 126 |
+
self.rng = np.random.default_rng(seed)
|
| 127 |
+
self.distribution_weight = distribution_weight
|
| 128 |
+
self.fallback_top_k = fallback_top_k
|
| 129 |
+
|
| 130 |
+
self._n_real = len(self.vocab)
|
| 131 |
+
self._real_token_ids = torch.arange(self._n_real, device=self.device) + OFFSET
|
| 132 |
+
self._log_uniform = -np.log(max(self._n_real, 1))
|
| 133 |
+
self._special_ids = {PAD_ID, BOS_ID, EOS_ID}
|
| 134 |
+
|
| 135 |
+
self._group_ids = np.full(self._n_real, -1, dtype=np.int32)
|
| 136 |
+
for i, tag in enumerate(self.vocab.tags):
|
| 137 |
+
for gid, suffix in enumerate(self.GROUP_SUFFIXES):
|
| 138 |
+
if tag.endswith(suffix):
|
| 139 |
+
self._group_ids[i] = gid
|
| 140 |
+
break
|
| 141 |
+
self._group_ids_tensor = torch.from_numpy(self._group_ids).to(self.device)
|
| 142 |
+
|
| 143 |
+
def _target_log_prob(self, idx: int, alpha: float) -> float:
|
| 144 |
+
log_emp = np.log(max(self.vocab.freqs[idx], 1e-12))
|
| 145 |
+
return (1.0 - 2.0 * alpha) * log_emp
|
| 146 |
+
|
| 147 |
+
def _num_people(self, tag_idxs: Set[int]) -> int:
|
| 148 |
+
n = 0
|
| 149 |
+
for idx in tag_idxs:
|
| 150 |
+
tag = self.vocab.tags[idx]
|
| 151 |
+
if tag == "solo":
|
| 152 |
+
n = max(n, 1)
|
| 153 |
+
elif tag in ("1girl", "1boy", "1other"):
|
| 154 |
+
n = max(n, 1)
|
| 155 |
+
elif tag in ("2girls", "2boys", "2others"):
|
| 156 |
+
n = max(n, 2)
|
| 157 |
+
elif tag in ("3girls", "3boys", "3others"):
|
| 158 |
+
n = max(n, 3)
|
| 159 |
+
elif tag in ("4girls", "4boys", "4others"):
|
| 160 |
+
n = max(n, 4)
|
| 161 |
+
elif tag in ("5girls", "5boys", "5others"):
|
| 162 |
+
n = max(n, 5)
|
| 163 |
+
elif tag in (
|
| 164 |
+
"6+girls", "6+boys", "6+others",
|
| 165 |
+
"multiple_girls", "multiple_boys", "multiple_others",
|
| 166 |
+
):
|
| 167 |
+
n = max(n, 2)
|
| 168 |
+
return max(n, 1)
|
| 169 |
+
|
| 170 |
+
def generate(
|
| 171 |
+
self,
|
| 172 |
+
alpha: float,
|
| 173 |
+
count: int,
|
| 174 |
+
length: int = 30,
|
| 175 |
+
anchor: Optional[List[str]] = None,
|
| 176 |
+
blacklist: Optional[List[str]] = None,
|
| 177 |
+
min_prob: float = 0.0,
|
| 178 |
+
temperature: float = 1.0,
|
| 179 |
+
top_k: int = 0,
|
| 180 |
+
top_p: float = 1.0,
|
| 181 |
+
rating: Optional[str] = None,
|
| 182 |
+
) -> List[List[str]]:
|
| 183 |
+
anchor = anchor or []
|
| 184 |
+
anchor_indices = [self.vocab.tag_to_idx[t] for t in anchor if t in self.vocab.tag_to_idx]
|
| 185 |
+
blacklist = blacklist or []
|
| 186 |
+
blacklist_indices = {self.vocab.tag_to_idx[t] for t in blacklist if t in self.vocab.tag_to_idx}
|
| 187 |
+
rating_token = _RATING_TOKENS.get((rating or "g").lower(), RATING_G)
|
| 188 |
+
|
| 189 |
+
target_bias = torch.tensor(
|
| 190 |
+
[self._target_log_prob(i, alpha) for i in range(self._n_real)],
|
| 191 |
+
dtype=torch.float32,
|
| 192 |
+
device=self.device,
|
| 193 |
+
)
|
| 194 |
+
|
| 195 |
+
results: List[List[str]] = []
|
| 196 |
+
with torch.no_grad():
|
| 197 |
+
for _ in range(count):
|
| 198 |
+
prompt_tokens: List[int] = [BOS_ID, rating_token] + [idx + OFFSET for idx in anchor_indices]
|
| 199 |
+
present_tag_idxs: Set[int] = set(anchor_indices)
|
| 200 |
+
excluded_tag_idxs: Set[int] = set(blacklist_indices)
|
| 201 |
+
group_counts: Dict[int, int] = {}
|
| 202 |
+
for idx in anchor_indices:
|
| 203 |
+
excluded_tag_idxs.update(self.graph.mutex[idx])
|
| 204 |
+
gid = self._group_ids[idx]
|
| 205 |
+
if gid >= 0:
|
| 206 |
+
group_counts[gid] = group_counts.get(gid, 0) + 1
|
| 207 |
+
|
| 208 |
+
target_len = max(length, len(anchor_indices))
|
| 209 |
+
max_len = getattr(self.model, "max_len", 256)
|
| 210 |
+
if target_len > max_len - 2:
|
| 211 |
+
target_len = max_len - 2
|
| 212 |
+
|
| 213 |
+
while len(prompt_tokens) - 1 < target_len:
|
| 214 |
+
input_ids = torch.tensor([prompt_tokens], dtype=torch.long, device=self.device)
|
| 215 |
+
logits = self.model(input_ids)[:, -1, :]
|
| 216 |
+
|
| 217 |
+
model_probs = torch.softmax(logits / max(temperature, 1e-6), dim=-1).squeeze(0)
|
| 218 |
+
allowed = model_probs >= min_prob
|
| 219 |
+
|
| 220 |
+
real_allowed = allowed[self._real_token_ids]
|
| 221 |
+
if not real_allowed.any():
|
| 222 |
+
k = min(self.fallback_top_k, self._n_real)
|
| 223 |
+
topk = torch.topk(model_probs[self._real_token_ids], k=k).indices
|
| 224 |
+
real_allowed = torch.zeros(self._n_real, dtype=torch.bool, device=self.device)
|
| 225 |
+
real_allowed[topk] = True
|
| 226 |
+
allowed = allowed.clone()
|
| 227 |
+
allowed[self._real_token_ids] = real_allowed
|
| 228 |
+
|
| 229 |
+
max_people = self._num_people(present_tag_idxs)
|
| 230 |
+
full_group_ids = [gid for gid, c in group_counts.items() if c >= max_people]
|
| 231 |
+
|
| 232 |
+
biased_logits = logits.squeeze(0).clone()
|
| 233 |
+
biased_logits[self._real_token_ids] += self.distribution_weight * target_bias
|
| 234 |
+
|
| 235 |
+
biased_logits[~allowed] = -float("inf")
|
| 236 |
+
for tid in self._special_ids:
|
| 237 |
+
biased_logits[tid] = -float("inf")
|
| 238 |
+
for idx in present_tag_idxs:
|
| 239 |
+
biased_logits[idx + OFFSET] = -float("inf")
|
| 240 |
+
for idx in excluded_tag_idxs:
|
| 241 |
+
biased_logits[idx + OFFSET] = -float("inf")
|
| 242 |
+
if full_group_ids:
|
| 243 |
+
full_groups_tensor = torch.tensor(full_group_ids, dtype=torch.int32, device=self.device)
|
| 244 |
+
group_full_mask = torch.isin(self._group_ids_tensor, full_groups_tensor)
|
| 245 |
+
biased_logits[self._real_token_ids[group_full_mask]] = -float("inf")
|
| 246 |
+
|
| 247 |
+
probs = torch.softmax(biased_logits / max(temperature, 1e-6), dim=0)
|
| 248 |
+
probs = _apply_top_k_top_p(probs, top_k, top_p)
|
| 249 |
+
if not torch.isfinite(probs).all() or probs.sum() <= 0:
|
| 250 |
+
break
|
| 251 |
+
probs = probs / probs.sum()
|
| 252 |
+
probs = probs.cpu().numpy()
|
| 253 |
+
probs = probs / probs.sum()
|
| 254 |
+
|
| 255 |
+
token_id = int(self.rng.choice(len(probs), p=probs))
|
| 256 |
+
if token_id in self._special_ids:
|
| 257 |
+
break
|
| 258 |
+
tag_idx = token_id - OFFSET
|
| 259 |
+
if tag_idx < 0 or tag_idx >= self._n_real:
|
| 260 |
+
break
|
| 261 |
+
if tag_idx in present_tag_idxs or tag_idx in excluded_tag_idxs:
|
| 262 |
+
break
|
| 263 |
+
|
| 264 |
+
prompt_tokens.append(token_id)
|
| 265 |
+
present_tag_idxs.add(tag_idx)
|
| 266 |
+
excluded_tag_idxs.update(self.graph.mutex[tag_idx])
|
| 267 |
+
gid = self._group_ids[tag_idx]
|
| 268 |
+
if gid >= 0:
|
| 269 |
+
group_counts[gid] = group_counts.get(gid, 0) + 1
|
| 270 |
+
|
| 271 |
+
results.append([
|
| 272 |
+
self.vocab.tags[i - OFFSET]
|
| 273 |
+
for i in prompt_tokens
|
| 274 |
+
if i >= OFFSET
|
| 275 |
+
])
|
| 276 |
+
|
| 277 |
+
return results
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:79fa338fbdd18c502f462b43117109b279775a17bf25b3764fdbb56a094f62a1
|
| 3 |
+
size 65743184
|
mutex.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
requirements.txt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
torch>=2.12
|
| 2 |
+
numpy>=1.26
|
| 3 |
+
safetensors>=0.4
|
vocab.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|