LoliRimuru commited on
Commit
8b27df3
·
verified ·
1 Parent(s): c61ee15

Upload 10 files

Browse files
Files changed (10) hide show
  1. README.md +68 -3
  2. config.json +18 -0
  3. counts.json +0 -0
  4. inference.py +100 -0
  5. model.pt +3 -0
  6. model.py +277 -0
  7. model.safetensors +3 -0
  8. mutex.json +0 -0
  9. requirements.txt +3 -0
  10. vocab.json +0 -0
README.md CHANGED
@@ -1,3 +1,68 @@
1
- ---
2
- license: apache-2.0
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