import os import csv import json import ast import time import logging from pathlib import Path from typing import Dict, List, Tuple, Any, Optional import numpy as np import gradio as gr from PIL import Image from huggingface_hub import hf_hub_download # Configure Logging logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(message)s') logger = logging.getLogger("PixAITagger") # Constants MODEL_REPO = "deepghs/pixai-tagger-v0.9-onnx" MODEL_FILENAME = "model.onnx" TAGS_FILENAME = "selected_tags.csv" INPUT_SIZE = 448 class HybridEngine: """ Handles inference, allowing dynamic switching between OpenVINO and ONNX Runtime. """ def __init__(self, model_path: str): self.model_path = model_path self.ov_session = None self.ort_session = None self.ov_available = False # Check OpenVINO availability once try: import openvino as ov self.ov_available = True except ImportError: self.ov_available = False def _get_ort_session(self): """Lazy load ONNX Runtime session""" if self.ort_session is None: import onnxruntime as ort sess_options = ort.SessionOptions() sess_options.log_severity_level = 3 logger.info("Engine: Initializing ONNX Runtime...") self.ort_session = ort.InferenceSession(self.model_path, sess_options=sess_options, providers=["CPUExecutionProvider"]) return self.ort_session def _get_ov_session(self): """Lazy load OpenVINO session""" if not self.ov_available: raise ImportError("OpenVINO not installed") if self.ov_session is None: import openvino as ov core = ov.Core() # ov.log.set_level(ov.log.Level.ERR) logger.info("Engine: Compiling OpenVINO model...") model_ov = core.read_model(self.model_path) self.ov_session = core.compile_model(model_ov, "CPU") return self.ov_session def run(self, input_data: np.ndarray, expected_dim: int) -> Tuple[np.ndarray, str]: """ Runs inference. Returns (logits, provider_name). Tries OpenVINO first, falls back to ONNX Runtime if needed. """ # Try OpenVINO if available if self.ov_available: try: sess = self._get_ov_session() # OpenVINO inference request = sess.create_infer_request() results = request.infer(input_data) # Find output with matching shape output_tensor = None for res_data in results.values(): if res_data.shape[1] == expected_dim: output_tensor = res_data[0] break if output_tensor is None: # Fallback to largest output output_tensor = max(results.values(), key=lambda x: x.shape[1])[0] return output_tensor, "OpenVINO (CPU)" except Exception as e: logger.warning(f"OpenVINO execution failed: {e}. Falling back to ONNX Runtime.") # Fall through to ORT # ONNX Runtime Fallback sess = self._get_ort_session() input_name = sess.get_inputs()[0].name outputs = sess.run(None, {input_name: input_data}) output_tensor = None for out in outputs: if out.shape[1] == expected_dim: output_tensor = out[0] break if output_tensor is None: output_tensor = max(outputs, key=lambda x: x.shape[1])[0] provider = sess.get_providers()[0] return output_tensor, f"ONNX Runtime ({provider})" class PixAITagger: def __init__(self): self.model_path = None self.tags_list = [] self._load_resources() self.engine = HybridEngine(str(self.model_path)) 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: try: ips = ast.literal_eval(ips_raw) except: 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) -> Tuple[Dict, Dict, str, str]: input_tensor = self.preprocess(image) # Run inference - always try OpenVINO first, fallback to ONNX logits, provider_name = self.engine.run(input_tensor, self.num_classes) # Sigmoid probs = 1 / (1 + np.exp(-logits)) # General Tags 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 # Character Tags & IPs 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 # --- UI Setup --- 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 and prevents Gradio return-value warnings.""" get_tagger() return None def run_inference(image, gen_thresh, char_thresh, resolve_mapping): if image is None: return "", "", "", {}, {}, "" try: model = get_tagger() start_time = time.time() # Internal predict handles OpenVINO -> ONNX fallback gen_tags, char_tags, ip_str, provider = model.predict(image, gen_thresh, char_thresh, resolve_mapping) taken = time.time() - start_time char_str = ", ".join(char_tags.keys()).replace("_", " ") gen_str = ", ".join(gen_tags.keys()).replace("_", " ") if not resolve_mapping: ip_disp = "Mapping Disabled" elif not ip_str: ip_disp = "No specific copyright detected" else: ip_disp = ip_str.replace("_", " ") time_disp = f"- **Provider:** {provider} | **Time taken:** {taken:.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)}" css = """ .container { max-width: 1200px; margin: 0 auto; } /* Align headers by removing default top margins */ .header-col h3 { margin-top: 0 !important; margin-bottom: 10px !important; } .image-container img { max-height: 80dvh !important; width: auto !important; margin: 0 auto; object-fit: contain; } #time-display { margin-top: 10px; padding-top: 5px; border-top: 1px solid rgba(128,128,128,0.1); font-size: 0.85em; opacity: 0.8; } footer { visibility: hidden; } """ with gr.Blocks(title="PixAI Tagger v0.9") as demo: # Header Row - Ensures Input and Results start at the exact same height with gr.Row(elem_classes=["container"]): with gr.Column(scale=1, elem_classes=["header-col"]): gr.Markdown("### Input Image") with gr.Column(scale=1, elem_classes=["header-col"]): gr.Markdown("### Results") with gr.Row(elem_classes=["container"]): # LEFT COLUMN with gr.Column(scale=1): input_img = gr.Image( type="pil", label="", show_label=False, elem_classes=["image-container"] ) with gr.Group(): run_btn = gr.Button("Analyze Image", variant="primary") gr.Markdown("### Configuration") with gr.Group(): char_slider = gr.Slider(0.0, 1.0, value=0.85, step=0.05, label="Character Threshold") gen_slider = gr.Slider(0.0, 1.0, value=0.30, step=0.05, label="General Threshold") map_checkbox = gr.Checkbox(value=True, label="Resolve Copyright Mapping") # RIGHT COLUMN with gr.Column(scale=1): # Combined Output Group with gr.Group(): # In Gradio 6.3+, 'buttons' is the standard way to add copy functionality 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 display embedded at the bottom of the group time_info = gr.Markdown(elem_id="time-display") gr.Markdown("### Confidence") char_plot = gr.Label(label="Character Probabilities", num_top_classes=50) gen_plot = gr.Label(label="General Probabilities", num_top_classes=50) # Click Logic run_btn.click( fn=run_inference, inputs=[input_img, gen_slider, char_slider, map_checkbox], outputs=[char_box, ip_box, gen_box, char_plot, gen_plot, time_info] ) # Load logic demo.load(fn=init_app) if __name__ == "__main__": # Gradio 6.x migration requires theme and css in launch() demo.launch( theme=gr.themes.Base(), css=css, api_open=False )