LoliRimuru commited on
Commit
c6dd74a
·
verified ·
1 Parent(s): eb4e70c

Upload 10 files

Browse files
Files changed (10) hide show
  1. README.md +39 -14
  2. app.py +124 -0
  3. config.json +10 -0
  4. counts.json +0 -0
  5. inference.py +93 -0
  6. model.py +251 -0
  7. model.safetensors +3 -0
  8. mutex.json +0 -0
  9. requirements.txt +4 -0
  10. vocab.json +0 -0
README.md CHANGED
@@ -1,14 +1,39 @@
1
- ---
2
- title: BooruPromptRNG
3
- emoji: 📚
4
- colorFrom: blue
5
- colorTo: pink
6
- sdk: gradio
7
- sdk_version: 6.19.0
8
- python_version: '3.13'
9
- app_file: app.py
10
- pinned: false
11
- license: apache-2.0
12
- ---
13
-
14
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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