From 375d86debbe994699dce8a3ed3fa52e80528633c Mon Sep 17 00:00:00 2001 From: spawner1145 Date: Sun, 19 Oct 2025 13:35:18 +0800 Subject: [PATCH] 111 --- README.md | 5 +- api_example/cancel.py | 42 +++++++++ api_example/generate.py | 86 +++++++++++++++++ backend_lsnet/__init__.py | 1 + backend_lsnet/api.py | 183 +++++++++++++++++++++++++++++++++++++ backend_lsnet/inference.py | 103 +++++++++++++++++++++ backend_lsnet/ui.py | 108 ++++++++++++++++++++++ install.py | 57 ++++++++++++ pyproject.toml | 2 +- scripts/app.py | 67 ++++++++++++++ 单独启动.bat | 2 + 11 files changed, 654 insertions(+), 2 deletions(-) create mode 100644 api_example/cancel.py create mode 100644 api_example/generate.py create mode 100644 backend_lsnet/__init__.py create mode 100644 backend_lsnet/api.py create mode 100644 backend_lsnet/inference.py create mode 100644 backend_lsnet/ui.py create mode 100644 install.py create mode 100644 scripts/app.py create mode 100644 单独启动.bat diff --git a/README.md b/README.md index bea18b1..26affa8 100644 --- a/README.md +++ b/README.md @@ -8,6 +8,8 @@ > “*Kaloscope*”(万花筒)致敬万花筒写轮眼,象征忍术(画风)复刻能力 +> 该插件支持comfyui插件,webui插件和单独启动三种方式运行 + ## 核心能力 基于 *LSNet* 技术核心,本工具聚焦两大核心场景: @@ -39,7 +41,6 @@ huggingface仓库地址:https://huggingface.co/heathcliff01/Kaloscope/tree/mai 目录结构示例: - ``` ComfyUI/ @@ -84,6 +85,8 @@ pip install triton-windows 2. 在画风分析工作流中,调用 LSNet 相关节点,即可触发 “画风分类” 或 “画风聚类” 功能 +> ps:目前版本该插件已经可以单独启动和作为webui插件启动,单独启动在项目根目录运行python -m scripts.app,模型路径为根目录下models/lsnet文件夹内,webui插件启动和comfyui差不多 + ### 使用示例 ![基础推理界面(含LSNet核心节点)](https://github.com/user-attachments/assets/28cc2820-ff5d-4290-8ac2-339763947e91) diff --git a/api_example/cancel.py b/api_example/cancel.py new file mode 100644 index 0000000..ce0534a --- /dev/null +++ b/api_example/cancel.py @@ -0,0 +1,42 @@ +import requests +import logging + +# 设置日志 +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + +# API 配置 +API_URL = "http://127.0.0.1:7871/lsnet/v1/cancel" +USERNAME = "user" # 替换为你的用户名,如果未启用认证可留空 +PASSWORD = "password" # 替换为你的密码,如果未启用认证可留空 + +def cancel_inference(): + """调用 /cancel 端点取消推理""" + try: + # 设置认证 + auth = None + if USERNAME and PASSWORD: + auth = (USERNAME, PASSWORD) + + # 发送请求 + response = requests.post(API_URL, auth=auth) + response.raise_for_status() + + result = response.json() + logger.info(f"Cancel result: {result['info']}") + return result['info'] + + except requests.exceptions.RequestException as e: + logger.error(f"API request failed: {str(e)}") + raise + except Exception as e: + logger.error(f"Cancel failed: {str(e)}") + raise + +# 示例使用 +if __name__ == "__main__": + try: + result = cancel_inference() + print(f"Cancel Result: {result}") + except Exception as e: + print(f"Error: {str(e)}") \ No newline at end of file diff --git a/api_example/generate.py b/api_example/generate.py new file mode 100644 index 0000000..6e74000 --- /dev/null +++ b/api_example/generate.py @@ -0,0 +1,86 @@ +import requests +import base64 +import json +import os +from PIL import Image +from io import BytesIO +import logging + +# 设置日志 +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + +# API 配置 +API_URL = "http://127.0.0.1:7871/lsnet/v1/infer" +USERNAME = "user" # 替换为你的用户名,如果未启用认证可留空 +PASSWORD = "password" # 替换为你的密码,如果未启用认证可留空 +OUTPUT_DIR = "outputs" +os.makedirs(OUTPUT_DIR, exist_ok=True) + +def encode_image_to_base64(image_path: str) -> str: + """将图片编码为 Base64 字符串""" + try: + with Image.open(image_path) as img: + img = img.convert("RGB") + buffered = BytesIO() + img.save(buffered, format="PNG") + return base64.b64encode(buffered.getvalue()).decode("utf-8") + except Exception as e: + logger.error(f"Failed to encode image {image_path}: {str(e)}") + raise + +def perform_inference(image_path: str, model_name='Kaloscope', **kwargs): + """调用 /infer 端点进行推理""" + try: + input_image_base64 = encode_image_to_base64(image_path) + + # 准备请求数据 + data = { + "input_image": input_image_base64, + "model_name": model_name, + **kwargs + } + + # 设置认证 + auth = None + if USERNAME and PASSWORD: + auth = (USERNAME, PASSWORD) + + # 发送请求 + response = requests.post(API_URL, json=data, auth=auth) + response.raise_for_status() + + result = response.json() + logger.info(f"Inference completed: {result['info']}") + return result['results'] + + except requests.exceptions.RequestException as e: + logger.error(f"API request failed: {str(e)}") + raise + except Exception as e: + logger.error(f"Inference failed: {str(e)}") + raise + +# 示例使用 +if __name__ == "__main__": + # 示例参数,请根据你的模型调整 + image_path = "test_image.png" # 替换为你的测试图片路径 + model_name = "Kaloscope" # 替换为你的模型文件夹名 + + try: + results = perform_inference( + image_path=image_path, + model_name=model_name, + top_k=5 + ) + print("Inference Results:") + print(json.dumps(results, indent=2, ensure_ascii=False)) + + # 保存结果 + output_file = os.path.join(OUTPUT_DIR, "inference_result.json") + with open(output_file, 'w', encoding='utf-8') as f: + json.dump(results, f, indent=2, ensure_ascii=False) + print(f"Results saved to {output_file}") + + except Exception as e: + print(f"Error: {str(e)}") \ No newline at end of file diff --git a/backend_lsnet/__init__.py b/backend_lsnet/__init__.py new file mode 100644 index 0000000..f4d407a --- /dev/null +++ b/backend_lsnet/__init__.py @@ -0,0 +1 @@ +# Backend for LSNet Artist Inference \ No newline at end of file diff --git a/backend_lsnet/api.py b/backend_lsnet/api.py new file mode 100644 index 0000000..1b53305 --- /dev/null +++ b/backend_lsnet/api.py @@ -0,0 +1,183 @@ +import base64 +import logging +from typing import Callable +from threading import Lock +from secrets import compare_digest +from io import BytesIO +import asyncio +import concurrent.futures +import os +import glob + +from fastapi import FastAPI, Depends, HTTPException +from fastapi.security import HTTPBasic, HTTPBasicCredentials +from pydantic import BaseModel, Field +from PIL import Image +import numpy as np +from backend_lsnet.inference import process_image_from_pil + +try: + from modules import shared + from modules.call_queue import queue_lock as webui_queue_lock + IN_WEBUI = True +except ImportError: + IN_WEBUI = False + shared = type('Shared', (), {'cmd_opts': type('CmdOpts', (), {'api_auth': None})()})() + webui_queue_lock = None + +# 设置日志 +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + +def get_available_checkpoints(model_name): + """Get available checkpoint files for the model""" + models_dir = "models/lsnet" + model_dir = os.path.join(models_dir, model_name) + if os.path.exists(model_dir): + checkpoints = [] + for ext in ['*.pth', '*.ckpt', '*.safetensors']: + checkpoints.extend(glob.glob(os.path.join(model_dir, ext))) + return [os.path.basename(f) for f in checkpoints] + return [] + +def get_available_csv(model_name): + """Get available CSV files for the model""" + models_dir = "models/lsnet" + model_dir = os.path.join(models_dir, model_name) + if os.path.exists(model_dir): + csv_files = glob.glob(os.path.join(model_dir, "*.csv")) + return [os.path.basename(f) for f in csv_files] + return [] + +def get_checkpoint_path(model_name, checkpoint_name): + """Get full checkpoint path""" + models_dir = "models/lsnet" + return os.path.join(models_dir, model_name, checkpoint_name) + +class InferenceRequest(BaseModel): + input_image: str = Field(..., description="Input image as Base64 encoded string") + model_name: str = Field('Kaloscope', description="Model name (subfolder in models/lsnet/)") + device: str = Field('cuda', description="Device to use") + top_k: int = Field(5, ge=1, le=20, description="Number of top predictions") + threshold: float = Field(0.0, ge=0.0, le=1.0, description="Probability threshold") + +class InferenceResponse(BaseModel): + results: dict = Field(..., description="Inference results") + info: str = Field(..., description="Additional information") + +class CancelResponse(BaseModel): + info: str = Field(..., description="Cancel operation result") + +class Api: + def __init__(self, app: FastAPI, queue_lock: Lock = None, prefix: str = "/lsnet/v1"): + self.app = app + self.queue_lock = queue_lock or Lock() + self.prefix = prefix + self.credentials = {} + self.executor = concurrent.futures.ThreadPoolExecutor(max_workers=1) + + if IN_WEBUI and shared.cmd_opts.api_auth: + for auth in shared.cmd_opts.api_auth.split(","): + user, password = auth.split(":") + self.credentials[user] = password + + self.add_api_route( + "infer", + self.endpoint_infer, + methods=["POST"], + response_model=InferenceResponse, + summary="Perform artist style inference", + description="Classify or cluster an image using LSNet artist model." + ) + self.add_api_route( + "cancel", + self.endpoint_cancel, + methods=["POST"], + response_model=CancelResponse, + summary="Cancel the current inference task", + description="Terminates the ongoing inference task." + ) + + def auth(self, creds: HTTPBasicCredentials = Depends(HTTPBasic())): + if not self.credentials: + return True + if creds.username in self.credentials: + if compare_digest(creds.password, self.credentials[creds.username]): + return True + raise HTTPException( + status_code=401, + detail="Incorrect username or password", + headers={"WWW-Authenticate": "Basic"} + ) + + def add_api_route(self, path: str, endpoint: Callable, **kwargs): + path = f"{self.prefix}/{path}" if self.prefix else path + if self.credentials: + return self.app.add_api_route(path, endpoint, dependencies=[Depends(self.auth)], **kwargs) + return self.app.add_api_route(path, endpoint, **kwargs) + + def decode_base64_image(self, base64_str: str) -> Image.Image: + try: + img_data = base64.b64decode(base64_str, validate=True) + img = Image.open(BytesIO(img_data)).convert("RGB") + return img + except base64.binascii.Error: + raise HTTPException(400, "Invalid Base64 string format") + except Exception as e: + raise HTTPException(400, f"Failed to decode image: {str(e)}") + + async def run_inference(self, image, **kwargs): + """Run inference in a separate thread""" + loop = asyncio.get_event_loop() + try: + return await loop.run_in_executor(self.executor, lambda: process_image_from_pil(image, **kwargs)) + except Exception as e: + logger.error(f"Inference execution failed: {str(e)}") + raise + + async def endpoint_infer(self, req: InferenceRequest): + logger.info(f"Received inference request: model_name={req.model_name}") + try: + with self.queue_lock: + input_image = self.decode_base64_image(req.input_image) + + checkpoints = get_available_checkpoints(req.model_name) + if not checkpoints: + raise HTTPException(400, f"No checkpoints found for model {req.model_name}") + checkpoint_name = checkpoints[0] # use first available + checkpoint = get_checkpoint_path(req.model_name, checkpoint_name) + if not os.path.exists(checkpoint): + raise HTTPException(400, f"Checkpoint not found: {checkpoint}") + + # Prepare inference arguments + csv_files = get_available_csv(req.model_name) + class_csv = None + if csv_files: + class_csv = os.path.join("models/lsnet", req.model_name, csv_files[0]) # use first available + infer_args = { + "model": "lsnet_xl_artist", # 固定模型架构 + "checkpoint": checkpoint, + "mode": "classify", # default to classify + "device": req.device, + "top_k": req.top_k, + "threshold": req.threshold, + "class_csv": class_csv + } + + # Run inference + results = await self.run_inference(input_image, **infer_args) + + return InferenceResponse(results=results, info="Inference completed successfully") + except Exception as e: + logger.error(f"Inference failed: {str(e)}") + raise HTTPException(500, f"Inference failed: {str(e)}") + + async def endpoint_cancel(self): + # For simplicity, just return a message since inference is quick + return CancelResponse(info="No active inference to cancel") + +def on_app_started(_demux, app, _server): + """Called when the webui app starts""" + queue_lock = webui_queue_lock or Lock() + api = Api(app, queue_lock) + logger.info("LSNet API routes added to webui") \ No newline at end of file diff --git a/backend_lsnet/inference.py b/backend_lsnet/inference.py new file mode 100644 index 0000000..98b1d8c --- /dev/null +++ b/backend_lsnet/inference.py @@ -0,0 +1,103 @@ +import os +import json +import tempfile +from pathlib import Path +import torch +from PIL import Image +import numpy as np +from inference_artist import ( + get_args_parser, load_checkpoint_state, normalize_state_dict_keys, + resolve_num_classes, resolve_feature_dim, load_model, process_single_image, + load_class_mapping +) +from timm.data import resolve_data_config +from timm.data.transforms_factory import create_transform + +def process_image(image_path, model='lsnet_t_artist', checkpoint='', num_classes=None, feature_dim=None, mode='classify', class_csv=None, device='cuda', top_k=5, threshold=0.0): + """ + Process a single image for artist style inference. + + Args: + image_path (str): Path to the input image + model (str): Model architecture + checkpoint (str): Path to model checkpoint + num_classes (int): Number of classes + feature_dim (int): Feature dimension + mode (str): Inference mode ('classify', 'cluster', 'both') + class_csv (str): Path to class mapping CSV + device (str): Device to use + top_k (int): Number of top predictions + threshold (float): Probability threshold + + Returns: + dict: Inference results + """ + # Create temporary directory for output + with tempfile.TemporaryDirectory() as temp_dir: + output_dir = Path(temp_dir) / "output" + output_dir.mkdir(exist_ok=True) + + # Prepare arguments + args = get_args_parser().parse_args([ + '--model', model, + '--checkpoint', checkpoint, + '--input', image_path, + '--output', str(output_dir), + '--device', device, + '--top-k', str(top_k), + '--threshold', str(threshold), + '--mode', mode + ]) + + if num_classes is not None: + args.num_classes = num_classes + if feature_dim is not None: + args.feature_dim = feature_dim + if class_csv is not None: + args.class_csv = class_csv + + # Load checkpoint and state + state_dict = load_checkpoint_state(checkpoint) + state_dict = normalize_state_dict_keys(state_dict) + + # Load class mapping + class_mapping = load_class_mapping(class_csv) if class_csv else None + + # Resolve num_classes + args.num_classes = resolve_num_classes(num_classes, class_mapping, state_dict) + + # Resolve feature_dim + args.feature_dim = resolve_feature_dim(feature_dim, state_dict) + + # Load model + model_obj = load_model(args, state_dict) + + # Prepare transform + config = resolve_data_config({}, model=model_obj) + transform = create_transform(**config) + + # Process single image + results = process_single_image(args, model_obj, transform, class_mapping) + + return results + +def process_image_from_pil(image, **kwargs): + """ + Process a PIL image for artist style inference. + + Args: + image (PIL.Image): Input image + **kwargs: Other arguments for process_image + + Returns: + dict: Inference results + """ + with tempfile.NamedTemporaryFile(suffix='.png', delete=False) as temp_file: + image.save(temp_file.name) + try: + return process_image(temp_file.name, **kwargs) + finally: + try: + os.unlink(temp_file.name) + except OSError: + pass # Ignore if file is still in use \ No newline at end of file diff --git a/backend_lsnet/ui.py b/backend_lsnet/ui.py new file mode 100644 index 0000000..f6e77de --- /dev/null +++ b/backend_lsnet/ui.py @@ -0,0 +1,108 @@ +import gradio as gr +from backend_lsnet.inference import process_image_from_pil +import os +import json +import glob + +def get_available_models(): + """Get available model folders from models/lsnet/""" + models_dir = "models/lsnet" + if os.path.exists(models_dir): + subdirs = [d for d in os.listdir(models_dir) if os.path.isdir(os.path.join(models_dir, d))] + if subdirs: + return subdirs + +def get_available_checkpoints(model_name): + """Get available checkpoint files for the model""" + models_dir = "models/lsnet" + model_dir = os.path.join(models_dir, model_name) + if os.path.exists(model_dir): + checkpoints = [] + for ext in ['*.pth', '*.ckpt', '*.safetensors']: + checkpoints.extend(glob.glob(os.path.join(model_dir, ext))) + return [os.path.basename(f) for f in checkpoints] + return [] + +def get_available_csv(model_name): + """Get available CSV files for the model""" + models_dir = "models/lsnet" + model_dir = os.path.join(models_dir, model_name) + if os.path.exists(model_dir): + csv_files = glob.glob(os.path.join(model_dir, "*.csv")) + return [os.path.basename(f) for f in csv_files] + return [] + +def get_checkpoint_path(model_name, checkpoint_name): + """Get full checkpoint path""" + models_dir = "models/lsnet" + return os.path.join(models_dir, model_name, checkpoint_name) + +def create_ui(): + css = """ + .contain-image img { + object-fit: contain !important; + width: 100% !important; + height: 100% !important; + background: #222; + } + """ + + block = gr.Blocks(css=css) + with block: + gr.Markdown('# LSNet Artist Inference') + with gr.Tabs(): + with gr.TabItem("Inference"): + with gr.Row(): + with gr.Column(): + input_image = gr.Image(sources='upload', type="pil", label="Input Image", height=320, elem_classes="contain-image") + model = gr.Dropdown( + choices=get_available_models(), + label="Model Folder", value='Kaloscope' + ) + device = gr.Dropdown(['cuda', 'cpu'], label="Device", value='cuda') + top_k = gr.Slider(label="Top K", minimum=1, maximum=20, value=5, step=1) + threshold = gr.Slider(label="Threshold", minimum=0.0, maximum=1.0, value=0.0, step=0.01) + infer_button = gr.Button(value="Infer") + with gr.Column(): + tag_string = gr.Textbox(label="Formatted Tags", lines=3, interactive=False) + result_json = gr.Textbox(label="JSON Results", lines=15, interactive=False) + error_message = gr.Markdown("", visible=False) + + def infer(image, model, device, top_k, threshold): + if image is None: + return "Please upload an image.", "", gr.update(visible=True) + checkpoints = get_available_checkpoints(model) + if not checkpoints: + return f"No checkpoints found for model {model}.", "", gr.update(visible=True) + checkpoint_name = checkpoints[0] # use first available + checkpoint = get_checkpoint_path(model, checkpoint_name) + if not os.path.exists(checkpoint): + return f"Checkpoint not found: {checkpoint}", "", gr.update(visible=True) + try: + csv_files = get_available_csv(model) + class_csv = None + if csv_files: + class_csv = os.path.join("models/lsnet", model, csv_files[0]) # use first available + kwargs = { + 'model': 'lsnet_xl_artist', # 固定模型架构 + 'checkpoint': checkpoint, + 'mode': 'classify', # default to classify + 'device': device, + 'top_k': top_k, + 'threshold': threshold, + 'class_csv': class_csv + } + results = process_image_from_pil(image, **kwargs) + tag_string = ",".join([r['class_name'] for r in results.get('classification', [])]) + json_str = json.dumps({r['class_name']: r['probability'] for r in results.get('classification', [])}, ensure_ascii=False) + return tag_string, json_str, gr.update(visible=False) + except Exception as e: + return str(e), "", gr.update(visible=True) + + infer_button.click( + infer, + inputs=[input_image, model, device, top_k, threshold], + outputs=[tag_string, result_json, error_message] + ) + + return block \ No newline at end of file diff --git a/install.py b/install.py new file mode 100644 index 0000000..01cea21 --- /dev/null +++ b/install.py @@ -0,0 +1,57 @@ +import launch +import importlib +from packaging.version import Version +from packaging.requirements import Requirement +import platform + +def is_installed(pip_package): + """ + Check if a package is installed and meets version requirements specified in pip-style format. + + Args: + pip_package (str): Package name in pip-style format (e.g., "numpy>=1.22.0"). + + Returns: + bool: True if the package is installed and meets the version requirement, False otherwise. + """ + try: + # Parse the pip-style package name and version constraints + requirement = Requirement(pip_package) + package_name = requirement.name + specifier = requirement.specifier # e.g., >=1.22.0 + + # Check if the package is installed + dist = importlib.metadata.distribution(package_name) + installed_version = Version(dist.version) + + # Check version constraints + if specifier.contains(installed_version): + return True + else: + print(f"Installed version of {package_name} ({installed_version}) does not satisfy the requirement ({specifier}).") + return False + except importlib.metadata.PackageNotFoundError: + print(f"Package {pip_package} is not installed.") + return False + + +# Read requirements from requirements.txt +with open("requirements.txt", "r") as f: + requirements = [line.strip() for line in f if line.strip() and not line.startswith("#") and not line.startswith(";")] + +# Add webui packages +webui_packages = [ + "gradio", + "fastapi", + "uvicorn", + "python-multipart" +] +requirements.extend(webui_packages) + +# Add platform-specific packages +if platform.system() == "Windows": + requirements.append("triton-windows") + +for req in requirements: + if not is_installed(req): + launch.run_pip(f"install {req}", f"sd-webui-lsnet requirement: {req}") \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 5004660..7233fc1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "lsnet" description = "" -version = "1.0.0" +version = "1.0.1" license = {file = "LICENSE"} # classifiers = [ # # For OS-independent nodes (works on all operating systems) diff --git a/scripts/app.py b/scripts/app.py new file mode 100644 index 0000000..d018e92 --- /dev/null +++ b/scripts/app.py @@ -0,0 +1,67 @@ +import logging +import threading +from threading import Lock +from fastapi import FastAPI +from backend_lsnet.inference import * +from backend_lsnet.ui import * +from backend_lsnet.api import Api +import uvicorn +import gradio as gr +import os +import argparse + +logging.basicConfig(level=logging.INFO) + +def parse_args(): + parser = argparse.ArgumentParser(description="LSNet Artist Inference WebUI") + parser.add_argument("--host", type=str, default="127.0.0.1", help="Server host") + parser.add_argument("--port", type=int, default=7860, help="Server port") + return parser.parse_args() + +try: + from modules import script_callbacks, shared + IN_WEBUI = True +except ImportError: + IN_WEBUI = False + shared = type('Shared', (), {'opts': type('Opts', (), { + 'outdir_samples': '', + 'outdir_txt2img_samples': '', + 'outdir_img2img_samples': '' + })})() + +if IN_WEBUI: + from backend_lsnet.api import on_app_started + def on_ui_tabs(): + block = create_ui() + return [(block, "LSNet Artist", "lsnet_tab")] + script_callbacks.on_ui_tabs(on_ui_tabs) + script_callbacks.on_app_started(on_app_started) +else: + if __name__ == "__main__": + args = parse_args() + # Create models directory + os.makedirs("models/lsnet", exist_ok=True) + + app = FastAPI(docs_url="/docs", openapi_url="/openapi.json") + queue_lock = Lock() + api = Api(app, queue_lock, prefix="/lsnet/v1") + logging.info("API 路由已挂载到 FastAPI 实例") + + block = create_ui() + logging.info("Gradio 界面已创建") + + # Mount Gradio app to FastAPI + app = gr.mount_gradio_app(app, block, path="") + + print(f"应用启动在:http://{args.host}:{args.port}") + print(f"API 文档:http://{args.host}:{args.port}/docs") + + try: + uvicorn.run( + app, + host=args.host, + port=args.port, + log_level="info" + ) + except Exception as e: + logging.error(f"启动失败: {str(e)}") \ No newline at end of file diff --git a/单独启动.bat b/单独启动.bat new file mode 100644 index 0000000..1938e1a --- /dev/null +++ b/单独启动.bat @@ -0,0 +1,2 @@ +python -m scripts.app +pause \ No newline at end of file