Spaces:
Running on Zero
Running on Zero
| #!/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)) | |
| 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 "<i>No results yet.</i>", "" | |
| 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 "<i>No tags pass the filters.</i>", "" | |
| rows = [] | |
| for tag, score in items: | |
| m = meta[tag] | |
| link = ( | |
| f'<a href="https://danbooru.donmai.us/posts?tags={tag}"' | |
| f' target="_blank">{tag}</a>' | |
| ) | |
| display = format_tag(tag, _use_underscore) | |
| rows.append( | |
| f"<tr>" | |
| f"<td>{link}</td>" | |
| f"<td>{display}</td>" | |
| f"<td style='text-align:right'>{score:.4f}</td>" | |
| f"<td><code>{m['category_name']}</code></td>" | |
| f"</tr>" | |
| ) | |
| table = ( | |
| '<table style="width:100%">' | |
| '<thead><tr>' | |
| '<th>Link</th><th>Tag</th>' | |
| '<th style="text-align:right">Score</th>' | |
| '<th>Category</th>' | |
| '</tr></thead>' | |
| '<tbody>' + "".join(rows) + '</tbody></table>' | |
| ) | |
| 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 "<i>Error loading URL.</i>", "", f"β {exc}" | |
| state["current_image"] = pil | |
| fn = state["predict_fn"] | |
| if fn is None: | |
| return "<i>No model loaded.</i>", "", "β 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, "", "<i>No results yet.</i>", "Cleared.", "", "", "<i>No results yet.</i>" | |
| 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 "<i>No matching tags.</i>" | |
| rows = "".join( | |
| f"<tr><td>{t}</td><td>{s:.4f}</td>" | |
| f"<td><code>{meta[t]['category_name']}</code></td></tr>" | |
| for t, s in matches | |
| ) | |
| return (f"<table style='width:100%'>" | |
| f"<tr><th>Tag</th><th>Score</th><th>Category</th></tr>" | |
| f"{rows}</table>") | |
| 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) | |
| 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() | |