#!/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 = ( '' '' '' '' '' '' '' + "".join(rows) + '
LinkTagScoreCategory
' ) 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"" f"{rows}
TagScoreCategory
") 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()