#!/usr/bin/env python3
"""
DanbooruTagQuery — Hugging Face ZeroGPU Space
Usage:
python app.py # download model from HF hub
python app.py /path/to/model.onnx # use local model
MODEL_DIR=/path/to python app.py # env var with model dir
"""
from __future__ import annotations
import json
import os
import sys
import tempfile
import time
from pathlib import Path
try:
import spaces
except ImportError:
class MockSpaces:
def GPU(self, func=None, **kwargs):
if func is not None and callable(func):
# Used as @spaces.GPU
return func
# Used as @spaces.GPU(...)
def decorator(f):
return f
return decorator
spaces = MockSpaces()
import gradio as gr
import numpy as np
from PIL import Image
# ── optional deps (loaded on demand) ────────────────────────────────────────
_hf_hub = None
def _import_hf_hub():
global _hf_hub
if _hf_hub is None:
import huggingface_hub as h
_hf_hub = h
return _hf_hub
# ── constants ───────────────────────────────────────────────────────────────
HF_REPO = "realphongha/danbooru-tag-query"
MODELS_DIR = "models"
CATEGORY_JSON = "tag_category.json"
IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
CATEGORY_MAP = {
0: "general",
1: "artist",
3: "copyright",
4: "character",
5: "meta",
}
DEFAULT_TOP_K = None
DEFAULT_MIN_SCORE = 0.2
# ── category lookup ─────────────────────────────────────────────────────────
def load_category_map(checkpoint: str | Path) -> dict[str, int]:
"""Load tag→category map. Tries HF hub download, then sidecar file.
Returns {tag: category_id} — all tags default to 0 (general).
"""
ckpt = Path(checkpoint)
# 1) sidecar: model.onnx → tag_category.json beside it
if ckpt.suffix == ".onnx":
sidecar = ckpt.with_name(CATEGORY_JSON)
if sidecar.exists():
return json.loads(sidecar.read_text())
# 2) parent dir: dir/model.onnx → dir/tag_category.json
parent_sidecar = ckpt.parent / CATEGORY_JSON
if parent_sidecar.exists():
return json.loads(parent_sidecar.read_text())
# 3) HF hub: download alongside model variant
if ckpt.suffix == ".onnx":
# try to infer variant from path
parts = ckpt.parts
for i, p in enumerate(parts):
if p == MODELS_DIR and i + 2 < len(parts):
variant = parts[i + 1]
try:
hf = _import_hf_hub()
path = hf.hf_hub_download(
repo_id=HF_REPO,
filename=f"{MODELS_DIR}/{variant}/{CATEGORY_JSON}",
repo_type="model",
)
return json.loads(Path(path).read_text())
except Exception:
pass
break
return {}
def get_category_name(cat_map: dict[str, int], tag: str) -> str:
cat_id = cat_map.get(tag, 0)
return CATEGORY_MAP.get(cat_id, "general")
# ── image preprocessing ────────────────────────────────────────────────────
def preprocess(image: Image.Image, image_size: int = 448) -> np.ndarray:
w, h = image.size
scale = image_size / max(w, h)
new_w = int(w * scale)
new_h = int(h * scale)
image = image.resize((new_w, new_h), Image.BILINEAR)
canvas = Image.new("RGB", (image_size, image_size), (0, 0, 0))
left = (image_size - new_w) // 2
top = (image_size - new_h) // 2
canvas.paste(image, (left, top))
arr = np.asarray(canvas, dtype=np.float32).transpose(2, 0, 1) / 255.0
arr[0] = (arr[0] - IMAGENET_MEAN[0]) / IMAGENET_STD[0]
arr[1] = (arr[1] - IMAGENET_MEAN[1]) / IMAGENET_STD[1]
arr[2] = (arr[2] - IMAGENET_MEAN[2]) / IMAGENET_STD[2]
return arr[np.newaxis, ...]
# ── sidecar loading ────────────────────────────────────────────────────────
def load_tag_to_id(checkpoint: str | Path) -> dict[str, int]:
ckpt = Path(checkpoint)
path = _sidecar_path(ckpt, ".tag_to_id.json")
if not path.exists():
path = ckpt.parent / "tag_to_id.json"
if not path.exists():
raise FileNotFoundError(f"Missing tag map: {path}")
return json.loads(path.read_text())
def load_config(checkpoint: str | Path) -> dict:
ckpt = Path(checkpoint)
path = _sidecar_path(ckpt, ".config.json")
if not path.exists():
path = ckpt.parent / "config.json"
if not path.exists():
return {"image_size": 448}
return json.loads(path.read_text())
def _sidecar_path(checkpoint: Path, suffix: str) -> Path:
if checkpoint.suffix == ".onnx":
return checkpoint.with_name(checkpoint.stem + suffix)
return checkpoint / suffix.lstrip(".")
# ── Predictor (ONNX) ──────────────────────────────────────────────────────
class Predictor:
def __init__(self, checkpoint: str | Path):
import onnxruntime as ort
self.checkpoint = str(checkpoint)
self.tag_to_id = load_tag_to_id(self.checkpoint)
cfg = load_config(self.checkpoint)
self.image_size = cfg.get("image_size", 448)
self.cat_map = load_category_map(self.checkpoint)
providers = [
("CUDAExecutionProvider", {}),
"CPUExecutionProvider",
]
try:
self._sess = ort.InferenceSession(self.checkpoint, providers=providers)
except Exception:
self._sess = ort.InferenceSession(
self.checkpoint, providers=["CPUExecutionProvider"]
)
self._input_name = self._sess.get_inputs()[0].name
self._output_name = self._sess.get_outputs()[0].name
def run(self, pixel_values: np.ndarray) -> np.ndarray:
raw = self._sess.run([self._output_name], {self._input_name: pixel_values})[0]
return 1.0 / (1.0 + np.exp(-raw))
@property
def num_classes(self) -> int:
return len(self.tag_to_id)
def category_name(self, tag: str) -> str:
return get_category_name(self.cat_map, tag)
# ── model discovery & loading (HF hub) ─────────────────────────────────────
def discover_model_variants() -> list[str]:
try:
hf = _import_hf_hub()
api = hf.HfApi()
siblings = api.list_repo_files(HF_REPO, repo_type="model")
variants: set[str] = set()
for path in siblings:
if path.startswith(f"{MODELS_DIR}/") and "/" in path[len(MODELS_DIR) + 1:]:
variant = path.split("/")[1]
if variant:
variants.add(variant)
return sorted(variants, reverse=True)
except Exception as exc:
print(f"Warning: could not discover models on hub: {exc}")
return []
def download_model_variant(variant: str) -> Path:
hf = _import_hf_hub()
onnx_path = hf.hf_hub_download(
repo_id=HF_REPO,
filename=f"{MODELS_DIR}/{variant}/model.onnx",
repo_type="model",
)
# download sidecars (config, tag_to_id, tag_category) if they exist
for sidecar in ["config.json", "tag_to_id.json", CATEGORY_JSON]:
try:
hf.hf_hub_download(
repo_id=HF_REPO,
filename=f"{MODELS_DIR}/{variant}/{sidecar}",
repo_type="model",
)
except Exception:
pass # optional — missing is fine
return Path(onnx_path)
# ── Gradio UI ──────────────────────────────────────────────────────────────
def build_app(predict_fn, model_choices: list[str]) -> gr.Blocks:
state = {
"all_logits": None,
"tag_metadata": None,
"current_image": None,
"predictor": None,
"predict_fn": predict_fn,
}
css = """
#csv-wrap { position: relative; }
#copy-csv-btn { position: absolute; top: 4px; right: 4px; z-index: 10;
min-width: 0; padding: 0 6px; height: 24px;
font-size: 13px; line-height: 24px; }
"""
category_names = sorted(CATEGORY_MAP.values())
with gr.Blocks(title="DanbooruTagQuery", theme=gr.themes.Soft(), css=css) as app:
gr.Markdown("# 🏷️ DanbooruTagQuery")
with gr.Row():
with gr.Column(scale=1):
image_input = gr.Image(
label="Image",
type="pil",
sources=["upload", "clipboard"],
height=300,
)
url_input = gr.Textbox(
label="Image URL",
placeholder="Paste image URL and press Enter",
)
with gr.Row():
analyze_btn = gr.Button("🔍 Analyze", variant="primary", scale=2)
clear_btn = gr.Button("🗑️ Clear", scale=1)
with gr.Column():
gr.Markdown('[📄 **Model Card**](https://huggingface.co/realphongha/danbooru-tag-query)')
gr.Markdown("### 🤖 Model")
model_dropdown = gr.Dropdown(
choices=model_choices,
value=model_choices[0] if model_choices else None,
label="Model variant",
interactive=True,
)
model_status = gr.Markdown("Ready")
with gr.Column(scale=1):
top_k = gr.Number(
label="Top-K", value=DEFAULT_TOP_K, minimum=0, step=1
)
min_score = gr.Slider(
label="Min Score",
value=DEFAULT_MIN_SCORE,
minimum=0.0,
maximum=1.0,
step=0.01,
)
sort_by = gr.Radio(
label="Sort by", choices=["score", "name"], value="score"
)
use_underscore = gr.Checkbox(
label="Use underscore (_)", value=False
)
categories = gr.CheckboxGroup(
label="Categories",
choices=category_names,
value=["general"],
)
with gr.Tabs():
with gr.TabItem("📋 Tag list"):
tag_table = gr.HTML(label="Tags")
with gr.TabItem("📝 Comma-separated"):
with gr.Column(elem_id="csv-wrap"):
tag_string = gr.Textbox(label="Tags", lines=6, elem_id="csv-text")
copy_btn = gr.Button("📋", elem_id="copy-csv-btn")
with gr.Row():
status = gr.Markdown("Ready. Load an image and click **Analyze**.")
gr.Markdown("### 🔍 Tag score query")
with gr.Row():
tag_query = gr.Textbox(
label="Search tag",
placeholder="Type to search…",
scale=3,
)
tag_query_output = gr.HTML(label="Results")
# ── callbacks ──────────────────────────────────────────────────────
def refresh_results(
_top_k, _min_score, _sort_by, _use_underscore, _categories,
):
if state["all_logits"] is None or state["tag_metadata"] is None:
return "No results yet.", ""
meta = state["tag_metadata"]
all_tags = list(meta.keys())
if _categories:
all_tags = [
t for t in all_tags
if meta[t]["category_name"] in _categories
]
items = [(t, meta[t]["score"]) for t in all_tags]
if _sort_by == "name":
items.sort(key=lambda x: format_tag(x[0], _use_underscore))
else:
items.sort(key=lambda x: x[1], reverse=True)
items = [(t, s) for t, s in items if s >= _min_score]
if _top_k is not None and _top_k > 0:
items = items[:_top_k]
if not items:
return "No tags pass the filters.", ""
rows = []
for tag, score in items:
m = meta[tag]
link = (
f'{tag}'
)
display = format_tag(tag, _use_underscore)
rows.append(
f"
"
f"| {link} | "
f"{display} | "
f"{score:.4f} | "
f"{m['category_name']} | "
f"
"
)
table = (
''
''
'| Link | Tag | '
'Score | '
'Category | '
'
'
'' + "".join(rows) + '
'
)
csv = ", ".join(format_tag(t, _use_underscore) for t, _ in items)
return table, csv
def on_analyze(image, url):
if image is None and not url:
return *refresh_results(
top_k.value, min_score.value,
sort_by.value, use_underscore.value,
categories.value,
), "⚠️ No image loaded."
pil = image
if pil is None and url:
import requests as std_requests
try:
resp = std_requests.get(url, timeout=30)
resp.raise_for_status()
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
tmp.write(resp.content)
tmp.close()
pil = Image.open(tmp.name).convert("RGB")
except Exception as exc:
return "Error loading URL.", "", f"❌ {exc}"
state["current_image"] = pil
fn = state["predict_fn"]
if fn is None:
return "No model loaded.", "", "❌ No model loaded."
t0 = time.time()
all_logits = fn(pil)
state["all_logits"] = all_logits
state["tag_metadata"] = enrich_tags(
all_logits, state["predictor"].cat_map if state["predictor"] else {}
)
table, csv = refresh_results(
top_k.value, min_score.value,
sort_by.value, use_underscore.value,
categories.value,
)
elapsed = time.time() - t0
n = len(state["tag_metadata"])
return table, csv, f"✅ {n} tags · {elapsed:.2f}s"
analyze_btn.click(
fn=on_analyze,
inputs=[image_input, url_input],
outputs=[tag_table, tag_string, status],
)
url_input.submit(
fn=on_analyze,
inputs=[image_input, url_input],
outputs=[tag_table, tag_string, status],
)
def on_clear():
state["all_logits"] = None
state["tag_metadata"] = None
state["current_image"] = None
return None, "", "No results yet.", "Cleared.", "", "", "No results yet."
clear_btn.click(
fn=on_clear,
inputs=[],
outputs=[image_input, url_input, tag_table, tag_string,
status, tag_query, tag_query_output],
)
for widget in [top_k, min_score, sort_by, use_underscore, categories]:
widget.change(
fn=refresh_results,
inputs=[top_k, min_score, sort_by, use_underscore, categories],
outputs=[tag_table, tag_string],
)
def query_tag_score(query):
meta = state.get("tag_metadata")
if not meta or not query:
return ""
query_l = query.lower()
matches = sorted(
[(t, meta[t]["score"]) for t in meta if query_l in t.lower()],
key=lambda x: x[1], reverse=True,
)[:20]
if not matches:
return "No matching tags."
rows = "".join(
f"| {t} | {s:.4f} | "
f"{meta[t]['category_name']} |
"
for t, s in matches
)
return (f""
f"| Tag | Score | Category |
"
f"{rows}
")
tag_query.change(
fn=query_tag_score,
inputs=[tag_query],
outputs=[tag_query_output],
)
copy_btn.click(
fn=lambda: None,
inputs=[],
outputs=[],
js="""() => {
const tb = document.querySelector('#csv-text textarea');
if (tb) { navigator.clipboard.writeText(tb.value); }
}"""
)
# HF ZeroGPU requires @spaces.GPU somewhere in code
gpu_hidden_state = gr.State(value=None)
@spaces.GPU
def _gpu_dummy():
return None
gpu_hidden_state.change(fn=_gpu_dummy, inputs=[gpu_hidden_state], outputs=[gpu_hidden_state])
# ── model switcher ────────────────────────────────────────────────
def on_model_change(variant):
if not variant:
return "⚠️ No model selected"
try:
onnx_path = download_model_variant(variant)
predictor = Predictor(onnx_path)
state["predictor"] = predictor
def new_predict_fn(image: Image.Image) -> list[tuple[str, float]]:
tensor = preprocess(image, predictor.image_size)
logits = predictor.run(tensor)[0]
inv = {v: k for k, v in predictor.tag_to_id.items()}
indices = np.argsort(logits)[::-1]
return [(inv[int(i)], float(logits[i])) for i in indices]
state["predict_fn"] = new_predict_fn
state["all_logits"] = None
state["tag_metadata"] = None
return f"✅ Switched to {variant} ({predictor.num_classes} tags)"
except Exception as exc:
return f"❌ Failed to load model: {exc}"
model_dropdown.change(
fn=on_model_change,
inputs=[model_dropdown],
outputs=[model_status],
)
return app
# ── helpers ────────────────────────────────────────────────────────────────
def enrich_tags(
tags_scores: list[tuple[str, float]], cat_map: dict[str, int]
) -> dict[str, dict]:
result: dict[str, dict] = {}
for tag, score in tags_scores:
cat_id = cat_map.get(tag, 0)
result[tag] = {
"score": score,
"category": cat_id,
"category_name": CATEGORY_MAP.get(cat_id, "general"),
}
return result
def format_tag(tag: str, use_underscore: bool) -> str:
return tag if use_underscore else tag.replace("_", " ")
# ── main ────────────────────────────────────────────────────────────────────
def main():
model_arg = sys.argv[1] if len(sys.argv) > 1 else None
model_env = os.environ.get("MODEL_DIR")
model_variants: list[str] = []
initial_predict_fn = None
if model_arg:
onnx = Path(model_arg)
if not onnx.exists():
print(f"Error: {onnx} not found", file=sys.stderr)
sys.exit(1)
print(f"Loading local model: {onnx}")
predictor = Predictor(onnx)
def _predict(image: Image.Image) -> list[tuple[str, float]]:
tensor = preprocess(image, predictor.image_size)
logits = predictor.run(tensor)[0]
inv = {v: k for k, v in predictor.tag_to_id.items()}
indices = np.argsort(logits)[::-1]
return [(inv[int(i)], float(logits[i])) for i in indices]
initial_predict_fn = _predict
elif model_env:
env_dir = Path(model_env)
if not env_dir.is_dir():
print(f"Error: MODEL_DIR {env_dir} is not a directory", file=sys.stderr)
sys.exit(1)
onnx_files = list(env_dir.glob("*.onnx"))
if not onnx_files:
print(f"Error: no .onnx files in {env_dir}", file=sys.stderr)
sys.exit(1)
onnx = onnx_files[0]
print(f"Loading local model from MODEL_DIR: {onnx}")
predictor = Predictor(onnx)
def _predict(image: Image.Image) -> list[tuple[str, float]]:
tensor = preprocess(image, predictor.image_size)
logits = predictor.run(tensor)[0]
inv = {v: k for k, v in predictor.tag_to_id.items()}
indices = np.argsort(logits)[::-1]
return [(inv[int(i)], float(logits[i])) for i in indices]
initial_predict_fn = _predict
else:
print("Discovering model variants on HF hub …")
model_variants = discover_model_variants()
if not model_variants:
print("Warning: no models found on hub.")
else:
print(f"Found variants: {model_variants}")
default = model_variants[0]
print(f"Downloading default model: {default} …")
try:
onnx_path = download_model_variant(default)
predictor = Predictor(onnx_path)
def _predict(image: Image.Image) -> list[tuple[str, float]]:
tensor = preprocess(image, predictor.image_size)
logits = predictor.run(tensor)[0]
inv = {v: k for k, v in predictor.tag_to_id.items()}
indices = np.argsort(logits)[::-1]
return [(inv[int(i)], float(logits[i])) for i in indices]
initial_predict_fn = _predict
print(f"Loaded {default} ({predictor.num_classes} tags)")
except Exception as exc:
print(f"Error loading default model: {exc}")
app = build_app(initial_predict_fn, model_variants)
host = os.environ.get("GRADIO_SERVER_NAME") or os.environ.get("HOST")
port_str = os.environ.get("GRADIO_SERVER_PORT") or os.environ.get("PORT")
port = int(port_str) if port_str else None
app.launch(server_name=host, server_port=port, ssr_mode=False)
if __name__ == "__main__":
main()