Spaces:
Running
Running
Download app.py from IdlecloudX/Image_Inversion: direct link, hf CLI and curl.
- Browser
- Download file 32 kB
-
https://huggingface.co/spaces/IdlecloudX/Image_Inversion/resolve/main/app.py
- Command line
-
hf download hf://spaces/IdlecloudX/Image_Inversion/app.py
-
curl -L -o app.py https://huggingface.co/spaces/IdlecloudX/Image_Inversion/resolve/main/app.py
32 kB
| import os | |
| # 禁止 Hub 依赖隐式读取私有 token;受限模型只在下方下载调用中显式使用 Space Secret。 | |
| os.environ.setdefault("HF_HUB_DISABLE_IMPLICIT_TOKEN", "1") | |
| import json | |
| import time | |
| import threading | |
| import warnings | |
| from html import escape | |
| from pathlib import Path | |
| from typing import Optional | |
| import gradio as gr | |
| from huggingface_hub import snapshot_download | |
| from PIL import Image, ImageFile | |
| from handler import EndpointHandler | |
| from image_inversion_job_api import ( | |
| ImageInversionJobAPI, | |
| ImageInversionJobRequest, | |
| ImageInversionJobSettings, | |
| create_job_api_lifespan, | |
| ) | |
| from translator import translate_texts | |
| # ------------------------------------------------------------------ | |
| # 安全配置 | |
| # ------------------------------------------------------------------ | |
| # 1) 限制上传文件原始体积,拦截伪装图片/图片中塞入额外数据/高熵噪声导致的超大文件 | |
| MAX_UPLOAD_BYTES = 8 * 1024 * 1024 # 8 MB | |
| # 2) 限制单边尺寸,避免异常超大分辨率 | |
| MAX_IMAGE_SIDE = 4096 | |
| # 3) 限制总像素数,防止“像素炸弹”或解码后内存占用过高 | |
| MAX_IMAGE_PIXELS = 20_000_000 # 2000 万像素 | |
| # 4) 限制解码后的估算内存占用 | |
| MAX_DECOMPRESSED_BYTES = 160 * 1024 * 1024 # 160 MB | |
| # 5) 仅允许常见安全图片格式 | |
| ALLOWED_IMAGE_FORMATS = {"PNG", "JPEG", "WEBP", "BMP", "GIF"} | |
| # Pillow 安全设置 | |
| Image.MAX_IMAGE_PIXELS = MAX_IMAGE_PIXELS | |
| ImageFile.LOAD_TRUNCATED_IMAGES = False | |
| warnings.simplefilter("error", Image.DecompressionBombWarning) | |
| class ImageValidationError(ValueError): | |
| """上传图片校验失败。""" | |
| def _format_size(num_bytes: int) -> str: | |
| if num_bytes < 1024: | |
| return f"{num_bytes} B" | |
| if num_bytes < 1024 * 1024: | |
| return f"{num_bytes / 1024:.2f} KB" | |
| return f"{num_bytes / (1024 * 1024):.2f} MB" | |
| def validate_and_open_image(image_path: str) -> Image.Image: | |
| """ | |
| 安全打开用户上传图片: | |
| - 校验原始文件体积 | |
| - 校验图片格式 | |
| - 校验宽高/总像素 | |
| - 校验解码后预估内存占用 | |
| - 拦截 Pillow 解压炸弹警告 | |
| """ | |
| if not image_path: | |
| raise ImageValidationError("未检测到上传文件。") | |
| if not os.path.isfile(image_path): | |
| raise ImageValidationError("上传文件不存在或无法访问。") | |
| file_size = os.path.getsize(image_path) | |
| if file_size <= 0: | |
| raise ImageValidationError("上传文件为空。") | |
| if file_size > MAX_UPLOAD_BYTES: | |
| raise ImageValidationError( | |
| f"图片文件过大:{_format_size(file_size)},超过限制 {_format_size(MAX_UPLOAD_BYTES)}。" | |
| ) | |
| try: | |
| with Image.open(image_path) as probe: | |
| img_format = (probe.format or "").upper() | |
| width, height = probe.size | |
| probe.verify() | |
| except Image.DecompressionBombWarning: | |
| raise ImageValidationError("图片疑似像素炸弹,已被拒绝处理。") | |
| except Exception as e: | |
| raise ImageValidationError(f"无法解析为有效图片文件:{e}") | |
| if img_format not in ALLOWED_IMAGE_FORMATS: | |
| raise ImageValidationError( | |
| f"不支持的图片格式:{img_format or '未知'}。仅允许:{', '.join(sorted(ALLOWED_IMAGE_FORMATS))}。" | |
| ) | |
| if width <= 0 or height <= 0: | |
| raise ImageValidationError("图片尺寸非法。") | |
| if width > MAX_IMAGE_SIDE or height > MAX_IMAGE_SIDE: | |
| raise ImageValidationError( | |
| f"图片尺寸过大:{width}×{height},单边不得超过 {MAX_IMAGE_SIDE} 像素。" | |
| ) | |
| total_pixels = width * height | |
| if total_pixels > MAX_IMAGE_PIXELS: | |
| raise ImageValidationError( | |
| f"图片总像素过大:{total_pixels:,},超过限制 {MAX_IMAGE_PIXELS:,}。" | |
| ) | |
| estimated_decompressed_bytes = total_pixels * 3 | |
| if estimated_decompressed_bytes > MAX_DECOMPRESSED_BYTES: | |
| raise ImageValidationError( | |
| "图片解码后的内存占用过高,已拒绝处理。" | |
| f" 预计占用约 {_format_size(estimated_decompressed_bytes)}," | |
| f"超过限制 {_format_size(MAX_DECOMPRESSED_BYTES)}。" | |
| ) | |
| try: | |
| with Image.open(image_path) as img: | |
| img.load() | |
| # 保留透明通道,让模型官方预处理统一合成白底。 | |
| img = img.copy() | |
| except Image.DecompressionBombWarning: | |
| raise ImageValidationError("图片在解码阶段触发像素炸弹保护,已拒绝处理。") | |
| except Exception as e: | |
| raise ImageValidationError(f"图片加载失败:{e}") | |
| return img | |
| # ------------------------------------------------------------------ | |
| # PixAI Tagger v1.0 模型配置 | |
| # ------------------------------------------------------------------ | |
| ASSETS_REPO_ID = os.environ.get("ASSETS_REPO_ID", "pixai-labs/pixai-tagger-v1.0") | |
| # 固定模型权重和自定义推理代码,避免上游 main 更新后混用不同版本。 | |
| ASSETS_REVISION = os.environ.get("ASSETS_REVISION", "9fe10addf9326e292da8a85a98ea74cd91b41771") | |
| ASSETS_HF_TOKEN = os.environ.get("HF_TOKEN") | |
| MODEL_DIR = os.environ.get("MODEL_DIR", "./assets-v1.0") | |
| REQUIRED_FILES = [ | |
| "config.json", | |
| "preprocessor_config.json", | |
| "tagger_pipeline.py", | |
| "model.safetensors", | |
| ] | |
| def ensure_assets(repo_id: str, revision: Optional[str], target_dir: str) -> str: | |
| """下载同一版本的模型、配置和官方推理代码,返回本地快照路径。 | |
| Args: | |
| repo_id: Hugging Face 模型仓库 ID。 | |
| revision: 可选模型版本;为空时使用仓库默认版本。 | |
| target_dir: Hugging Face 模型缓存目录。 | |
| Returns: | |
| 已完整下载的模型快照目录。 | |
| """ | |
| target = Path(target_dir) | |
| target.mkdir(parents=True, exist_ok=True) | |
| snapshot_path = snapshot_download( | |
| repo_id=repo_id, | |
| revision=revision, | |
| allow_patterns=REQUIRED_FILES, | |
| cache_dir=str(target), | |
| # token 仅用于 Hub 下载;模型加载与翻译请求均不传递此凭据。 | |
| token=ASSETS_HF_TOKEN or False, | |
| ) | |
| for fname in REQUIRED_FILES: | |
| src = Path(snapshot_path) / fname | |
| if not src.exists(): | |
| raise FileNotFoundError( | |
| f"模型资源缺失:'{fname}' 未在 {repo_id} @ {revision or 'default'} 中找到。" | |
| ) | |
| return snapshot_path | |
| # ------------------------------------------------------------------ | |
| # Tagger 类:使用新版 EndpointHandler | |
| # ------------------------------------------------------------------ | |
| class Tagger: | |
| def __init__(self): | |
| self.handler = None | |
| self.device = "unknown" | |
| self._load_model_and_labels() | |
| def _load_model_and_labels(self) -> None: | |
| try: | |
| snapshot_path = ensure_assets(ASSETS_REPO_ID, ASSETS_REVISION, MODEL_DIR) | |
| self.handler = EndpointHandler(snapshot_path) | |
| self.device = getattr(self.handler, "device", "unknown") | |
| print(f"✅ PixAI Tagger v1.0 加载成功,设备:{str(self.device).upper()}") | |
| except Exception as e: | |
| print(f"❌ PixAI Tagger v1.0 加载失败: {e}") | |
| raise RuntimeError(f"模型初始化失败: {e}") from e | |
| def _display_tag(tag: str) -> str: | |
| return str(tag).replace("_", " ") | |
| def predict(self, img: Image.Image, gen_th: float = 0.17, char_th: float = 0.27): | |
| """运行 v1.0 反推并生成页面和任务 API 共用的标签结果。 | |
| Args: | |
| img: 已校验的 PIL 图片,透明背景由官方预处理处理。 | |
| gen_th: 通用标签的置信度阈值。 | |
| char_th: 角色标签的置信度阈值。 | |
| Returns: | |
| 三类展示标签、对应翻译顺序,以及包含六类原始标签的元数据。 | |
| """ | |
| if self.handler is None: | |
| raise RuntimeError("模型未成功加载,无法进行预测。") | |
| if img is None: | |
| raise ValueError("输入图像不能为空。") | |
| params = { | |
| "general_threshold": float(gen_th), | |
| "character_threshold": float(char_th), | |
| "mode": "threshold", | |
| "topk_general": 25, | |
| "topk_character": 10, | |
| "include_scores": True, | |
| } | |
| data = { | |
| "inputs": img, | |
| "parameters": params, | |
| } | |
| started = time.time() | |
| out = self.handler(data) | |
| latency = round(time.time() - started, 4) | |
| # 保持页面和既有任务 API 的三个展示类别,IP 使用 v1.0 的真实置信度。 | |
| res = {} | |
| for display_category, model_category in ( | |
| ("general", "feature"), ("characters", "character"), ("ips", "ip") | |
| ): | |
| res[display_category] = { | |
| self._display_tag(tag): score | |
| for tag, score in sorted( | |
| out[f"{model_category}_scores"].items(), | |
| key=lambda item: item[1], reverse=True, | |
| ) | |
| } | |
| tag_categories_for_translation = { | |
| category: list(tags) for category, tags in res.items() | |
| } | |
| raw_meta = { | |
| "model": ASSETS_REPO_ID, | |
| "revision": ASSETS_REVISION, | |
| "device": str(self.device), | |
| "latency_s_total": latency, | |
| "_params": out["_params"], | |
| "_timings": out["_timings"], | |
| "tags": { | |
| "general": out["feature_scores"], | |
| "character": out["character_scores"], | |
| "copyright": out["ip_scores"], | |
| "style": out["style_scores"], | |
| "meta": out["meta_scores"], | |
| "rating": out["rating_scores"], | |
| }, | |
| } | |
| return res, tag_categories_for_translation, raw_meta | |
| # 全局 Tagger 实例 | |
| try: | |
| tagger_instance = Tagger() | |
| except RuntimeError as e: | |
| print(f"应用启动时 Tagger 初始化失败: {e}") | |
| tagger_instance = None | |
| DEVICE_LABEL = ( | |
| f"设备:{str(tagger_instance.device).upper()}" | |
| if tagger_instance is not None | |
| else "设备:UNKNOWN" | |
| ) | |
| # UI 与自定义任务 API 共用单推理门闩,避免 CPU 模型被并发调用导致资源争用。 | |
| INFERENCE_SLOT = threading.Lock() | |
| def predict_tags_with_shared_slot( | |
| image: Image.Image, | |
| general_threshold: float, | |
| character_threshold: float, | |
| ): | |
| """在共享单推理门闩内调用现有标签模型。 | |
| Args: | |
| image: 已完成安全校验的输入图片。 | |
| general_threshold: 通用标签阈值。 | |
| character_threshold: 角色标签阈值。 | |
| Returns: | |
| Tagger.predict 返回的标签、翻译顺序和元数据三元组。 | |
| """ | |
| if tagger_instance is None: | |
| raise RuntimeError("标签分析器尚未成功初始化。") | |
| if not INFERENCE_SLOT.acquire(blocking=False): | |
| raise RuntimeError("标签分析服务正忙,请稍后重试。") | |
| try: | |
| return tagger_instance.predict( | |
| image, | |
| general_threshold, | |
| character_threshold, | |
| ) | |
| finally: | |
| INFERENCE_SLOT.release() | |
| # ------------------------------------------------------------------ | |
| # Gradio UI | |
| # ------------------------------------------------------------------ | |
| custom_css = """ | |
| .label-container { | |
| max-height: 300px; | |
| overflow-y: auto; | |
| border: 1px solid #ddd; | |
| padding: 10px; | |
| border-radius: 5px; | |
| background-color: #f9f9f9; | |
| } | |
| .tag-item { | |
| display: flex; | |
| justify-content: space-between; | |
| align-items: center; | |
| margin: 2px 0; | |
| padding: 2px 5px; | |
| border-radius: 3px; | |
| background-color: #fff; | |
| transition: background-color 0.2s; | |
| } | |
| .tag-item:hover { | |
| background-color: #f0f0f0; | |
| } | |
| .tag-en { | |
| font-weight: bold; | |
| color: #333; | |
| cursor: pointer; | |
| } | |
| .tag-zh { | |
| color: #666; | |
| margin-left: 10px; | |
| } | |
| .tag-score { | |
| color: #999; | |
| font-size: 0.9em; | |
| white-space: nowrap; | |
| } | |
| .btn-analyze-container { | |
| margin-top: 15px; | |
| margin-bottom: 15px; | |
| } | |
| """ | |
| _js_functions = """ | |
| function copyToClipboard(text) { | |
| console.log('copyToClipboard function was called.'); | |
| console.log('Received text:', text); | |
| if (typeof text === 'undefined' || text === null) { | |
| console.warn('copyToClipboard was called with undefined or null text. Aborting this specific copy operation.'); | |
| return; | |
| } | |
| navigator.clipboard.writeText(text).then(() => { | |
| const feedback = document.createElement('div'); | |
| let displayText = String(text); | |
| displayText = displayText.substring(0, 30) + (displayText.length > 30 ? '...' : ''); | |
| feedback.textContent = '已复制: ' + displayText; | |
| feedback.style.position = 'fixed'; | |
| feedback.style.bottom = '20px'; | |
| feedback.style.left = '50%'; | |
| feedback.style.transform = 'translateX(-50%)'; | |
| feedback.style.backgroundColor = '#4CAF50'; | |
| feedback.style.color = 'white'; | |
| feedback.style.padding = '10px 20px'; | |
| feedback.style.borderRadius = '5px'; | |
| feedback.style.zIndex = '10000'; | |
| feedback.style.transition = 'opacity 0.5s ease-out'; | |
| document.body.appendChild(feedback); | |
| setTimeout(() => { | |
| feedback.style.opacity = '0'; | |
| setTimeout(() => { | |
| if (document.body.contains(feedback)) { | |
| document.body.removeChild(feedback); | |
| } | |
| }, 500); | |
| }, 1500); | |
| }).catch(err => { | |
| console.error('Failed to copy tag. Error:', err, 'Attempted to copy text:', text); | |
| const errorFeedback = document.createElement('div'); | |
| errorFeedback.textContent = '复制操作失败!'; | |
| errorFeedback.style.position = 'fixed'; | |
| errorFeedback.style.bottom = '20px'; | |
| errorFeedback.style.left = '50%'; | |
| errorFeedback.style.transform = 'translateX(-50%)'; | |
| errorFeedback.style.backgroundColor = '#D32F2F'; | |
| errorFeedback.style.color = 'white'; | |
| errorFeedback.style.padding = '10px 20px'; | |
| errorFeedback.style.borderRadius = '5px'; | |
| errorFeedback.style.zIndex = '10000'; | |
| errorFeedback.style.transition = 'opacity 0.5s ease-out'; | |
| document.body.appendChild(errorFeedback); | |
| setTimeout(() => { | |
| errorFeedback.style.opacity = '0'; | |
| setTimeout(() => { | |
| if (document.body.contains(errorFeedback)) { | |
| document.body.removeChild(errorFeedback); | |
| } | |
| }, 500); | |
| }, 2500); | |
| }); | |
| } | |
| """ | |
| with gr.Blocks(theme=gr.themes.Soft(), title="AI 图像标签分析器", css=custom_css, js=_js_functions) as demo: | |
| gr.Markdown("# 🖼️ AI 图像标签分析器") | |
| gr.Markdown( | |
| "上传图片自动识别标签,支持中英文显示和一键复制。" | |
| f"**当前模型:{ASSETS_REPO_ID}** | **{DEVICE_LABEL}**\n\n" | |
| "通用、角色和 IP 标签支持中英文展示;风格、元数据和评级标签可在下方“推理元数据”的 tags 中查看。" | |
| ) | |
| state_res = gr.State({}) | |
| state_translations_dict = gr.State({}) | |
| state_tag_categories_for_translation = gr.State({}) | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| img_in = gr.Image(type="filepath", image_mode=None, label="上传图片", height=300) | |
| btn = gr.Button("🚀 开始分析", variant="primary", elem_classes=["btn-analyze-container"]) | |
| with gr.Accordion("⚙️ 高级设置", open=False): | |
| gen_slider = gr.Slider( | |
| 0, | |
| 1, | |
| value=0.17, | |
| step=0.01, | |
| label="通用标签阈值", | |
| info="越高 → 标签更少更准", | |
| ) | |
| char_slider = gr.Slider( | |
| 0, | |
| 1, | |
| value=0.27, | |
| step=0.01, | |
| label="角色标签阈值", | |
| info="v1.0 推荐 0.27;提高阈值可减少误识别", | |
| ) | |
| show_tag_scores = gr.Checkbox( | |
| True, | |
| label="在列表中显示标签置信度", | |
| info="通用、角色和 IP 标签均显示模型置信度。", | |
| ) | |
| with gr.Accordion("📊 标签汇总设置", open=True): | |
| gr.Markdown("选择要包含在下方汇总文本框中的标签类别:") | |
| with gr.Row(): | |
| sum_general = gr.Checkbox(True, label="通用标签", min_width=50) | |
| sum_char = gr.Checkbox(True, label="角色标签", min_width=50) | |
| sum_ip = gr.Checkbox(False, label="IP 标签", min_width=50) | |
| sum_sep = gr.Dropdown(["逗号", "换行", "空格"], value="逗号", label="标签之间的分隔符") | |
| sum_show_zh = gr.Checkbox(False, label="在汇总中显示中文翻译") | |
| processing_info = gr.Markdown("", visible=False) | |
| with gr.Column(scale=2): | |
| with gr.Tabs(): | |
| with gr.TabItem("🏷️ 通用标签"): | |
| out_general = gr.HTML(label="General Tags") | |
| with gr.TabItem("👤 角色标签"): | |
| gr.Markdown("<p style='color:gray; font-size:small;'>提示:角色标签由模型推断,默认使用 v1.0 推荐阈值 0.27。</p>") | |
| out_char = gr.HTML(label="Character Tags") | |
| with gr.TabItem("🌐 IP 标签"): | |
| gr.Markdown("<p style='color:gray; font-size:small;'>提示:IP 标签来自 v1.0 的版权/作品类别,支持真实置信度。</p>") | |
| out_ip = gr.HTML(label="IP Tags") | |
| gr.Markdown("### 标签汇总结果") | |
| out_summary = gr.Textbox( | |
| label="标签汇总", | |
| placeholder="分析完成后,此处将显示汇总的英文标签...", | |
| lines=5, | |
| show_copy_button=True, | |
| ) | |
| with gr.Accordion("🧾 推理元数据", open=False): | |
| out_meta = gr.JSON(label="Metadata") | |
| # ----------------- 辅助函数 ----------------- | |
| def format_tags_html(tags_dict, translations_list, category_name, show_scores=True, show_translation_in_list=True): | |
| if not tags_dict: | |
| return "<p>暂无标签</p>" | |
| html = '<div class="label-container">' | |
| if not isinstance(translations_list, list): | |
| translations_list = [] | |
| tag_keys = list(tags_dict.keys()) | |
| for i, tag in enumerate(tag_keys): | |
| score = tags_dict[tag] | |
| safe_tag_text = escape(str(tag)) | |
| js_arg = json.dumps(str(tag), ensure_ascii=False) | |
| html += '<div class="tag-item">' | |
| tag_display_html = ( | |
| f'<span class="tag-en" onclick=\'copyToClipboard({js_arg})\'>{safe_tag_text}</span>' | |
| ) | |
| if show_translation_in_list and i < len(translations_list) and translations_list[i]: | |
| tag_display_html += f'<span class="tag-zh">({escape(str(translations_list[i]))})</span>' | |
| html += f"<div>{tag_display_html}</div>" | |
| if show_scores and isinstance(score, (int, float)): | |
| html += f'<span class="tag-score">{score:.3f}</span>' | |
| html += "</div>" | |
| html += "</div>" | |
| return html | |
| def generate_summary_text_content( | |
| current_res, | |
| current_translations_dict, | |
| s_gen, | |
| s_char, | |
| s_ip, | |
| s_sep_type, | |
| s_show_zh, | |
| ): | |
| if not current_res: | |
| return "请先分析图像或选择要汇总的标签类别。" | |
| summary_parts = [] | |
| separators = {"逗号": ", ", "换行": "\n", "空格": " "} | |
| separator = separators.get(s_sep_type, ", ") | |
| categories_to_summarize = [] | |
| if s_gen: | |
| categories_to_summarize.append("general") | |
| if s_char: | |
| categories_to_summarize.append("characters") | |
| if s_ip: | |
| categories_to_summarize.append("ips") | |
| if not categories_to_summarize: | |
| return "请至少选择一个标签类别进行汇总。" | |
| for cat_key in categories_to_summarize: | |
| if current_res.get(cat_key): | |
| tags_to_join = [] | |
| cat_tags_en = list(current_res[cat_key].keys()) | |
| cat_translations = current_translations_dict.get(cat_key, []) | |
| for i, en_tag in enumerate(cat_tags_en): | |
| if s_show_zh and i < len(cat_translations) and cat_translations[i]: | |
| tags_to_join.append(f"{en_tag}/*{cat_translations[i]}*/") | |
| else: | |
| tags_to_join.append(en_tag) | |
| if tags_to_join: | |
| summary_parts.append(separator.join(tags_to_join)) | |
| joiner = "\n\n" if separator != "\n" and len(summary_parts) > 1 else separator if separator == "\n" else " " | |
| final_summary = joiner.join(summary_parts) | |
| return final_summary if final_summary else "选定的类别中没有找到标签。" | |
| def process_image_and_generate_outputs( | |
| image_path, | |
| g_th, | |
| c_th, | |
| s_scores, | |
| s_gen, | |
| s_char, | |
| s_ip, | |
| s_sep, | |
| s_zh_in_sum, | |
| ): | |
| if image_path is None: | |
| yield ( | |
| gr.update(interactive=True, value="🚀 开始分析"), | |
| gr.update(visible=True, value="❌ 请先上传图片。"), | |
| "", | |
| "", | |
| "", | |
| "", | |
| {}, | |
| {}, | |
| {}, | |
| {}, | |
| ) | |
| return | |
| if tagger_instance is None: | |
| yield ( | |
| gr.update(interactive=True, value="🚀 开始分析"), | |
| gr.update(visible=True, value="❌ 分析器未成功初始化,请检查控制台错误。"), | |
| "", | |
| "", | |
| "", | |
| "", | |
| {}, | |
| {}, | |
| {}, | |
| {}, | |
| ) | |
| return | |
| yield ( | |
| gr.update(interactive=False, value="🔄 处理中..."), | |
| gr.update(visible=True, value="🔄 正在校验并分析图像,请稍候..."), | |
| gr.HTML(value="<p>分析中...</p>"), | |
| gr.HTML(value="<p>分析中...</p>"), | |
| gr.HTML(value="<p>分析中...</p>"), | |
| gr.update(value="分析中,请稍候..."), | |
| {}, | |
| {}, | |
| {}, | |
| {}, | |
| ) | |
| try: | |
| img = validate_and_open_image(image_path) | |
| res, tag_categories_original_order, meta = predict_tags_with_shared_slot( | |
| img, | |
| g_th, | |
| c_th, | |
| ) | |
| all_tags_to_translate = [] | |
| for cat_key in ["general", "characters", "ips"]: | |
| all_tags_to_translate.extend(tag_categories_original_order.get(cat_key, [])) | |
| all_translations_flat = [] | |
| if all_tags_to_translate: | |
| try: | |
| all_translations_flat = translate_texts(all_tags_to_translate, src_lang="auto", tgt_lang="zh") | |
| except Exception as translate_error: | |
| print(f"⚠️ 标签翻译失败,将仅显示英文标签:{translate_error}") | |
| all_translations_flat = [""] * len(all_tags_to_translate) | |
| current_translations_dict = {} | |
| offset = 0 | |
| for cat_key in ["general", "characters", "ips"]: | |
| cat_original_tags = tag_categories_original_order.get(cat_key, []) | |
| num_tags_in_cat = len(cat_original_tags) | |
| if num_tags_in_cat > 0: | |
| current_translations_dict[cat_key] = all_translations_flat[offset: offset + num_tags_in_cat] | |
| offset += num_tags_in_cat | |
| else: | |
| current_translations_dict[cat_key] = [] | |
| general_html = format_tags_html( | |
| res.get("general", {}), | |
| current_translations_dict.get("general", []), | |
| "general", | |
| s_scores, | |
| True, | |
| ) | |
| char_html = format_tags_html( | |
| res.get("characters", {}), | |
| current_translations_dict.get("characters", []), | |
| "characters", | |
| s_scores, | |
| True, | |
| ) | |
| ip_html = format_tags_html( | |
| res.get("ips", {}), | |
| current_translations_dict.get("ips", []), | |
| "ips", | |
| s_scores, | |
| True, | |
| ) | |
| summary_text = generate_summary_text_content( | |
| res, | |
| current_translations_dict, | |
| s_gen, | |
| s_char, | |
| s_ip, | |
| s_sep, | |
| s_zh_in_sum, | |
| ) | |
| yield ( | |
| gr.update(interactive=True, value="🚀 开始分析"), | |
| gr.update(visible=True, value="✅ 分析完成!"), | |
| general_html, | |
| char_html, | |
| ip_html, | |
| gr.update(value=summary_text), | |
| res, | |
| current_translations_dict, | |
| tag_categories_original_order, | |
| meta, | |
| ) | |
| except ImageValidationError as e: | |
| yield ( | |
| gr.update(interactive=True, value="🚀 开始分析"), | |
| gr.update(visible=True, value=f"❌ 上传图片未通过安全校验:{str(e)}"), | |
| "<p>图片已被安全策略拒绝</p>", | |
| "<p>图片已被安全策略拒绝</p>", | |
| "<p>图片已被安全策略拒绝</p>", | |
| gr.update(value=f"错误: {str(e)}", placeholder="上传图片未通过安全校验..."), | |
| {}, | |
| {}, | |
| {}, | |
| {}, | |
| ) | |
| except Exception as e: | |
| import traceback | |
| tb_str = traceback.format_exc() | |
| print(f"处理时发生错误: {e}\n{tb_str}") | |
| yield ( | |
| gr.update(interactive=True, value="🚀 开始分析"), | |
| gr.update(visible=True, value=f"❌ 处理失败: {str(e)}"), | |
| "<p>处理出错</p>", | |
| "<p>处理出错</p>", | |
| "<p>处理出错</p>", | |
| gr.update(value=f"错误: {str(e)}", placeholder="分析失败..."), | |
| {}, | |
| {}, | |
| {}, | |
| {}, | |
| ) | |
| def update_summary_display( | |
| s_gen, | |
| s_char, | |
| s_ip, | |
| s_sep, | |
| s_zh_in_sum, | |
| current_res_from_state, | |
| current_translations_from_state, | |
| ): | |
| if not current_res_from_state: | |
| return gr.update(placeholder="请先完成一次图像分析以生成汇总。", value="") | |
| new_summary_text = generate_summary_text_content( | |
| current_res_from_state, | |
| current_translations_from_state, | |
| s_gen, | |
| s_char, | |
| s_ip, | |
| s_sep, | |
| s_zh_in_sum, | |
| ) | |
| return gr.update(value=new_summary_text) | |
| btn.click( | |
| process_image_and_generate_outputs, | |
| inputs=[ | |
| img_in, | |
| gen_slider, | |
| char_slider, | |
| show_tag_scores, | |
| sum_general, | |
| sum_char, | |
| sum_ip, | |
| sum_sep, | |
| sum_show_zh, | |
| ], | |
| outputs=[ | |
| btn, | |
| processing_info, | |
| out_general, | |
| out_char, | |
| out_ip, | |
| out_summary, | |
| state_res, | |
| state_translations_dict, | |
| state_tag_categories_for_translation, | |
| out_meta, | |
| ], | |
| ) | |
| summary_controls = [sum_general, sum_char, sum_ip, sum_sep, sum_show_zh] | |
| for ctrl in summary_controls: | |
| ctrl.change( | |
| fn=update_summary_display, | |
| inputs=summary_controls + [state_res, state_translations_dict], | |
| outputs=[out_summary], | |
| ) | |
| def execute_image_inversion_job( | |
| payload: ImageInversionJobRequest, | |
| input_image: Image.Image, | |
| job_id: str, | |
| ) -> dict: | |
| """复用现有 Tagger、翻译与展示格式生成自定义 API 结果。 | |
| Args: | |
| payload: 已通过公开 Schema 校验的标签分析参数。 | |
| input_image: 已由任务 API 安全抓取并解码的图片。 | |
| job_id: 当前异步任务标识,仅用于隔离任务上下文。 | |
| Returns: | |
| 包含现有前端所需六个展示字段的 JSON 对象。 | |
| """ | |
| del job_id | |
| if tagger_instance is None: | |
| raise RuntimeError("标签分析器尚未成功初始化。") | |
| # ImageInversionJobAPI 已持有 INFERENCE_SLOT,此处不能重复加锁。 | |
| result, tag_order, metadata = tagger_instance.predict( | |
| input_image, | |
| payload.general_threshold, | |
| payload.character_threshold, | |
| ) | |
| all_tags: list[str] = [] | |
| for category_name in ("general", "characters", "ips"): | |
| all_tags.extend(tag_order.get(category_name, [])) | |
| translated_tags: list[str] = [] | |
| if all_tags: | |
| try: | |
| translated_tags = translate_texts( | |
| all_tags, | |
| src_lang="auto", | |
| tgt_lang="zh", | |
| ) | |
| except Exception as exc: | |
| print(f"标签翻译失败,将仅返回英文标签:{type(exc).__name__}") | |
| translated_tags = [""] * len(all_tags) | |
| translations: dict[str, list[str]] = {} | |
| offset = 0 | |
| for category_name in ("general", "characters", "ips"): | |
| category_tags = tag_order.get(category_name, []) | |
| tag_count = len(category_tags) | |
| translations[category_name] = translated_tags[offset : offset + tag_count] | |
| offset += tag_count | |
| separator_names = { | |
| "comma": "逗号", | |
| "newline": "换行", | |
| "space": "空格", | |
| } | |
| summary_text = generate_summary_text_content( | |
| result, | |
| translations, | |
| payload.show_general, | |
| payload.show_character, | |
| payload.show_ip, | |
| separator_names[payload.separator], | |
| payload.show_chinese, | |
| ) | |
| return { | |
| "status_markdown": "✅ 分析完成!", | |
| "general_tags_html": format_tags_html( | |
| result.get("general", {}), | |
| translations.get("general", []), | |
| "general", | |
| payload.show_confidence, | |
| True, | |
| ), | |
| "character_tags_html": format_tags_html( | |
| result.get("characters", {}), | |
| translations.get("characters", []), | |
| "characters", | |
| payload.show_confidence, | |
| True, | |
| ), | |
| "ip_tags_html": format_tags_html( | |
| result.get("ips", {}), | |
| translations.get("ips", []), | |
| "ips", | |
| payload.show_confidence, | |
| True, | |
| ), | |
| "summary_text": summary_text, | |
| "metadata": metadata, | |
| } | |
| IMAGE_INVERSION_JOB_SETTINGS = ImageInversionJobSettings.from_env() | |
| IMAGE_INVERSION_JOB_API = ImageInversionJobAPI( | |
| settings=IMAGE_INVERSION_JOB_SETTINGS, | |
| executor=execute_image_inversion_job, | |
| inference_slot=INFERENCE_SLOT, | |
| ) | |
| JOB_API_LIFESPAN = create_job_api_lifespan(IMAGE_INVERSION_JOB_API) | |
| if __name__ == "__main__": | |
| if tagger_instance is None: | |
| print("CRITICAL: Tagger 未能初始化,应用功能将受限。请检查之前的错误信息。") | |
| demo.queue(max_size=8).launch( | |
| server_name="0.0.0.0", | |
| server_port=7860, | |
| # 关闭 Gradio SSR,避免其 SvelteKit 中间件在 FastAPI 路由匹配前截获自定义 GET API。 | |
| ssr_mode=False, | |
| app_kwargs={"lifespan": JOB_API_LIFESPAN}, | |
| ) | |