| import ast |
| import csv |
| import json |
| import logging |
| import time |
| from pathlib import Path |
| from typing import Dict, Tuple |
|
|
| import gradio as gr |
| import numpy as np |
| from fetch_url_util import fetch_image_from_url |
| from huggingface_hub import hf_hub_download |
| from PIL import Image |
|
|
| |
| logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(name)s - %(message)s") |
| logger = logging.getLogger("PixAITagger") |
|
|
| logging.getLogger().setLevel(logging.WARNING) |
| logging.getLogger("PixAITagger").setLevel(logging.INFO) |
|
|
| |
| MODEL_REPO = "deepghs/pixai-tagger-v0.9-onnx" |
| MODEL_FILENAME = "model.onnx" |
| TAGS_FILENAME = "selected_tags.csv" |
| INPUT_SIZE = 448 |
|
|
|
|
| class HybridEngine: |
| """ |
| Handles inference with dynamic backend selection. |
| Priority: OpenVINO INT8 -> OpenVINO FP32 -> ONNX Runtime |
| """ |
|
|
| def __init__(self, model_path: str, use_int8: bool = True): |
| self.model_path = model_path |
| self.session = None |
| self.use_openvino = False |
| self.provider_name = None |
| self.use_int8 = use_int8 |
|
|
| self._init_backend() |
|
|
| def _init_backend(self): |
| |
| try: |
| import openvino as ov |
| logger.info("Engine: OpenVINO available, reading model...") |
| core = ov.Core() |
| model = core.read_model(self.model_path) |
| |
| |
| if self.use_int8: |
| try: |
| import nncf |
| logger.info("Engine: Compressing weights to INT8 using NNCF...") |
| model = nncf.compress_weights(model) |
| self.provider_name = "OpenVINO (INT8 Weights)" |
| except ImportError: |
| logger.warning("Engine: NNCF not installed. Falling back to FP32.") |
| self.provider_name = "OpenVINO (FP32 - NNCF missing)" |
| except Exception as e: |
| logger.warning(f"Engine: NNCF Compression failed ({e}). Falling back to FP32.") |
| self.provider_name = "OpenVINO (FP32 - Compression error)" |
| else: |
| self.provider_name = "OpenVINO (FP32)" |
|
|
| |
| self.session = core.compile_model(model, "CPU") |
| self.use_openvino = True |
| logger.info(f"Engine: Success using {self.provider_name}") |
| return |
|
|
| except Exception as e: |
| logger.warning(f"Engine: OpenVINO initialization failed ({e}). Falling back to ONNX Runtime.") |
|
|
| |
| try: |
| import onnxruntime as ort |
| sess_options = ort.SessionOptions() |
| sess_options.log_severity_level = 3 |
|
|
| self.session = ort.InferenceSession( |
| self.model_path, |
| sess_options=sess_options, |
| providers=["CPUExecutionProvider"], |
| ) |
| self.use_openvino = False |
| self.provider_name = f"ONNX Runtime ({self.session.get_providers()[0]})" |
| logger.info("Engine: Using ONNX Runtime backend") |
| except Exception as e: |
| logger.error(f"Engine: FATAL - ONNX Runtime also failed: {e}") |
| self.provider_name = "Error: No backend available" |
| self.session = None |
|
|
| def run(self, input_data: np.ndarray, expected_dim: int): |
| if not self.session: |
| raise RuntimeError("Engine not initialized") |
|
|
| if self.use_openvino: |
| results = self.session(input_data) |
| outputs = list(results.values()) |
| else: |
| input_name = self.session.get_inputs()[0].name |
| outputs = self.session.run(None, {input_name: input_data}) |
|
|
| |
| for out in outputs: |
| if out.shape[1] == expected_dim: |
| return out[0], self.provider_name |
|
|
| |
| out = max(outputs, key=lambda x: x.shape[1]) |
| return out[0], self.provider_name |
|
|
|
|
| class PixAITagger: |
| def __init__(self): |
| self.model_path = None |
| self.tags_list = [] |
| self._load_resources() |
| |
| |
| self.engine = None |
| self.current_int8_mode = None |
|
|
| def _load_resources(self): |
| logger.info(f"Downloading resources from {MODEL_REPO}...") |
| self.model_path = hf_hub_download(repo_id=MODEL_REPO, filename=MODEL_FILENAME) |
|
|
| try: |
| tags_path = hf_hub_download(repo_id=MODEL_REPO, filename=TAGS_FILENAME) |
| self._load_tags_csv(Path(tags_path)) |
| except Exception as e: |
| raise FileNotFoundError(f"Could not load tags file: {e}") |
|
|
| def _load_tags_csv(self, csv_path: Path): |
| self.tags_list = [] |
| with csv_path.open("r", encoding="utf-8") as f: |
| reader = csv.DictReader(f) |
| for row in reader: |
| try: |
| idx = int(row.get("id")) |
| name = row.get("name") |
| category = int(row.get("category", 0)) |
| ips_raw = row.get("ips", "[]") |
| ips = [] |
| if ips_raw and ips_raw != "[]": |
| try: |
| ips = json.loads(ips_raw) |
| except Exception: |
| try: |
| ips = ast.literal_eval(ips_raw) |
| except Exception: |
| pass |
| self.tags_list.append( |
| {"id": idx, "name": name, "is_char": category == 4, "ips": [str(ip) for ip in ips]} |
| ) |
| except ValueError: |
| continue |
|
|
| self.tags_list.sort(key=lambda x: x["id"]) |
|
|
| self.id_to_tag = {} |
| self.char_indices = [] |
| self.gen_indices = [] |
| self.mapping = {} |
|
|
| for item in self.tags_list: |
| idx = item["id"] |
| name = item["name"] |
| self.id_to_tag[idx] = name |
|
|
| if item["is_char"]: |
| self.char_indices.append(idx) |
| if item["ips"]: |
| self.mapping[name] = item["ips"] |
| else: |
| self.gen_indices.append(idx) |
|
|
| self.num_classes = len(self.tags_list) |
| logger.info(f"Loaded {self.num_classes} tags.") |
|
|
| def preprocess(self, image: Image.Image) -> np.ndarray: |
| if image.mode != "RGB": |
| image = image.convert("RGB") |
| image = image.resize((INPUT_SIZE, INPUT_SIZE), Image.BICUBIC) |
| img = np.array(image).astype(np.float32) / 255.0 |
| img = (img - 0.5) / 0.5 |
| img = img.transpose(2, 0, 1) |
| return np.expand_dims(img, 0) |
|
|
| def predict( |
| self, image: Image.Image, gen_threshold: float, char_threshold: float, resolve_mapping: bool, use_int8_weights: bool |
| ) -> Tuple[Dict, Dict, str, str, float]: |
| |
| |
| if self.engine is None or self.current_int8_mode != use_int8_weights: |
| logger.info(f"Reloading engine. New mode INT8: {use_int8_weights}") |
| self.engine = HybridEngine(str(self.model_path), use_int8=use_int8_weights) |
| self.current_int8_mode = use_int8_weights |
|
|
| input_tensor = self.preprocess(image) |
| |
| infer_start = time.time() |
| logits, provider_name = self.engine.run(input_tensor, self.num_classes) |
| infer_time = time.time() - infer_start |
| |
| |
| probs = 1 / (1 + np.exp(-logits)) |
| |
| |
| gen_tags = {} |
| for idx in self.gen_indices: |
| if idx < len(probs): |
| score = float(probs[idx]) |
| if score >= gen_threshold: |
| gen_tags[self.id_to_tag[idx]] = score |
| |
| |
| char_tags = {} |
| detected_ips = set() |
| |
| for idx in self.char_indices: |
| if idx < len(probs): |
| score = float(probs[idx]) |
| if score >= char_threshold: |
| name = self.id_to_tag[idx] |
| char_tags[name] = score |
| |
| if resolve_mapping and name in self.mapping: |
| for ip in self.mapping[name]: |
| detected_ips.add(ip) |
| |
| gen_tags = dict(sorted(gen_tags.items(), key=lambda x: x[1], reverse=True)) |
| char_tags = dict(sorted(char_tags.items(), key=lambda x: x[1], reverse=True)) |
| |
| ip_text = ", ".join(sorted(list(detected_ips))) if detected_ips else "" |
| |
| return gen_tags, char_tags, ip_text, provider_name, infer_time |
|
|
|
|
| |
|
|
| tagger_instance = None |
|
|
| def get_tagger(): |
| global tagger_instance |
| if tagger_instance is None: |
| tagger_instance = PixAITagger() |
| return tagger_instance |
|
|
| def init_app(): |
| """Warms up the model loader.""" |
| get_tagger() |
| return None |
|
|
| def run_inference(image, gen_thresh, char_thresh, resolve_mapping, use_int8): |
| if image is None: |
| return "", "", "", {}, {}, "" |
| try: |
| model = get_tagger() |
| start_time = time.time() |
|
|
| gen_tags, char_tags, ip_str, provider, infer_time = model.predict( |
| image, gen_thresh, char_thresh, resolve_mapping, use_int8 |
| ) |
|
|
| char_str = ", ".join(char_tags.keys()).replace("_", " ") |
| gen_str = ", ".join(gen_tags.keys()).replace("_", " ") |
|
|
| if not resolve_mapping: |
| ip_disp = "" |
| elif not ip_str: |
| ip_disp = "" |
| else: |
| ip_disp = ip_str.replace("_", " ") |
|
|
| time_disp = f"- **Provider:** {provider} | **Inference time:** {infer_time:.4f}s" |
| return char_str, ip_disp, gen_str, char_tags, gen_tags, time_disp |
| except Exception as e: |
| logger.error(f"Inference Error: {e}") |
| return "", f"Error: {str(e)}", "", {}, {}, f"Error: {str(e)}" |
|
|
|
|
| with gr.Blocks(title="PixAI Tagger v0.9 ONNX") as demo: |
|
|
| gr.Markdown( |
| '<a href="https://huggingface.co/pixai-labs/pixai-tagger-v0.9" ' |
| 'target="_blank" rel="noopener noreferrer">PixAI Tagger</a>' |
| ' is an iteration on top of ' |
| '<a href="https://huggingface.co/SmilingWolf/wd-eva02-large-tagger-v3" ' |
| 'target="_blank" rel="noopener noreferrer">SmilingWolf/wd-eva02-large-tagger-v3</a>' |
| ' with an updated dataset (2025-01). \n' |
| 'It should be noted that PixAI Tagger may be worse in accuracy over eva02 large. See the PixAI page for details' |
| ) |
|
|
| |
| with gr.Row(elem_classes=["container"]): |
| with gr.Column(scale=1, elem_classes=["header-col"]): |
| gr.Markdown("### Input") |
| with gr.Column(scale=1, elem_classes=["header-col"]): |
| gr.Markdown("### Configuration & Results") |
|
|
| with gr.Row(elem_classes=["container"]): |
| |
| with gr.Column(scale=1): |
| url_input = gr.Textbox( |
| label="Enter Image URL (not all may work) or upload an image below", |
| placeholder="https://example.com/image.jpg", |
| ) |
| input_img = gr.Image(type="pil", label="", show_label=False, elem_classes=["image-container"]) |
|
|
| |
| with gr.Column(scale=1): |
| |
| with gr.Row(): |
| run_btn = gr.Button("Run (Image Upload)", variant="primary") |
| url_btn = gr.Button("Run (URL)", variant="secondary") |
|
|
| |
| with gr.Group(): |
| with gr.Row(elem_id="d-row-container"): |
| char_box = gr.Textbox(label="Character Tags", interactive=False, buttons=["copy"]) |
| ip_box = gr.Textbox(label="Character - Copyright Mapping", interactive=False, buttons=["copy"]) |
|
|
| gen_box = gr.Textbox(label="General Tags", interactive=False, buttons=["copy"]) |
| time_info = gr.Markdown(elem_id="time-display") |
|
|
| |
| with gr.Group(): |
| with gr.Row(elem_id="d-row-container"): |
| char_slider = gr.Slider(0.0, 1.0, value=0.75, step=0.05, label="Character Threshold") |
| gen_slider = gr.Slider(0.0, 1.0, value=0.30, step=0.05, label="General Threshold") |
| |
| |
| with gr.Row(): |
| map_checkbox = gr.Checkbox(value=True, label="Resolve Copyright Mapping") |
| int8_checkbox = gr.Checkbox(value=True, label="INT8 Weights") |
|
|
| |
| gr.Markdown("### Confidence Scores") |
| char_plot = gr.Label(label="Character Probabilities", num_top_classes=50) |
| gen_plot = gr.Label(label="General Probabilities", num_top_classes=500) |
|
|
| gr.Markdown( |
| 'Model sourced from ' |
| '<a href="https://huggingface.co/deepghs/pixai-tagger-v0.9-onnx" ' |
| 'target="_blank" rel="noopener noreferrer">deepghs/pixai-tagger-v0.9-onnx</a>. \n' |
| 'OpenVINO™ will be used to accelerate CPU inference with ONNX CPUExecutionProvider as fallback. \n' |
| 'INT8 weights option may improve inference times with a deviation of roughly +-0.005 in scores.' |
| ) |
|
|
| |
| |
| url_btn.click(fn=fetch_image_from_url, inputs=[url_input], outputs=[input_img]).then( |
| fn=run_inference, |
| inputs=[input_img, gen_slider, char_slider, map_checkbox, int8_checkbox], |
| outputs=[char_box, ip_box, gen_box, char_plot, gen_plot, time_info], |
| ) |
|
|
| run_btn.click( |
| fn=run_inference, |
| inputs=[input_img, gen_slider, char_slider, map_checkbox, int8_checkbox], |
| outputs=[char_box, ip_box, gen_box, char_plot, gen_plot, time_info], |
| ) |
|
|
| |
| demo.load(fn=init_app) |
|
|
| if __name__ == "__main__": |
| demo.launch( |
| theme=gr.themes.Base(), |
| css=""" |
| * { box-sizing: border-box; } |
| @media (max-width: 1022px) { |
| #d-row-container { |
| flex-direction: column !important; |
| } |
| #d-row-container > * { |
| width: 100% !important; |
| } |
| #d-row-container .block { |
| width: 100% !important; |
| } |
| }""", |
| ) |